trellis_sparse_mrp_mlx/tests/test_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

156 lines
5.8 KiB
Python

"""DiT block tests against torch.
These layers all have torch counterparts, so unlike the submanifold conv they are checked
against upstream's real semantics. The upstream forward bodies are reproduced verbatim
from pixal3d/modules/{attention/modules.py,transformer/modulated.py}.
"""
import sys
from pathlib import Path
import mlx.core as mx
import numpy as np
import torch
import torch.nn.functional as F
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from trellis_sparse_mlx.dit import ( # noqa: E402
DiTAttention,
ModulatedTransformerCrossBlock,
MultiHeadRMSNorm,
TimestepEmbedder,
apply_rope,
)
def test_rms_norm_matches_upstream(dim=32, heads=4, n=17):
"""Upstream writes it as F.normalize(x)*gamma*sqrt(dim); we implement RMS directly."""
rng = np.random.default_rng(0)
x = rng.standard_normal((n, heads, dim)).astype(np.float32)
g = rng.standard_normal((heads, dim)).astype(np.float32)
m = MultiHeadRMSNorm(dim, heads)
m.gamma = mx.array(g)
got = np.asarray(m(mx.array(x)))
tx = torch.tensor(x)
want = (F.normalize(tx.float(), dim=-1) * torch.tensor(g) * dim**0.5).numpy()
err = np.abs(got - want).max()
assert err < 2e-4, f"rms norm err {err:.3g}"
return err
def test_rms_norm_is_scale_invariant():
"""RMS norm must remove input magnitude — catches a plain scale-by-gamma stub."""
m = MultiHeadRMSNorm(8, 2)
x = np.random.default_rng(1).standard_normal((5, 2, 8)).astype(np.float32)
a = np.asarray(m(mx.array(x)))
b = np.asarray(m(mx.array(x * 37.0)))
err = np.abs(a - b).max()
assert err < 1e-3, f"not scale-invariant: {err:.3g}"
return err
def test_rope_is_a_rotation():
"""RoPE must preserve the norm of each rotary pair."""
rng = np.random.default_rng(2)
q = rng.standard_normal((2, 6, 4, 16)).astype(np.float32)
ph = rng.uniform(0, 2 * np.pi, (2, 6, 4, 8)).astype(np.float32)
rq, _ = apply_rope(mx.array(q), mx.array(q), mx.array(ph))
rq = np.asarray(rq)
n0 = np.sqrt(q[..., 0::2] ** 2 + q[..., 1::2] ** 2)
n1 = np.sqrt(rq[..., 0::2] ** 2 + rq[..., 1::2] ** 2)
err = np.abs(n0 - n1).max()
assert err < 1e-4, f"rope changed pair norms by {err:.3g}"
# and it must actually rotate
assert np.abs(rq - q).max() > 1e-3, "rope was a no-op"
return err
def test_attention_matches_torch(ch=64, heads=8, n=13, b=2):
rng = np.random.default_rng(3)
x = rng.standard_normal((b, n, ch)).astype(np.float32)
wq = rng.standard_normal((ch * 3, ch)).astype(np.float32) * 0.05
bq = rng.standard_normal((ch * 3,)).astype(np.float32) * 0.05
wo = rng.standard_normal((ch, ch)).astype(np.float32) * 0.05
bo = rng.standard_normal((ch,)).astype(np.float32) * 0.05
a = DiTAttention(ch, heads)
a.to_qkv.weight, a.to_qkv.bias = mx.array(wq), mx.array(bq)
a.to_out.weight, a.to_out.bias = mx.array(wo), mx.array(bo)
got = np.asarray(a(mx.array(x)))
d = ch // heads
tx = torch.tensor(x)
qkv = F.linear(tx, torch.tensor(wq), torch.tensor(bq)).reshape(b, n, 3, heads, d)
q, k, v = qkv[:, :, 0], qkv[:, :, 1], qkv[:, :, 2]
o = F.scaled_dot_product_attention(
q.permute(0, 2, 1, 3), k.permute(0, 2, 1, 3), v.permute(0, 2, 1, 3)
)
o = o.permute(0, 2, 1, 3).reshape(b, n, ch)
want = F.linear(o, torch.tensor(wo), torch.tensor(bo)).detach().numpy()
err = np.abs(got - want).max()
assert err < 2e-4, f"attention err {err:.3g}"
return err
def test_modulation_gates_actually_gate(ch=32, heads=4, n=7, b=2):
"""gate=0 must make the block an identity on the modulated paths."""
blk = ModulatedTransformerCrossBlock(ch, ch, heads, share_mod=True)
x = mx.array(np.random.default_rng(4).standard_normal((b, n, ch)).astype(np.float32))
ctx = mx.array(np.zeros((b, 5, ch), dtype=np.float32))
# zero everything, then force the cross-attn output to zero via zero out-proj
blk.modulation = mx.zeros((6 * ch,))
blk.cross_attn.to_out.weight = mx.zeros_like(blk.cross_attn.to_out.weight)
blk.cross_attn.to_out.bias = mx.zeros_like(blk.cross_attn.to_out.bias)
mod = mx.zeros((b, 6 * ch))
out = np.asarray(blk(x, mod, ctx))
err = np.abs(out - np.asarray(x)).max()
assert err < 1e-4, f"gates did not zero the residual branches: {err:.3g}"
return err
def test_norm_affine_asymmetry():
"""norm1/norm3 non-affine, norm2 affine — silent if swapped."""
blk = ModulatedTransformerCrossBlock(16, 16, 2, share_mod=True)
assert not blk.norm1.affine, "norm1 should be non-affine"
assert blk.norm2.affine, "norm2 SHOULD be affine"
assert not blk.norm3.affine, "norm3 should be non-affine"
return 0.0
def test_timestep_embedder_shape():
te = TimestepEmbedder(64, 32)
out = te(mx.array(np.array([0.0, 0.5, 1.0], dtype=np.float32)))
assert out.shape == (3, 64), out.shape
assert np.isfinite(np.asarray(out)).all()
return 0.0
if __name__ == "__main__":
tests = [
("rms norm vs upstream", test_rms_norm_matches_upstream),
("rms scale invariance", test_rms_norm_is_scale_invariant),
("rope is a rotation", test_rope_is_a_rotation),
("dit attention vs torch", test_attention_matches_torch),
("modulation gates", test_modulation_gates_actually_gate),
("norm affine asymmetry", test_norm_affine_asymmetry),
("timestep embedder", test_timestep_embedder_shape),
]
failed = 0
for name, fn in tests:
try:
err = fn()
print(f" PASS {name:26s} (max err {err:.2e})")
except AssertionError as e:
print(f" FAIL {name:26s} {e}")
failed += 1
except Exception as e: # noqa: BLE001
print(f" ERROR {name:26s} {type(e).__name__}: {e}")
failed += 1
print(f"\n{len(tests)-failed}/{len(tests)} passed")
sys.exit(1 if failed else 0)