pixal3d_mrp_mlx/README.md
John 6ac586b947 SparseStructureDecoder verified at correlation 1.00000000
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.
2026-08-02 13:27:24 +10:00

111 lines
5.3 KiB
Markdown

# pixal3d_mrp_mlx
MLX port of [Pixal3D](https://github.com/TencentARC/Pixal3D) (TencentARC + Tsinghua,
SIGGRAPH 2026, MIT) — pixel-aligned single-image 3D generation — for Apple Silicon.
Built on **[`trellis_sparse_mlx`](../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
- [x] Scoped against the shared core
- [x] Weights downloaded (24 GB) — each ships a sibling `.json` with the exact config,
so unlike LATO.2 there is no architecture to infer
- [x] `upsample` (masked) + `downsample(mode=)` landed in the shared core
- [x] **Weight converter** — decoders remap KRSC→`[K³,in,out]`, dtype preserved,
remap verified a pure permutation. Flow models pass through untouched.
- [x] DiT blocks: RoPE (stored complex phases), qk_rms_norm, AdaLN modulation, proj conditioning
- [x] **All four flow models verified against upstream at correlation 1.00000000**
`ss_flow` (max diff 1.2e-5) and `slat_flow` (9.3e-6), 700/700 params each,
on the real 1.3B checkpoints
- [x] `SparseConvNeXtBlock3d`, `SparseResBlockC2S3d`, `spatial2channel`/`channel2spatial`
— all in the shared core, round-trip and selective-growth tested
- [x] **`ss_dec` verified at correlation 1.00000000** (74/74), completing the whole
structure stage: image -> ss_flow -> latent -> ss_dec -> 64^3 occupancy grid
- [ ] `shape_dec` / `tex_dec` model 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:
1. **A parameterless final `LayerNorm`** between the last block and `out_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).
2. **`rope_phases` is complex64**, built with `torch.polar`. The rotation is a complex
multiply and cos/sin are the phase's real/imaginary parts — taking `cos()` of a
complex phase is meaningless.
3. **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.