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)