Transformers
nano-transformers / model.py
nebulette's picture
Upload 4 files
069dce3 verified
Raw History Blame
14.9 kB
import math
from safetensors.torch import load_model
import torch
import torch.nn as nn
import torch.nn.functional as F
SPRINT_NUM_F = 2
SPRINT_NUM_H = 2
TEXT_EMBED_DIM = 640
def adaln(
x: torch.Tensor,
scale: torch.Tensor,
shift: torch.Tensor,
eps: float = 1e-6
):
normalized = F.layer_norm(
x,
normalized_shape=(x.shape[-1],),
weight=None,
bias=None,
eps=eps,
)
return normalized * (1 + scale) + shift
def rope_2d(head_dim, height, width, device):
"""Split-half 2D RoPE over centered token coordinates (x frequencies, then y), as (1, N, 1, head_dim // 2, 2, 2) rotations."""
axis_dim = head_dim // 2
inv_freq = 1.0 / (10000.0 ** (torch.arange(0, axis_dim, 2, dtype=torch.float32, device=device) / axis_dim))
y = torch.arange(height, dtype=torch.float32, device=device) - (height - 1) / 2
x = torch.arange(width, dtype=torch.float32, device=device) - (width - 1) / 2
y, x = torch.meshgrid(y, x, indexing="ij")
angles = torch.cat([torch.outer(x.flatten(), inv_freq), torch.outer(y.flatten(), inv_freq)], dim=-1)
cos, sin = torch.cos(angles), torch.sin(angles)
return torch.stack([cos, -sin, sin, cos], dim=-1).view(1, height * width, 1, axis_dim, 2, 2)
def modulate(x, shift, scale):
return torch.addcmul(shift, x, 1 + scale)
def rope_split_half(x, pe):
"""Differentiable ck.apply_rope_split_half1: pair k is (x[k], x[k + head_dim // 2])."""
x_ = x.reshape(*x.shape[:-1], 2, -1).movedim(-2, -1).unsqueeze(-2).to(pe.dtype)
return (pe[..., 0] * x_[..., 0] + pe[..., 1] * x_[..., 1]).movedim(-1, -2).reshape(x.shape).type_as(x)
def norm_rope(q, k, q_norm, k_norm, pe):
return rope_split_half(q_norm(q), pe), rope_split_half(k_norm(k), pe)
class Embed(nn.Module):
def __init__(self, in_dim, hidden_size, norm=False, dtype=None, device=None):
super().__init__()
self.proj = nn.Linear(in_dim, hidden_size, bias=True, dtype=dtype, device=device)
self.norm = nn.RMSNorm(hidden_size, eps=1e-6, dtype=dtype, device=device) if norm else nn.Identity()
def forward(self, x):
return self.norm(self.proj(x))
class FeedForward(nn.Module):
def __init__(self, dim, hidden_dim, dtype=None, device=None):
super().__init__()
self.w12 = nn.Linear(dim, hidden_dim * 2, bias=False, dtype=dtype, device=device)
self.w3 = nn.Linear(hidden_dim, dim, bias=False, dtype=dtype, device=device)
def forward(self, x):
x1, x2 = self.w12(x).chunk(2, dim=-1)
return self.w3(F.silu(x1) * x2)
class Attention(nn.Module):
"""Self-attention over image tokens with 2D RoPE, optionally joint with (non-rotated) text keys and values."""
def __init__(self, dim, num_heads, cross_attention, dtype=None, device=None):
super().__init__()
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.qkv_x = nn.Linear(dim, dim * 3, bias=False, dtype=dtype, device=device)
self.kv_y = nn.Linear(dim, dim * 2, bias=False, dtype=dtype, device=device) if cross_attention else None
self.q_norm = nn.RMSNorm(self.head_dim, eps=1e-6, dtype=dtype, device=device)
self.k_norm = nn.RMSNorm(self.head_dim, eps=1e-6, dtype=dtype, device=device)
self.proj = nn.Linear(dim, dim, bias=True, dtype=dtype, device=device)
def forward(self, x, txt, pe, txt_bias, transformer_options={}):
b, n, c = x.shape
q, k, v = self.qkv_x(x).view(b, n, 3, self.num_heads, self.head_dim).unbind(2)
q, k = norm_rope(q, k, self.q_norm, self.k_norm, pe)
mask = None
if self.kv_y is not None:
ky, vy = self.kv_y(txt).view(b, -1, 2, self.num_heads, self.head_dim).unbind(2)
k = torch.cat([k, self.k_norm(ky)], dim=1)
v = torch.cat([v, vy], dim=1)
# Image keys are always visible; text keys carry log(emphasis weight), -inf for padding.
mask = F.pad(txt_bias, (n, 0))[:, None, None, :]
# x = optimized_attention(q.reshape(b, n, c), k.reshape(b, -1, c), v.reshape(b, -1, c), self.num_heads, mask=mask, transformer_options=transformer_options)
q = q.transpose(1, 2)
k = k.transpose(1, 2)
v = v.transpose(1, 2)
# Mask from ldm/modules/attention, attention_pytorch() function.
if mask is not None:
# add a batch dimension if there isn't already one
if mask.ndim == 2:
mask = mask.unsqueeze(0)
# add a heads dimension if there isn't already one
if mask.ndim == 3:
mask = mask.unsqueeze(1)
x = F.scaled_dot_product_attention(q, k, v, attn_mask=mask)
x = x.transpose(1, 2).reshape(b, n, c)
return self.proj(x)
class DiTBlock(nn.Module):
"""Image block; its adaLN modulation is computed by the caller (adaLN-single)."""
def __init__(self, hidden_size, num_heads, mlp_hidden, cross_attention, dtype=None, device=None):
super().__init__()
self.norm1 = nn.RMSNorm(hidden_size, eps=1e-6, dtype=dtype, device=device)
self.attn = Attention(hidden_size, num_heads, cross_attention, dtype=dtype, device=device)
self.norm2 = nn.RMSNorm(hidden_size, eps=1e-6, dtype=dtype, device=device)
self.mlp = FeedForward(hidden_size, mlp_hidden, dtype=dtype, device=device)
def forward(self, x, txt, pe, mod, txt_bias, transformer_options={}):
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = mod.chunk(6, dim=-1)
x = torch.addcmul(x, gate_msa, self.attn(modulate(self.norm1(x), shift_msa, scale_msa), txt, pe, txt_bias, transformer_options=transformer_options))
return torch.addcmul(x, gate_mlp, self.mlp(modulate(self.norm2(x), shift_mlp, scale_mlp)))
class TextRefineAttention(nn.Module):
def __init__(self, dim, num_heads, dtype=None, device=None):
super().__init__()
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.qkv = nn.Linear(dim, dim * 3, bias=False, dtype=dtype, device=device)
self.q_norm = nn.RMSNorm(self.head_dim, eps=1e-6, dtype=dtype, device=device)
self.k_norm = nn.RMSNorm(self.head_dim, eps=1e-6, dtype=dtype, device=device)
self.proj = nn.Linear(dim, dim, bias=True, dtype=dtype, device=device)
def forward(self, x, mask, transformer_options={}):
b, n, c = x.shape
q, k, v = self.qkv(x).view(b, n, 3, self.num_heads, self.head_dim).unbind(2)
q, k = self.q_norm(q), self.k_norm(k)
# SDPA expects (B, H, N, D)
q = q.transpose(1, 2)
k = k.transpose(1, 2)
v = v.transpose(1, 2)
# x = optimized_attention(self.q_norm(q).reshape(b, n, c), self.k_norm(k).reshape(b, n, c), v.reshape(b, n, c), self.num_heads, mask=mask, transformer_options=transformer_options)
x = F.scaled_dot_product_attention(q, k, v)
x = x.transpose(1, 2).reshape(b, n, c)
return self.proj(x)
class TextRefineBlock(nn.Module):
def __init__(self, hidden_size, num_heads, mlp_hidden, dtype=None, device=None):
super().__init__()
self.norm1 = nn.RMSNorm(hidden_size, eps=1e-6, dtype=dtype, device=device)
self.attn = TextRefineAttention(hidden_size, num_heads, dtype=dtype, device=device)
self.norm2 = nn.RMSNorm(hidden_size, eps=1e-6, dtype=dtype, device=device)
self.mlp = FeedForward(hidden_size, mlp_hidden, dtype=dtype, device=device)
self.adaLN_modulation = nn.Sequential(nn.Linear(hidden_size, 6 * hidden_size, bias=True, dtype=dtype, device=device))
def forward(self, x, c, mask, keep, transformer_options={}):
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.adaLN_modulation(c).chunk(6, dim=-1)
x = torch.addcmul(x, gate_msa, self.attn(modulate(self.norm1(x), shift_msa, scale_msa), mask, transformer_options=transformer_options))
x = torch.addcmul(x, gate_mlp, self.mlp(modulate(self.norm2(x), shift_mlp, scale_mlp)))
return x * keep
class FinalLayer(nn.Module):
def __init__(self, hidden_size, out_channels, dtype=None, device=None):
super().__init__()
self.norm_final = nn.LayerNorm(hidden_size, elementwise_affine=False, eps=1e-6, dtype=dtype, device=device)
self.adaLN_modulation = nn.Linear(hidden_size, 2 * hidden_size, bias=True, dtype=dtype, device=device)
self.linear = nn.Linear(hidden_size, out_channels, bias=True, dtype=dtype, device=device)
def forward(self, x, c):
shift, scale = self.adaLN_modulation(c).chunk(2, dim=-1)
return self.linear(adaln(x, scale, shift, self.norm_final.eps))
class TimestepEmbedder(nn.Module):
"""
Embeds scalar timesteps into vector representations.
"""
def __init__(self, hidden_size, frequency_embedding_size=256, output_size=None, dtype=None, device=None, max_period=10000):
super().__init__()
if output_size is None:
output_size = hidden_size
self.mlp = nn.Sequential(
nn.Linear(frequency_embedding_size, hidden_size, bias=True, dtype=dtype, device=device),
nn.SiLU(),
nn.Linear(hidden_size, output_size, bias=True, dtype=dtype, device=device),
)
self.frequency_embedding_size = frequency_embedding_size
self.max_period = max_period
def forward(self, t, dtype, **kwargs):
t_freq = self.timestep_embedding(t).to(dtype)
t_emb = self.mlp(t_freq)
return t_emb
def timestep_embedding(self, timesteps, repeat_only=False):
"""
Create sinusoidal timestep embeddings.
:param timesteps: a 1-D Tensor of N indices, one per batch element.
These may be fractional.
:return: an [N x dim] Tensor of positional embeddings.
"""
dim = self.frequency_embedding_size
if not repeat_only:
half = dim // 2
freqs = torch.exp(
-math.log(self.max_period) * torch.arange(start=0, end=half, dtype=torch.float32, device=timesteps.device) / half
)
args = timesteps[:, None].float() * freqs[None]
embedding = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
if dim % 2:
embedding = torch.cat([embedding, torch.zeros_like(embedding[:, :1])], dim=-1)
else:
# embedding = repeat(timesteps, 'b -> b d', d=dim)
embedding = timesteps[:, None].expand(-1, dim)
return embedding
class Transformer2DModel(nn.Module):
def __init__(self, in_channels=64, hidden_size=1536, num_heads=16, num_blocks=18, num_text_blocks=2, mlp_hidden=4096, dtype=None, device=None):
super().__init__()
self.dtype = dtype
self.in_channels = in_channels
self.num_heads = num_heads
self.head_dim = hidden_size // num_heads
self.s_embedder = Embed(in_channels, hidden_size, dtype=dtype, device=device)
self.t_embedder = TimestepEmbedder(hidden_size, dtype=dtype, device=device)
self.y_embedder = Embed(TEXT_EMBED_DIM, hidden_size, norm=True, dtype=dtype, device=device)
self.shared_encoder_adaLN = nn.Sequential(nn.Linear(hidden_size, 6 * hidden_size, bias=True, dtype=dtype, device=device))
self.encoder_adaLN_offsets = nn.ParameterList([nn.Parameter(torch.empty(6 * hidden_size, dtype=dtype, device=device)) for _ in range(num_blocks)])
self.y_pool_proj = nn.Linear(hidden_size, hidden_size, bias=False, dtype=dtype, device=device)
self.sprint_out_proj = nn.Linear(hidden_size, hidden_size, bias=False, dtype=dtype, device=device)
self.blocks = nn.ModuleList([DiTBlock(hidden_size, num_heads, mlp_hidden, i % 2 == 0, dtype=dtype, device=device) for i in range(num_blocks)])
self.text_refine_blocks = nn.ModuleList([TextRefineBlock(hidden_size, num_heads, mlp_hidden, dtype=dtype, device=device) for _ in range(num_text_blocks)])
self.final_layer = FinalLayer(hidden_size, in_channels, dtype=dtype, device=device)
def run_block(self, i, x, txt, pe, mod, txt_bias, transformer_options):
offset = self.encoder_adaLN_offsets[i].to(dtype=mod.dtype)
return self.blocks[i](x, txt, pe, mod + offset, txt_bias, transformer_options=transformer_options)
def sparse_path(self, x, txt, pe, mod, txt_bias, transformer_options):
g = x
for i in range(SPRINT_NUM_F, len(self.blocks) - SPRINT_NUM_H):
g = self.run_block(i, g, txt, pe, mod, txt_bias, transformer_options)
return x + self.sprint_out_proj(g - x)
@torch.no_grad()
def forward(self, x, timestep, context, token_weights, sparse_skip=None, transformer_options={}, **kwargs):
"""sparse_skip: optional per-row flags; flagged rows skip the sparse middle blocks."""
b, _, h, w = x.shape
keep = token_weights > 0
txt_bias = torch.log(token_weights.clamp(min=1e-4)).masked_fill(~keep, float("-inf"))
# Refine attention is a hard mask; a row with no text key keeps its first one and is zeroed after.
refine_keep = keep.clone()
refine_keep[:, 0] |= ~keep.any(dim=1)
refine_mask = torch.zeros_like(txt_bias).masked_fill(~refine_keep, float("-inf"))[:, None, None, :]
keep = keep.unsqueeze(-1).to(x.dtype)
t = self.t_embedder(timestep * 1000.0, x.dtype).unsqueeze(1)
txt = self.y_embedder(context)
time_condition = F.silu(t)
for block in self.text_refine_blocks:
txt = block(txt, time_condition, refine_mask, keep, transformer_options=transformer_options)
weights = token_weights.unsqueeze(-1)
pooled = (txt * weights).sum(dim=1) / weights.sum(dim=1).clamp(min=1.0)
condition = F.silu(t + self.y_pool_proj(pooled).unsqueeze(1))
mod = self.shared_encoder_adaLN(condition)
pe = rope_2d(self.head_dim, h, w, x.device)
s = self.s_embedder(x.flatten(2).transpose(1, 2))
for i in range(SPRINT_NUM_F):
s = self.run_block(i, s, txt, pe, mod, txt_bias, transformer_options)
if sparse_skip is None or not any(sparse_skip):
s = self.sparse_path(s, txt, pe, mod, txt_bias, transformer_options)
elif not all(sparse_skip):
rows = [i for i, skip in enumerate(sparse_skip) if not skip]
s[rows] = self.sparse_path(s[rows], txt[rows], pe, mod[rows], txt_bias[rows], transformer_options)
for i in range(len(self.blocks) - SPRINT_NUM_H, len(self.blocks)):
s = self.run_block(i, s, txt, pe, mod, txt_bias, transformer_options)
x0 = self.final_layer(s, condition).transpose(1, 2).reshape(b, self.in_channels, h, w)
return (x - x0) / timestep.view(-1, 1, 1, 1)
@staticmethod
def from_safetensors(path, device=None):
model = Transformer2DModel()
load_model(model, path)
return model.to(device)