74/74 params, max abs diff 4.6e-4 on values around -163 (~3e-6 relative), 60ms for a 16^3 latent -> 64^3 occupancy grid. Fully dense, so no sparse ops involved - but MLX's Conv3d is channels-LAST where torch is NCDHW, so tensors are carried channels-last throughout and transposed only at the boundaries. The converter already emits [O,kz,ky,kx,I] to match. pixel_shuffle_3d had to be rewritten for that layout: the (H,s)(W,s)(D,s) interleave order is what matters and getting it wrong scrambles the grid while preserving its shape. This completes the structure stage end to end: image -> ss_flow -> latent -> ss_dec. |
||
|---|---|---|
| pixal3d_mlx | ||
| tests | ||
| .gitignore | ||
| CLAUDE.md | ||
| README.md | ||
pixal3d_mrp_mlx
MLX port of Pixal3D (TencentARC + Tsinghua, SIGGRAPH 2026, MIT) — pixel-aligned single-image 3D generation — for Apple Silicon.
Built on trellis_sparse_mlx, the shared TRELLIS-lineage
sparse core. Pixal3D and LATO.2 inherit the same sparse module from TRELLIS.2, so the
expensive part — submanifold sparse convolution, which has no Metal implementation — is
already done and tested there.
Why Pixal3D
It back-projects pixel features directly into 3D rather than injecting them through attention, so silhouettes stay exact to the source image. Different failure mode from TRELLIS/Hunyuan, and complementary to them.
Upstream needs ~24 GB VRAM, which is a wall on consumer Nvidia and a non-issue on a 128 GB+ Ultra.
Scope, measured against the shared core
| Need | Status |
|---|---|
SparseConv3d — every call is (c, out, 3), i.e. stride=1/padding=None → SubMConv3d |
✅ in shared core |
attn_mode='full' (the only mode used) |
✅ in shared core |
SparseLinear, norms, activations, ResBlock, transformer blocks |
✅ in shared core |
SparseDownsample(2) |
✅ in shared core |
VarLenTensor / SparseTensor split + get/register_spatial_cache |
✅ added to shared core |
dense nn.Conv3d(.., 2, stride=2) in sparse_structure_vae |
✅ maps to mlx.nn.Conv3d |
SparseUpsample(2) |
❌ to do — cache-paired inverse of a downsample |
SparseSpatial2Channel(2) |
❌ to do — sparse pixel-shuffle, spatial→channel |
Both of those are now done in the shared core.
Correction to the earlier scope
"the gap is two ops" was accurate about modules/sparse/ — the sparse primitives. It
undercounted the model blocks, which the configs revealed:
| Still needed | Where |
|---|---|
| RoPE positional embedding | all 4 flow models (pe_mode: "rope") — but rope_phases ships as a stored tensor, so phases are precomputed, not derived |
qk_rms_norm on q and k |
all 4 flow models |
AdaLN modulation (share_mod: true) |
all 4 flow models |
image_attn_mode: "proj" conditioning |
all 4 flow models |
SparseConvNeXtBlock3d |
shape_dec, tex_dec |
SparseResBlockC2S3d (channel↔spatial) |
shape_dec, tex_dec — uses SparseSpatial2Channel |
Offsetting that, a genuine simplification the tensors revealed: the four flow models
(~20 GB, the bulk of the download) contain ZERO 5-D tensors. ss_flow and the three
slat_flow DiTs are pure transformers — they never touch sparse convolution, so they
need none of the sparse core, just DiT blocks.
Model surface
pixal3d/models/
sparse_structure_vae.py dense Conv3d — voxel structure
sparse_structure_flow.py structure flow (SS)
structured_latent_flow.py SLAT flow
sc_vaes/sparse_unet_vae.py the only file using sparse conv
Weights: 24.04 GB across 19 files (1.3B DiTs at 512/1024 + shape/tex decoders).
Status
- Scoped against the shared core
- Weights downloaded (24 GB) — each ships a sibling
.jsonwith the exact config, so unlike LATO.2 there is no architecture to infer upsample(masked) +downsample(mode=)landed in the shared core- Weight converter — decoders remap KRSC→
[K³,in,out], dtype preserved, remap verified a pure permutation. Flow models pass through untouched. - DiT blocks: RoPE (stored complex phases), qk_rms_norm, AdaLN modulation, proj conditioning
- All four flow models verified against upstream at correlation 1.00000000
—
ss_flow(max diff 1.2e-5) andslat_flow(9.3e-6), 700/700 params each, on the real 1.3B checkpoints SparseConvNeXtBlock3d,SparseResBlockC2S3d,spatial2channel/channel2spatial— all in the shared core, round-trip and selective-growth testedss_decverified at correlation 1.00000000 (74/74), completing the whole structure stage: image -> ss_flow -> latent -> ss_dec -> 64^3 occupancy gridshape_dec/tex_decmodel graphs (the two sparse decoders)- End-to-end pipeline wiring
Numerical verification
The flow models contain no sparse convolution, which means upstream runs on CPU torch
here — swap flash-attn for F.scaled_dot_product_attention and it loads the real
checkpoint and runs. So unlike the sparse path (where spconv is uninstallable and the
oracle had to be hand-written), these are diffed against upstream directly:
MLX : mean +0.18490 std 0.86841
UPSTREAM : mean +0.18490 std 0.86841
max abs diff 1.216e-05 correlation 1.00000000
Getting there required finding three bugs that weight-key matching could not catch — the loader reported a perfect 700/700 with 0 missing and 0 unmapped through all of them:
- A parameterless final
LayerNormbetween the last block andout_layer. It has no weights, so it leaves no trace in the checkpoint. Without it the output was ~200x too large (std 187 vs 0.87). rope_phasesis complex64, built withtorch.polar. The rotation is a complex multiply and cos/sin are the phase's real/imaginary parts — takingcos()of a complex phase is meaningless.- qk RMS norm is applied BEFORE RoPE, not after. They do not commute. Reversed, a single block still correlated 0.9998 with upstream; over 30 blocks that compounds to 0.84. This one is invisible without an oracle.