lato.2_mrp_mlx/CLAUDE.md
John 97dcdfb54a MLX sparse core: SubMConv3d + SparseTensor + weight converter
The blocker for LATO.2 on Apple Silicon is one op, not the whole setup.sh --all
CUDA stack. Measured: 5 of 7 checkpoints are fully dense, and every SparseConv3d
in the model is constructed stride=1/padding=None, which upstream dispatches to
spconv's SubMConv3d. No strided or inverse sparse conv is ever instantiated.

- SubMConv3d in pure MLX via a sorted-key indice map (27 lookups/voxel vectorised,
  cached per coordinate set the way spconv uses indice_key)
- SparseTensor container + subdivide upsampling
- Converter handles the 5-D layout collision: spconv KRSC [O,kz,ky,kx,I] vs torch
  Conv3d [O,I,kz,ky,kx]. Rank alone is ambiguous; misreading it silently mangles
  the voxel encoder.
- 7/7 tests pass vs an independent naive reference, max err 3e-7. spconv has no
  Metal build so there is no upstream oracle; the reference shares no indexing code.

Kernel orientation (feats[c+d] vs feats[c-d]) remains unverified and is silent
when wrong; --flip-kernel builds the mirror for an end-to-end A/B.
2026-08-02 10:04:24 +10:00

27 lines
1.5 KiB
Markdown

# lato.2_mrp_mlx — working notes
MLX port of LATO.2 (factorised mesh gen: vertex flow -> connectivity flow) for Apple Silicon.
## Key facts established by measurement (don't re-derive)
- Upstream is CUDA-locked via `modules/sparse/__init__.py`: BACKEND accepts only
spconv/torchsparse, ATTN only xformers/flash_attn. No SDPA fallback exists upstream.
- The ONLY CUDA-locked op that matters is submanifold 3x3x3 conv. Every SparseConv3d in
the model uses default `stride=1, padding=None` -> SubMConv3d. No strided/inverse conv.
- Of 7 checkpoints, only vvae.pt has sparse conv kernels (18). vflow.pt is sparse-typed
but conv-free. voxel_encoder.pt is dense nn.Conv3d (29). Rest are fully dense.
- spconv weight layout is KRSC `[out, kz, ky, kx, in]` (verified against vvae.pt).
torch nn.Conv3d is `[out, in, kz, ky, kx]`. Both are 5-D -> `classify_5d()` disambiguates;
do not treat rank-5 as automatically sparse.
- Upsampling is `SparseSubdivide` (coords*2 + unit cube, feats replicated), not inverse conv.
## Unresolved
- Kernel orientation (correlation vs convolution) is unverified and numerically silent.
Settle it via V-VAE reconstruction quality, not by argument. `--flip-kernel` builds the
mirrored weights.
## Conventions
- Project venv at `.venv` (mlx + numpy + torch + trimesh). torch is CONVERSION-ONLY.
- `upstream/LATO.2` is vendored read-only reference; never edit it.
- Tests must compare against an independent implementation, not a second copy of the same
indexing logic. spconv is unavailable here so there is no upstream numerical oracle.