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. |
||
|---|---|---|
| bench | ||
| tests | ||
| trellis_sparse_mlx | ||
| trellis_sparse_mlx.egg-info | ||
| .gitignore | ||
| pyproject.toml | ||
| README.md | ||
trellis_sparse_mlx
The TRELLIS-lineage sparse module, in MLX, so the models built on it run on Apple Silicon.
Why this is a package and not vendored code
Microsoft's TRELLIS.2 sparse stack has been inherited near-verbatim by a growing family
of 3D generation models — LATO.2 and TencentARC's
Pixal3D among them. Every one of them hard-requires
spconv or torchsparse:
BACKEND = 'spconv' # accepts only ['spconv', 'torchsparse'] — both CUDA-only
ATTN = 'flash_attn' # accepts only ['xformers', 'flash_attn'] — both CUDA-only
Neither sparse backend has a Metal build, and there is no SDPA fallback. That single dependency is what keeps the entire lineage off Apple Silicon.
The blocker turns out to be one operation. Every SparseConv3d in these models is
constructed stride=1, padding=None, which spconv dispatches to SubMConv3d. Nothing
instantiates strided or inverse sparse conv. Implement submanifold convolution and the
remainder is ordinary linear / norm / attention work.
Since each new model in this family needs the same core, it lives here once — tested — rather than being copy-pasted per port.
What's in it
SubMConv3d |
submanifold 3×3×3 (and 1×1×1) sparse conv — the actual blocker |
SparseTensor |
coords [N,4] + feats [N,C], batch-contiguous, with an indice-map cache |
subdivide / downsample |
×2 coord expansion; ×2 max-pool (upstream says "average" in its docstring but passes reduce="amax") |
SparseLinear, LayerNorm32, SparseGroupNorm32, SparseSiLU/GELU |
|
SparseResBlock, SparseFeedForwardNet |
|
SparseMultiHeadAttention (self + cross), SparseTransformerBlock, SparseTransformerCrossBlock |
attn_mode="full", which is all these models instantiate |
Performance
SubMConv3d, 128³ grid / 128 channels / ~10% occupancy (209,715 voxels), warm indice cache:
| box | chip | GPU cores | conv_ms |
|---|---|---|---|
| m3ultra | M3 Ultra | 80 | 13.8 |
| m2max | M2 Max | 38 | 27.3 |
| m1ultra | M1 Ultra | 64 | 39.7 |
| m1max | M1 Max | 32 | 41.4 |
| m4pro | M4 Pro | ~20 | 55.1 |
The naive implementation — a loop accumulating gather(i) @ W[i] over the 27 offsets —
was 5.9× slower, and the giveaway was that throughput anti-correlated with GPU core
count: the 80-core Ultra lost to a 38-core M2 Max, because 2·K³ = 54 tiny dispatches per
layer never fill the machine and UltraFusion punishes small dispatches hardest.
Two fixes: fuse the taps into a single [N, K³·Cin] × [K³·Cin, Cout] matmul, and cache the
prepared device gather index (not just the raw indice map — rebuilding the sentinel
substitution and re-uploading cost ~26 ms against ~13 ms of real GPU work). Only after
both does core count predict performance, which is the sign it is finally compute-bound.
Correctness
spconv cannot be installed on Apple Silicon — that is the entire reason this exists — so
there is no upstream oracle for the conv. tests/test_sparse.py instead checks the
vectorised implementation against a deliberately naive one written from the definition
(dict lookup, per-voxel loop) sharing no indexing code, so an off-by-one cannot hide in
both. tests/test_ops.py compares every other layer against torch, which is
available. 13/13 pass.
One assumption is unverified: whether spconv gathers feats[c+d] (cross-correlation,
the deep-learning convention and what this implements) or feats[c-d]. A flipped kernel is
numerically silent — latent statistics were tried as a discriminator and do not separate
them. It needs an end-to-end reconstruction to settle.
Licence
MIT. Derived from the TRELLIS/TRELLIS.2 sparse module (Copyright Microsoft Corporation and VAST-AI-Research contributors, MIT).