pixal3d_mrp_mlx/README.md
John 18179bc479 Sparse VAE decoders: shape_dec and tex_dec
All seven Pixal3D models now load and run. shape_dec 292/292, tex_dec 284/284, both
with zero missing/unmapped/mismatched keys - the 8-param difference between them is
exactly the four to_subdiv layers, since tex_dec has pred_subdiv=False.

Behaviour is right: 12 latent voxels grow SELECTIVELY through four stages
(12 -> 91 -> 193 -> 358 -> 1150) rather than x8 each time, which would have reached
49,152. Scale lands at 1/16. shape_dec emits the hardcoded 7 channels and its vertex
head produces offsets inside the [-0.5, 1.5] band its sigmoid+voxel_margin allows.
tex_dec, guided by shape_dec's masks, reproduces exactly the same voxel count.

VERIFICATION CAVEAT, recorded in the README: these two are the only models using sparse
conv, so upstream cannot run here and there is no numerical oracle. Unlike the five
models verified at correlation 1.0, these are checked structurally and behaviourally
only. Weaker evidence, and labelled as such rather than presented alongside the
verified results.
2026-08-02 13:31:53 +10:00

129 lines
6.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
- [x] `shape_dec` / `tex_dec` — both load complete (292/292, 284/284) and run.
**Behaviourally** checked only; see the verification note below
- [ ] End-to-end pipeline wiring (image encoder -> flows -> decoders -> mesh export)
## Model status
| model | params | verification |
|---|---|---|
| `ss_flow` | 700/700 | **corr 1.00000000** vs upstream |
| `slat_flow` x3 | 700/700 | **corr 1.00000000** vs upstream |
| `ss_dec` | 74/74 | **corr 1.00000000** vs upstream |
| `shape_dec` | 292/292 | behavioural only — no oracle possible |
| `tex_dec` | 284/284 | behavioural only — no oracle possible |
The split is not arbitrary: the first five contain no sparse convolution, so upstream
runs on CPU torch and can be diffed directly. The two decoders do use it, spconv has no
Metal build, and so there is nothing to diff against. Their blocks are individually
tested, and the assembled graphs are checked for complete weight mapping, selective
growth, correct scale and a vertex head inside its valid band — but that is weaker
evidence than a correlation and should be read that way.
## 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.