"""Modulated (DiT-style) transformer pieces used by the flow models in this family. Pixal3D's four flow checkpoints (~20GB — the bulk of its download) are pure transformers with no sparse convolution at all, built from these blocks. LATO.2 has flow models too (`vertex_structured_flow`, `topo_flow`), so they live in the shared core rather than in either port. Details that are silent when wrong, taken from upstream rather than assumed: * In `ModulatedTransformerCrossBlock`, `norm1` and `norm3` are NON-affine but `norm2` IS affine (`elementwise_affine=True`). Same shapes either way. * `MultiHeadRMSNorm` is written upstream as `F.normalize(x, dim=-1) * gamma * sqrt(dim)`. `F.normalize` is L2 (x/‖x‖), and multiplying by √d turns it into RMS norm — x/‖x‖·√d == x/rms(x). Implemented directly as RMS to avoid the double indirection. * RoPE phases are NOT derived here. Pixal3D ships `rope_phases` as a stored tensor in the checkpoint, so phases are passed in. """ from __future__ import annotations import math from typing import Optional, Tuple import mlx.core as mx import mlx.nn as nn class MultiHeadRMSNorm(nn.Module): """Per-head RMS norm over head_dim, with a [heads, dim] gain.""" def __init__(self, dim: int, heads: int): super().__init__() self.gamma = mx.ones((heads, dim)) self.eps = 1e-12 def __call__(self, x: mx.array) -> mx.array: # x: [..., heads, dim] dt = x.dtype f = x.astype(mx.float32) rms = mx.sqrt(mx.mean(f * f, axis=-1, keepdims=True) + self.eps) return ((f / rms) * self.gamma).astype(dt) def apply_rope(q: mx.array, k: mx.array, phases: mx.array) -> Tuple[mx.array, mx.array]: """Rotate q/k by precomputed `phases`. `phases` is [..., head_dim/2] (or broadcastable) giving the angle per rotary pair. Pairs are (even, odd) along the last axis, matching upstream's interleaved layout. """ cos, sin = mx.cos(phases), mx.sin(phases) def rot(t: mx.array) -> mx.array: dt = t.dtype f = t.astype(mx.float32) a, b = f[..., 0::2], f[..., 1::2] ra = a * cos - b * sin rb = a * sin + b * cos out = mx.stack([ra, rb], axis=-1).reshape(f.shape) return out.astype(dt) return rot(q), rot(k) class TimestepEmbedder(nn.Module): """Sinusoidal timestep embedding -> 2-layer MLP, as in DiT.""" def __init__(self, hidden_size: int, frequency_embedding_size: int = 256): super().__init__() self.frequency_embedding_size = frequency_embedding_size self.mlp_0 = nn.Linear(frequency_embedding_size, hidden_size) self.mlp_2 = nn.Linear(hidden_size, hidden_size) @staticmethod def timestep_embedding(t: mx.array, dim: int, max_period: int = 10000) -> mx.array: half = dim // 2 freqs = mx.exp( -math.log(max_period) * mx.arange(half, dtype=mx.float32) / half ) args = t.astype(mx.float32)[:, None] * freqs[None] emb = mx.concatenate([mx.cos(args), mx.sin(args)], axis=-1) if dim % 2: emb = mx.concatenate([emb, mx.zeros((emb.shape[0], 1))], axis=-1) return emb def __call__(self, t: mx.array) -> mx.array: e = self.timestep_embedding(t, self.frequency_embedding_size) return self.mlp_2(nn.silu(self.mlp_0(e))) class _LN(nn.Module): """LayerNorm computed in fp32, optional affine — upstream's LayerNorm32.""" def __init__(self, dim: int, affine: bool, eps: float = 1e-6): super().__init__() self.eps = eps self.affine = affine if affine: self.weight = mx.ones((dim,)) self.bias = mx.zeros((dim,)) def __call__(self, x: mx.array) -> mx.array: dt = x.dtype f = x.astype(mx.float32) f = (f - mx.mean(f, -1, keepdims=True)) * mx.rsqrt( mx.var(f, -1, keepdims=True) + self.eps ) if self.affine: f = f * self.weight + self.bias return f.astype(dt) class DiTAttention(nn.Module): """Self or cross attention with optional RoPE and per-head q/k RMS norm.""" def __init__( self, channels: int, num_heads: int, ctx_channels: Optional[int] = None, attn_type: str = "self", qkv_bias: bool = True, use_rope: bool = False, qk_rms_norm: bool = False, ): super().__init__() self.channels, self.num_heads = channels, num_heads self.head_dim = channels // num_heads self.scale = self.head_dim**-0.5 self._type = attn_type self.use_rope = use_rope self.qk_rms_norm = qk_rms_norm ctx = ctx_channels if ctx_channels is not None else channels if attn_type == "self": self.to_qkv = nn.Linear(channels, channels * 3, bias=qkv_bias) else: self.to_q = nn.Linear(channels, channels, bias=qkv_bias) self.to_kv = nn.Linear(ctx, channels * 2, bias=qkv_bias) if qk_rms_norm: self.q_rms_norm = MultiHeadRMSNorm(self.head_dim, num_heads) self.k_rms_norm = MultiHeadRMSNorm(self.head_dim, num_heads) self.to_out = nn.Linear(channels, channels) def __call__( self, x: mx.array, context: Optional[mx.array] = None, phases: Optional[mx.array] = None, ) -> mx.array: b, n, _ = x.shape h, d = self.num_heads, self.head_dim if self._type == "self": qkv = self.to_qkv(x).reshape(b, n, 3, h, d) q, k, v = qkv[:, :, 0], qkv[:, :, 1], qkv[:, :, 2] else: if context is None: raise ValueError("cross attention needs a context") m = context.shape[1] q = self.to_q(x).reshape(b, n, h, d) kv = self.to_kv(context).reshape(b, m, 2, h, d) k, v = kv[:, :, 0], kv[:, :, 1] if self.use_rope and phases is not None: q, k = apply_rope(q, k, phases) if self.qk_rms_norm: q, k = self.q_rms_norm(q), self.k_rms_norm(k) q = q.transpose(0, 2, 1, 3) k = k.transpose(0, 2, 1, 3) v = v.transpose(0, 2, 1, 3) o = mx.fast.scaled_dot_product_attention(q, k, v, scale=self.scale) return self.to_out(o.transpose(0, 2, 1, 3).reshape(b, n, self.channels)) class ModulatedTransformerCrossBlock(nn.Module): """AdaLN-modulated self-attn -> cross-attn -> FFN. Modulation is six chunks (shift/scale/gate for MSA and MLP). With `share_mod` the block holds its own `modulation` parameter that is ADDED to the incoming `mod`; otherwise it derives them via its own `adaLN_modulation`. Only the self-attention and MLP are modulated — the cross-attention path is not, which is why `norm2` is the affine one. """ def __init__( self, channels: int, ctx_channels: int, num_heads: int, mlp_ratio: float = 4.0, share_mod: bool = False, use_rope: bool = False, qk_rms_norm: bool = False, qk_rms_norm_cross: bool = False, ): super().__init__() self.share_mod = share_mod self.norm1 = _LN(channels, affine=False, eps=1e-6) self.norm2 = _LN(channels, affine=True, eps=1e-6) self.norm3 = _LN(channels, affine=False, eps=1e-6) self.self_attn = DiTAttention( channels, num_heads, use_rope=use_rope, qk_rms_norm=qk_rms_norm ) self.cross_attn = DiTAttention( channels, num_heads, ctx_channels=ctx_channels, attn_type="cross", qk_rms_norm=qk_rms_norm_cross, ) hidden = int(channels * mlp_ratio) self.mlp_0 = nn.Linear(channels, hidden) self.mlp_2 = nn.Linear(hidden, channels) if share_mod: self.modulation = mx.zeros((6 * channels,)) else: self.adaLN_modulation_1 = nn.Linear(channels, 6 * channels) def __call__( self, x: mx.array, mod: mx.array, context: mx.array, phases: Optional[mx.array] = None, ) -> mx.array: if self.share_mod: m = (self.modulation + mod).astype(mod.dtype) else: m = self.adaLN_modulation_1(nn.silu(mod)) c = m.shape[-1] // 6 sh_msa, sc_msa, g_msa, sh_mlp, sc_mlp, g_mlp = ( m[..., i * c : (i + 1) * c] for i in range(6) ) h = self.norm1(x) * (1 + sc_msa[:, None]) + sh_msa[:, None] x = x + self.self_attn(h, phases=phases) * g_msa[:, None] x = x + self.cross_attn(self.norm2(x), context) h = self.norm3(x) * (1 + sc_mlp[:, None]) + sh_mlp[:, None] h = self.mlp_2(nn.gelu_approx(self.mlp_0(h))) return x + h * g_mlp[:, None]