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

1.5 KiB

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.