trellis_sparse_mrp_mlx/trellis_sparse_mlx/dit.py
John 384e87fca1 DiT blocks: RoPE, per-head RMS norm, AdaLN modulation (23 tests total)
Unlocks all four Pixal3D flow checkpoints (~20GB) at once - they are pure transformers
with no sparse conv. LATO.2 has flow models too (vertex_structured_flow, topo_flow), so
these belong in the shared core rather than either port.

Three details taken from upstream rather than assumed, each silent when wrong:
- norm1/norm3 are NON-affine but norm2 IS affine in the modulated cross block. There is
  an explicit test asserting that asymmetry.
- MultiHeadRMSNorm is written upstream as F.normalize(x)*gamma*sqrt(dim). F.normalize is
  L2, and the sqrt(d) turns it into RMS - implemented directly as RMS and verified equal
  to the upstream formulation to 9.5e-7.
- RoPE phases are NOT derived: Pixal3D ships rope_phases as a stored tensor, so they are
  passed in. Tested that the rotation preserves per-pair norms and is not a no-op.

Also tested: gates at zero make the block an identity on its residual branches.
2026-08-02 11:33:04 +10:00

241 lines
8.6 KiB
Python

"""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]