commit 70cd436eb9ea33971d5e84f2da51fd164585c5b9 Author: John Date: Sun Aug 2 10:31:07 2026 +1000 Shared TRELLIS-lineage sparse core in MLX Extracted from lato.2_mrp_mlx. The same sparse module underlies LATO.2, Pixal3D and the rest of the TRELLIS.2 family, and all of them are blocked on Apple Silicon by the same single op - submanifold conv - so it belongs in one tested package rather than vendored per port. 13/13 tests: 7 for the conv against a hand-written reference (spconv is uninstallable here so there is no upstream oracle), 6 for the remaining layers against torch. SubMConv3d runs 13.8ms at 128^3/128ch on m3ultra, 5.9x faster than the obvious per-offset loop. diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..0ce2e53 --- /dev/null +++ b/.gitignore @@ -0,0 +1,3 @@ +.venv +__pycache__/ +*.pyc diff --git a/README.md b/README.md new file mode 100644 index 0000000..55522bb --- /dev/null +++ b/README.md @@ -0,0 +1,78 @@ +# trellis_sparse_mlx + +The TRELLIS-lineage sparse module, in MLX, so the models built on it run on Apple Silicon. + +## Why this is a package and not vendored code + +Microsoft's TRELLIS.2 sparse stack has been inherited near-verbatim by a growing family +of 3D generation models — [LATO.2](https://github.com/LoHhhha/LATO.2) and TencentARC's +[Pixal3D](https://github.com/TencentARC/Pixal3D) among them. Every one of them hard-requires +`spconv` or `torchsparse`: + +```python +BACKEND = 'spconv' # accepts only ['spconv', 'torchsparse'] — both CUDA-only +ATTN = 'flash_attn' # accepts only ['xformers', 'flash_attn'] — both CUDA-only +``` + +Neither sparse backend has a Metal build, and there is no SDPA fallback. That single +dependency is what keeps the entire lineage off Apple Silicon. + +The blocker turns out to be **one operation**. Every `SparseConv3d` in these models is +constructed `stride=1, padding=None`, which spconv dispatches to `SubMConv3d`. Nothing +instantiates strided or inverse sparse conv. Implement submanifold convolution and the +remainder is ordinary linear / norm / attention work. + +Since each new model in this family needs the same core, it lives here once — tested — +rather than being copy-pasted per port. + +## What's in it + +| | | +|---|---| +| `SubMConv3d` | submanifold 3×3×3 (and 1×1×1) sparse conv — the actual blocker | +| `SparseTensor` | coords `[N,4]` + feats `[N,C]`, batch-contiguous, with an indice-map cache | +| `subdivide` / `downsample` | ×2 coord expansion; ×2 **max**-pool (upstream says "average" in its docstring but passes `reduce="amax"`) | +| `SparseLinear`, `LayerNorm32`, `SparseGroupNorm32`, `SparseSiLU/GELU` | | +| `SparseResBlock`, `SparseFeedForwardNet` | | +| `SparseMultiHeadAttention` (self + cross), `SparseTransformerBlock`, `SparseTransformerCrossBlock` | `attn_mode="full"`, which is all these models instantiate | + +## Performance + +`SubMConv3d`, 128³ grid / 128 channels / ~10% occupancy (209,715 voxels), warm indice cache: + +| box | chip | GPU cores | conv_ms | +|---|---|---|---| +| m3ultra | M3 Ultra | 80 | **13.8** | +| m2max | M2 Max | 38 | 27.3 | +| m1ultra | M1 Ultra | 64 | 39.7 | +| m1max | M1 Max | 32 | 41.4 | +| m4pro | M4 Pro | ~20 | 55.1 | + +The naive implementation — a loop accumulating `gather(i) @ W[i]` over the 27 offsets — +was **5.9× slower**, and the giveaway was that throughput *anti-correlated* with GPU core +count: the 80-core Ultra lost to a 38-core M2 Max, because 2·K³ = 54 tiny dispatches per +layer never fill the machine and UltraFusion punishes small dispatches hardest. + +Two fixes: fuse the taps into a single `[N, K³·Cin] × [K³·Cin, Cout]` matmul, and cache the +*prepared device* gather index (not just the raw indice map — rebuilding the sentinel +substitution and re-uploading cost ~26 ms against ~13 ms of real GPU work). Only after +both does core count predict performance, which is the sign it is finally compute-bound. + +## Correctness + +`spconv` cannot be installed on Apple Silicon — that is the entire reason this exists — so +there is no upstream oracle for the conv. `tests/test_sparse.py` instead checks the +vectorised implementation against a deliberately naive one written from the definition +(dict lookup, per-voxel loop) sharing no indexing code, so an off-by-one cannot hide in +both. `tests/test_ops.py` compares every other layer against **torch**, which *is* +available. 13/13 pass. + +**One assumption is unverified**: whether spconv gathers `feats[c+d]` (cross-correlation, +the deep-learning convention and what this implements) or `feats[c-d]`. A flipped kernel is +numerically silent — latent statistics were tried as a discriminator and do not separate +them. It needs an end-to-end reconstruction to settle. + +## Licence + +MIT. Derived from the TRELLIS/TRELLIS.2 sparse module (Copyright Microsoft Corporation and +VAST-AI-Research contributors, MIT). diff --git a/bench/bench_subm.py b/bench/bench_subm.py new file mode 100644 index 0000000..f79fa05 --- /dev/null +++ b/bench/bench_subm.py @@ -0,0 +1,128 @@ +"""Benchmark SubMConv3d across the fleet. Needs only mlx + numpy. + +Submanifold conv is the hot op in the V-VAE, so its scaling across Apple Silicon +generations is what decides which box should run LATO.2 once the port lands. Two +distinct costs are timed separately because they scale differently: + + indice map — coordinate hashing + binary search, runs on CPU in numpy, cacheable + per coordinate set (upstream reuses it via `indice_key`) + conv — 27 gathers + 27 matmuls on the GPU, paid every layer + +Occupancy is set to ~10% of the grid, roughly what a voxelised surface gives you. +""" + +from __future__ import annotations + +import argparse +import json +import platform +import subprocess +import sys +import time +from pathlib import Path + +import mlx.core as mx +import numpy as np + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) +from trellis_sparse_mlx.conv import SubMConv3d, build_indice_map # noqa: E402 +from trellis_sparse_mlx.tensor import SparseTensor # noqa: E402 + + +def machine() -> dict: + def sysctl(k): + try: + return subprocess.run( + ["sysctl", "-n", k], capture_output=True, text=True, timeout=5 + ).stdout.strip() + except Exception: # noqa: BLE001 + return "?" + + mem = sysctl("hw.memsize") + return { + "host": platform.node().split(".")[0], + "chip": sysctl("machdep.cpu.brand_string"), + "ram_gb": round(int(mem) / 1e9) if mem.isdigit() else "?", + "mlx": mx.__version__, + } + + +def make_sparse(res: int, channels: int, occupancy: float, seed: int = 0): + rng = np.random.default_rng(seed) + n_total = res**3 + n = max(64, int(n_total * occupancy)) + flat = rng.choice(n_total, size=n, replace=False) + flat.sort() + z, rem = np.divmod(flat, res * res) + y, x = np.divmod(rem, res) + coords = np.stack([np.zeros_like(z), z, y, x], axis=1).astype(np.int32) + feats = rng.standard_normal((n, channels)).astype(np.float32) + return mx.array(coords), mx.array(feats) + + +def time_it(fn, warmup=1, iters=3): + for _ in range(warmup): + fn() + ts = [] + for _ in range(iters): + t0 = time.perf_counter() + fn() + ts.append(time.perf_counter() - t0) + return min(ts) + + +def main() -> None: + ap = argparse.ArgumentParser() + ap.add_argument("--out", default=None, help="write JSON result here") + ap.add_argument("--occupancy", type=float, default=0.10) + ap.add_argument("--iters", type=int, default=3) + a = ap.parse_args() + + info = machine() + print(f"== {info['host']} {info['chip']} {info['ram_gb']}GB mlx {info['mlx']}") + print(f"{'res':>5} {'chan':>5} {'voxels':>9} {'imap_ms':>9} {'conv_ms':>9} {'Mvox/s':>8}") + + rows = [] + for res, ch in [(32, 64), (64, 64), (64, 128), (128, 128)]: + try: + coords, feats = make_sparse(res, ch, a.occupancy) + n = coords.shape[0] + conv = SubMConv3d(ch, ch, 3, bias=True, indice_key=f"b{res}") + + t_imap = time_it(lambda: build_indice_map(coords, 3), warmup=0, iters=1) + + x = SparseTensor(feats, coords) + _ = conv(x) # populate the indice cache so we time conv alone + mx.eval(_.feats) + + def run(): + mx.eval(conv(x).feats) + + t_conv = time_it(run, iters=a.iters) + rows.append( + { + "res": res, + "channels": ch, + "voxels": n, + "imap_ms": t_imap * 1e3, + "conv_ms": t_conv * 1e3, + "mvox_s": n / t_conv / 1e6, + } + ) + print( + f"{res:>5} {ch:>5} {n:>9,} {t_imap*1e3:>9.1f} {t_conv*1e3:>9.1f} " + f"{n/t_conv/1e6:>8.2f}" + ) + except Exception as e: # noqa: BLE001 + print(f"{res:>5} {ch:>5} FAILED: {type(e).__name__}: {e}") + rows.append({"res": res, "channels": ch, "error": f"{type(e).__name__}: {e}"}) + + result = {"machine": info, "occupancy": a.occupancy, "rows": rows} + if a.out: + Path(a.out).parent.mkdir(parents=True, exist_ok=True) + Path(a.out).write_text(json.dumps(result, indent=2)) + print(f"\nwrote {a.out}") + + +if __name__ == "__main__": + main() diff --git a/bench/m1max.json b/bench/m1max.json new file mode 100644 index 0000000..72a5290 --- /dev/null +++ b/bench/m1max.json @@ -0,0 +1,43 @@ +{ + "machine": { + "host": "M1MAX", + "chip": "Apple M1 Max", + "ram_gb": 34, + "mlx": "0.32.0" + }, + "occupancy": 0.1, + "rows": [ + { + "res": 32, + "channels": 64, + "voxels": 3276, + "imap_ms": 4.740834003314376, + "conv_ms": 0.7942919619381428, + "mvox_s": 4.12442798993744 + }, + { + "res": 64, + "channels": 64, + "voxels": 26214, + "imap_ms": 36.672832909971476, + "conv_ms": 2.2925420198589563, + "mvox_s": 11.434468713298768 + }, + { + "res": 64, + "channels": 128, + "voxels": 26214, + "imap_ms": 35.68366612307727, + "conv_ms": 5.524750100448728, + "mvox_s": 4.744829996540633 + }, + { + "res": 128, + "channels": 128, + "voxels": 209715, + "imap_ms": 319.13441605865955, + "conv_ms": 41.41504201106727, + "mvox_s": 5.063739883300329 + } + ] +} \ No newline at end of file diff --git a/bench/m1ultra.json b/bench/m1ultra.json new file mode 100644 index 0000000..f2e6b9d --- /dev/null +++ b/bench/m1ultra.json @@ -0,0 +1,43 @@ +{ + "machine": { + "host": "ultra", + "chip": "Apple M1 Ultra", + "ram_gb": 137, + "mlx": "0.32.0" + }, + "occupancy": 0.1, + "rows": [ + { + "res": 32, + "channels": 64, + "voxels": 3276, + "imap_ms": 5.284500017296523, + "conv_ms": 2.137709001544863, + "mvox_s": 1.532481735181227 + }, + { + "res": 64, + "channels": 64, + "voxels": 26214, + "imap_ms": 39.58837501704693, + "conv_ms": 3.456457983702421, + "mvox_s": 7.584064416116698 + }, + { + "res": 64, + "channels": 128, + "voxels": 26214, + "imap_ms": 38.462458993308246, + "conv_ms": 4.72516700392589, + "mvox_s": 5.547740424459111 + }, + { + "res": 128, + "channels": 128, + "voxels": 209715, + "imap_ms": 349.49354198761284, + "conv_ms": 39.74312503123656, + "mvox_s": 5.276761699920983 + } + ] +} \ No newline at end of file diff --git a/bench/m2max.json b/bench/m2max.json new file mode 100644 index 0000000..f4a15f7 --- /dev/null +++ b/bench/m2max.json @@ -0,0 +1,43 @@ +{ + "machine": { + "host": "m2max", + "chip": "Apple M2 Max", + "ram_gb": 103, + "mlx": "0.32.0" + }, + "occupancy": 0.1, + "rows": [ + { + "res": 32, + "channels": 64, + "voxels": 3276, + "imap_ms": 5.0670419877860695, + "conv_ms": 0.9743330010678619, + "mvox_s": 3.362300154474423 + }, + { + "res": 64, + "channels": 64, + "voxels": 26214, + "imap_ms": 42.181833006907254, + "conv_ms": 2.3582500289194286, + "mvox_s": 11.115869682406615 + }, + { + "res": 64, + "channels": 128, + "voxels": 26214, + "imap_ms": 41.914000001270324, + "conv_ms": 3.8609580078627914, + "mvox_s": 6.789506631933196 + }, + { + "res": 128, + "channels": 128, + "voxels": 209715, + "imap_ms": 356.8419170042034, + "conv_ms": 27.282333001494408, + "mvox_s": 7.6868426167407575 + } + ] +} \ No newline at end of file diff --git a/bench/m3ultra.json b/bench/m3ultra.json new file mode 100644 index 0000000..2c6def4 --- /dev/null +++ b/bench/m3ultra.json @@ -0,0 +1,43 @@ +{ + "machine": { + "host": "m3ultra", + "chip": "Apple M3 Ultra", + "ram_gb": 275, + "mlx": "0.32.0" + }, + "occupancy": 0.1, + "rows": [ + { + "res": 32, + "channels": 64, + "voxels": 3276, + "imap_ms": 3.659042005892843, + "conv_ms": 0.4588330048136413, + "mvox_s": 7.1398525512143 + }, + { + "res": 64, + "channels": 64, + "voxels": 26214, + "imap_ms": 28.264666034374386, + "conv_ms": 1.1872079921886325, + "mvox_s": 22.080376962148115 + }, + { + "res": 64, + "channels": 128, + "voxels": 26214, + "imap_ms": 28.218958992511034, + "conv_ms": 3.0656250310130417, + "mvox_s": 8.550947925727739 + }, + { + "res": 128, + "channels": 128, + "voxels": 209715, + "imap_ms": 232.5495420373045, + "conv_ms": 13.753084000200033, + "mvox_s": 15.2485798819341 + } + ] +} \ No newline at end of file diff --git a/bench/m3ultra_fused.json b/bench/m3ultra_fused.json new file mode 100644 index 0000000..b9a927a --- /dev/null +++ b/bench/m3ultra_fused.json @@ -0,0 +1,43 @@ +{ + "machine": { + "host": "m3ultra", + "chip": "Apple M3 Ultra", + "ram_gb": 275, + "mlx": "0.32.0" + }, + "occupancy": 0.1, + "rows": [ + { + "res": 32, + "channels": 64, + "voxels": 3276, + "imap_ms": 3.81183298304677, + "conv_ms": 1.5459590358659625, + "mvox_s": 2.1190729663577157 + }, + { + "res": 64, + "channels": 64, + "voxels": 26214, + "imap_ms": 28.528874972835183, + "conv_ms": 2.9240419971756637, + "mvox_s": 8.964987515678686 + }, + { + "res": 64, + "channels": 128, + "voxels": 26214, + "imap_ms": 28.55800004908815, + "conv_ms": 6.438167009036988, + "mvox_s": 4.071655793210163 + }, + { + "res": 128, + "channels": 128, + "voxels": 209715, + "imap_ms": 237.24404198583215, + "conv_ms": 39.25516700837761, + "mvox_s": 5.342354038520428 + } + ] +} \ No newline at end of file diff --git a/bench/m4pro.json b/bench/m4pro.json new file mode 100644 index 0000000..e1eac8e --- /dev/null +++ b/bench/m4pro.json @@ -0,0 +1,43 @@ +{ + "machine": { + "host": "m4pro", + "chip": "Apple M4 Pro", + "ram_gb": 26, + "mlx": "0.32.0" + }, + "occupancy": 0.1, + "rows": [ + { + "res": 32, + "channels": 64, + "voxels": 3276, + "imap_ms": 13.317167002242059, + "conv_ms": 1.4369579730555415, + "mvox_s": 2.2798161542845454 + }, + { + "res": 64, + "channels": 64, + "voxels": 26214, + "imap_ms": 25.363416003528982, + "conv_ms": 3.218833007849753, + "mvox_s": 8.143945316849939 + }, + { + "res": 64, + "channels": 128, + "voxels": 26214, + "imap_ms": 25.877792038954794, + "conv_ms": 8.792624983470887, + "mvox_s": 2.9813622267842965 + }, + { + "res": 128, + "channels": 128, + "voxels": 209715, + "imap_ms": 204.2026249691844, + "conv_ms": 55.10258302092552, + "mvox_s": 3.8059014387829975 + } + ] +} \ No newline at end of file diff --git a/bench/run_fleet.sh b/bench/run_fleet.sh new file mode 100755 index 0000000..2302c15 --- /dev/null +++ b/bench/run_fleet.sh @@ -0,0 +1,58 @@ +#!/bin/zsh +# Run the SubMConv3d benchmark on every Apple Silicon box in the fleet. +# Ships code only (no weights); each box gets a small mlx+numpy venv under ~/lato-bench. +set -u +SRC=/Users/m3ultra/Documents/lato.2_mrp_mlx +OUT=$SRC/bench +SSH_OPTS=(-o ConnectTimeout=10 -o BatchMode=yes -o StrictHostKeyChecking=accept-new) +SSH_STR="ssh -o ConnectTimeout=10 -o BatchMode=yes -o StrictHostKeyChecking=accept-new" + +# host:label (m3ultra runs locally, already done) +HOSTS=( + "m1max@100.92.78.24:m1max" + "m2max@100.120.83.110:m2max" + "m4pro@100.69.21.128:m4pro" + "johnking@100.91.239.7:m1ultra" +) + +hb(){ mkdir -p ~/.jobs; echo "$(date '+%F %T') | $1 | $2" > ~/.jobs/lato-fleet-bench.status; } +hb "starting" "${#HOSTS[@]} hosts" + +for entry in $HOSTS; do + host=${entry%%:*}; label=${entry##*:} + echo "===== $label ($host)" + hb "$label" "connecting" + + if ! ssh "${SSH_OPTS[@]}" "$host" 'true' 2>/dev/null; then + echo " SKIP unreachable"; continue + fi + + # m1max is disk-critical (~20GB free per fleet notes) - refuse rather than fill it + free_mb=$(ssh "${SSH_OPTS[@]}" "$host" "df -m / | tail -1 | awk '{print \$4}'" 2>/dev/null) + if [ -n "$free_mb" ] && [ "$free_mb" -lt 3000 ]; then + echo " SKIP only ${free_mb}MB free - not installing a venv here"; continue + fi + + ssh "${SSH_OPTS[@]}" "$host" 'mkdir -p ~/lato-bench' 2>/dev/null + rsync -a -e "$SSH_STR" --delete \ + "$SRC/lato_mlx" "$SRC/bench" "$host:~/lato-bench/" 2>&1 | tail -1 + + hb "$label" "installing venv" + ssh "${SSH_OPTS[@]}" "$host" ' + cd ~/lato-bench || exit 1 + UV=$(command -v uv || echo /opt/homebrew/bin/uv) + if [ ! -x "$UV" ]; then echo " no uv on this host"; exit 3; fi + [ -d .venv ] || "$UV" venv --python 3.12 .venv >/dev/null 2>&1 + VIRTUAL_ENV=.venv "$UV" pip install -q mlx numpy >/dev/null 2>&1 + .venv/bin/python -c "import mlx.core" 2>/dev/null || { echo " mlx install failed"; exit 4; } + ' || { echo " SKIP setup failed (rc=$?)"; continue; } + + hb "$label" "benchmarking" + ssh "${SSH_OPTS[@]}" "$host" \ + 'cd ~/lato-bench && .venv/bin/python bench/bench_subm.py --out bench/result.json' 2>&1 | tail -8 + rsync -a -e "$SSH_STR" "$host:~/lato-bench/bench/result.json" "$OUT/$label.json" 2>/dev/null \ + && echo " -> $OUT/$label.json" || echo " (no result pulled)" +done + +hb "DONE" "results in $OUT" +echo "===== done" diff --git a/pyproject.toml b/pyproject.toml new file mode 100644 index 0000000..b2e59f7 --- /dev/null +++ b/pyproject.toml @@ -0,0 +1,14 @@ +[project] +name = "trellis-sparse-mlx" +version = "0.1.0" +description = "MLX implementation of the TRELLIS-lineage sparse module (SubMConv3d et al) for Apple Silicon" +requires-python = ">=3.10" +dependencies = ["mlx>=0.20", "numpy>=1.24"] +license = { text = "MIT" } + +[build-system] +requires = ["setuptools>=68"] +build-backend = "setuptools.build_meta" + +[tool.setuptools] +packages = ["trellis_sparse_mlx"] diff --git a/tests/test_ops.py b/tests/test_ops.py new file mode 100644 index 0000000..51b1aae --- /dev/null +++ b/tests/test_ops.py @@ -0,0 +1,171 @@ +"""Sparse layer tests against torch. + +Unlike the submanifold conv — where spconv is uninstallable and the oracle had to be +hand-written — every layer here has a real torch counterpart, so these compare against +upstream's actual semantics rather than a paraphrase of them. The upstream forward +bodies are reproduced verbatim (see modules/sparse/{norm,linear,nonlinearity}.py), +including the [N_b,C] -> [1,C,N_b] GroupNorm reshape, which is the one that would fail +silently if guessed. +""" + +import sys +from pathlib import Path + +import mlx.core as mx +import numpy as np +import torch +import torch.nn.functional as F + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from trellis_sparse_mlx.ops import ( # noqa: E402 + LayerNorm32, + SparseGroupNorm32, + SparseMultiHeadAttention, + SparseTransformerBlock, +) +from trellis_sparse_mlx.tensor import SparseTensor # noqa: E402 + + +def make_batched(n_per_batch, channels, seed=0): + """Batch-contiguous coords, as upstream requires.""" + rng = np.random.default_rng(seed) + coords, feats = [], [] + for b, nb in enumerate(n_per_batch): + for i in range(nb): + coords.append((b, i // 16, (i // 4) % 4, i % 4)) + feats.append(rng.standard_normal((nb, channels)).astype(np.float32)) + return np.array(coords, dtype=np.int32), np.concatenate(feats, 0) + + +def test_group_norm_matches_torch(groups=8, channels=32): + n_per_batch = [37, 51] + coords, feats = make_batched(n_per_batch, channels, seed=1) + rng = np.random.default_rng(2) + w = rng.standard_normal(channels).astype(np.float32) + b = rng.standard_normal(channels).astype(np.float32) + + gn = SparseGroupNorm32(groups, channels) + gn.weight, gn.bias = mx.array(w), mx.array(b) + got = np.asarray(gn(SparseTensor(mx.array(feats), mx.array(coords))).feats) + + # upstream: per batch item, [N_b,C] -> permute -> [1,C,N_b] -> nn.GroupNorm + tg = torch.nn.GroupNorm(groups, channels, eps=1e-5, affine=True) + tg.weight.data = torch.tensor(w) + tg.bias.data = torch.tensor(b) + want = np.zeros_like(feats) + off = 0 + for nb in n_per_batch: + bf = torch.tensor(feats[off : off + nb]) + bf = bf.permute(1, 0).reshape(1, channels, -1) + bf = tg(bf) + want[off : off + nb] = bf.reshape(channels, -1).permute(1, 0).detach().numpy() + off += nb + + err = np.abs(got - want).max() + assert err < 2e-4, f"group norm err {err:.3g}" + return err + + +def test_group_norm_is_not_per_voxel(groups=8, channels=32): + """Guard the easy-to-miss distinction: GroupNorm here is NOT a per-voxel norm.""" + coords, feats = make_batched([40], channels, seed=5) + gn = SparseGroupNorm32(groups, channels) + got = np.asarray(gn(SparseTensor(mx.array(feats), mx.array(coords))).feats) + per_voxel = torch.nn.functional.group_norm( + torch.tensor(feats).reshape(40, channels), groups + ).numpy() + assert np.abs(got - per_voxel).max() > 1e-3, "matched per-voxel norm — reshape lost" + return 0.0 + + +def test_layer_norm_matches_torch(channels=64): + rng = np.random.default_rng(3) + feats = rng.standard_normal((50, channels)).astype(np.float32) + ln = LayerNorm32(channels, affine=False, eps=1e-6) + got = np.asarray(ln(mx.array(feats))) + want = F.layer_norm(torch.tensor(feats), (channels,), eps=1e-6).numpy() + err = np.abs(got - want).max() + assert err < 1e-5, f"layer norm err {err:.3g}" + return err + + +def test_self_attention_matches_torch(channels=64, heads=8): + n_per_batch = [23, 31] + coords, feats = make_batched(n_per_batch, channels, seed=4) + rng = np.random.default_rng(6) + wq = rng.standard_normal((channels * 3, channels)).astype(np.float32) * 0.05 + bq = rng.standard_normal((channels * 3,)).astype(np.float32) * 0.05 + wo = rng.standard_normal((channels, channels)).astype(np.float32) * 0.05 + bo = rng.standard_normal((channels,)).astype(np.float32) * 0.05 + + attn = SparseMultiHeadAttention(channels, heads) + attn.to_qkv.weight, attn.to_qkv.bias = mx.array(wq), mx.array(bq) + attn.to_out.weight, attn.to_out.bias = mx.array(wo), mx.array(bo) + got = np.asarray(attn(SparseTensor(mx.array(feats), mx.array(coords))).feats) + + # reference: attention strictly within each batch item + d = channels // heads + want = np.zeros_like(feats) + off = 0 + for nb in n_per_batch: + f = torch.tensor(feats[off : off + nb]) + qkv = F.linear(f, torch.tensor(wq), torch.tensor(bq)).reshape(nb, 3, heads, d) + q, k, v = qkv[:, 0], qkv[:, 1], qkv[:, 2] + o = F.scaled_dot_product_attention( + q.permute(1, 0, 2)[None], k.permute(1, 0, 2)[None], v.permute(1, 0, 2)[None] + ) + o = o[0].permute(1, 0, 2).reshape(nb, channels) + want[off : off + nb] = ( + F.linear(o, torch.tensor(wo), torch.tensor(bo)).detach().numpy() + ) + off += nb + + err = np.abs(got - want).max() + assert err < 2e-4, f"attention err {err:.3g}" + return err + + +def test_attention_does_not_cross_batches(channels=32, heads=4): + """Perturbing batch 1 must never change batch 0's output.""" + coords, feats = make_batched([12, 12], channels, seed=7) + attn = SparseMultiHeadAttention(channels, heads) + a = np.asarray(attn(SparseTensor(mx.array(feats), mx.array(coords))).feats) + f2 = feats.copy() + f2[12:] += 10.0 + b = np.asarray(attn(SparseTensor(mx.array(f2), mx.array(coords))).feats) + assert np.abs(a[:12] - b[:12]).max() < 1e-5, "batch 0 changed — attention leaked" + return 0.0 + + +def test_transformer_block_shape(channels=64, heads=8): + coords, feats = make_batched([20, 20], channels, seed=8) + blk = SparseTransformerBlock(channels, heads) + out = blk(SparseTensor(mx.array(feats), mx.array(coords))) + assert out.feats.shape == (40, channels) + assert np.isfinite(np.asarray(out.feats)).all(), "non-finite output" + return 0.0 + + +if __name__ == "__main__": + tests = [ + ("group norm vs torch", test_group_norm_matches_torch), + ("group norm != per-voxel", test_group_norm_is_not_per_voxel), + ("layer norm vs torch", test_layer_norm_matches_torch), + ("self-attn vs torch", test_self_attention_matches_torch), + ("attn batch isolation", test_attention_does_not_cross_batches), + ("transformer block", test_transformer_block_shape), + ] + failed = 0 + for name, fn in tests: + try: + err = fn() + print(f" PASS {name:26s} (max err {err:.2e})") + except AssertionError as e: + print(f" FAIL {name:26s} {e}") + failed += 1 + except Exception as e: # noqa: BLE001 + print(f" ERROR {name:26s} {type(e).__name__}: {e}") + failed += 1 + print(f"\n{len(tests)-failed}/{len(tests)} passed") + sys.exit(1 if failed else 0) diff --git a/tests/test_sparse.py b/tests/test_sparse.py new file mode 100644 index 0000000..0c069e6 --- /dev/null +++ b/tests/test_sparse.py @@ -0,0 +1,162 @@ +"""Correctness tests for the MLX sparse core. + +spconv cannot be installed on this machine — that is the entire reason this port exists — +so there is no way to diff against upstream numerically here. Instead the vectorised MLX +implementation is checked against a deliberately naive, obviously-correct reference +written straight from the definition of submanifold convolution (dict lookup, per-voxel +Python loop). The two share no indexing code, so an off-by-one in the fast path cannot +hide in both. +""" + +import sys +from pathlib import Path + +import mlx.core as mx +import numpy as np + +sys.path.insert(0, str(Path(__file__).resolve().parents[1])) + +from trellis_sparse_mlx.conv import SubMConv3d, build_indice_map, _kernel_offsets +from trellis_sparse_mlx.tensor import SparseTensor, subdivide + + +def reference_subm_conv(coords, feats, weight, bias, k): + """Definition of submanifold conv, written for obviousness, not speed.""" + occupied = {tuple(c): i for i, c in enumerate(coords.tolist())} + offsets = _kernel_offsets(k) + n, out_c = coords.shape[0], weight.shape[2] + out = np.zeros((n, out_c), dtype=np.float64) + for i, c in enumerate(coords.tolist()): + for oi, d in enumerate(offsets): + nb = (c[0], c[1] + d[0], c[2] + d[1], c[3] + d[2]) + j = occupied.get(nb) + if j is not None: + out[i] += feats[j].astype(np.float64) @ weight[oi].astype(np.float64) + if bias is not None: + out += bias.astype(np.float64) + return out + + +def random_sparse(n_vox, channels, batch=2, res=8, seed=0): + rng = np.random.default_rng(seed) + seen, coords = set(), [] + while len(coords) < n_vox: + b = int(rng.integers(0, batch)) + z, y, x = (int(v) for v in rng.integers(0, res, 3)) + if (b, z, y, x) in seen: + continue + seen.add((b, z, y, x)) + coords.append((b, z, y, x)) + coords.sort() # upstream requires batch-contiguous rows + coords = np.array(coords, dtype=np.int32) + feats = rng.standard_normal((n_vox, channels)).astype(np.float32) + return coords, feats + + +def test_subm_conv_matches_reference(k=3, n=180, cin=12, cout=7): + coords, feats = random_sparse(n, cin, seed=1) + rng = np.random.default_rng(2) + w = rng.standard_normal((k**3, cin, cout)).astype(np.float32) * 0.1 + b = rng.standard_normal((cout,)).astype(np.float32) + + conv = SubMConv3d(cin, cout, k, bias=True, indice_key="t") + conv.weight = mx.array(w) + conv.bias = mx.array(b) + + got = np.asarray(conv(SparseTensor(mx.array(feats), mx.array(coords))).feats) + want = reference_subm_conv(coords, feats, w, b, k) + err = np.abs(got - want).max() + assert err < 2e-4, f"k={k} max abs err {err:.3g}" + return err + + +def test_kernel1_is_pointwise(n=64, cin=8, cout=5): + coords, feats = random_sparse(n, cin, seed=3) + rng = np.random.default_rng(4) + w = rng.standard_normal((1, cin, cout)).astype(np.float32) + conv = SubMConv3d(cin, cout, 1, bias=False, indice_key="p") + conv.weight = mx.array(w) + got = np.asarray(conv(SparseTensor(mx.array(feats), mx.array(coords))).feats) + err = np.abs(got - feats @ w[0]).max() + assert err < 1e-4, f"k=1 err {err:.3g}" + return err + + +def test_isolated_voxel_sees_only_itself(): + """A voxel with no occupied neighbours must reduce to the centre tap alone.""" + coords = np.array([[0, 0, 0, 0], [0, 50, 50, 50]], dtype=np.int32) + feats = np.ones((2, 3), dtype=np.float32) + w = np.zeros((27, 3, 3), dtype=np.float32) + centre = 13 # index of (0,0,0) in centred C-order offsets + assert tuple(_kernel_offsets(3)[centre]) == (0, 0, 0) + w[centre] = np.eye(3) + conv = SubMConv3d(3, 3, 3, bias=False, indice_key="iso") + conv.weight = mx.array(w) + got = np.asarray(conv(SparseTensor(mx.array(feats), mx.array(coords))).feats) + assert np.abs(got - feats).max() < 1e-6, got + return 0.0 + + +def test_batches_do_not_leak(): + """Same spatial cell in two batch items must not become neighbours.""" + coords = np.array([[0, 1, 1, 1], [1, 1, 1, 1]], dtype=np.int32) + feats = np.array([[1.0], [100.0]], dtype=np.float32) + w = np.ones((27, 1, 1), dtype=np.float32) # sum every occupied neighbour + conv = SubMConv3d(1, 1, 3, bias=False, indice_key="b") + conv.weight = mx.array(w) + got = np.asarray(conv(SparseTensor(mx.array(feats), mx.array(coords))).feats) + assert np.allclose(got, [[1.0], [100.0]]), f"batch leak: {got}" + return 0.0 + + +def test_indice_map_hit_rate(): + """A fully dense block: interior voxels must find all 27 neighbours.""" + coords = np.array( + [[0, z, y, x] for z in range(4) for y in range(4) for x in range(4)], + dtype=np.int32, + ) + imap = build_indice_map(mx.array(coords), 3) + lin = {tuple(c): i for i, c in enumerate(coords.tolist())} + interior = [lin[(0, z, y, x)] for z in (1, 2) for y in (1, 2) for x in (1, 2)] + assert (imap[:, interior] != -1).all(), "interior voxel missing a neighbour" + corner = lin[(0, 0, 0, 0)] + assert (imap[:, corner] != -1).sum() == 8, "corner should see exactly 8 of 27" + return 0.0 + + +def test_subdivide(): + coords = np.array([[0, 1, 2, 3]], dtype=np.int32) + feats = np.array([[5.0, 6.0]], dtype=np.float32) + out = subdivide(SparseTensor(mx.array(feats), mx.array(coords))) + c = np.asarray(out.coords) + assert c.shape == (8, 4) and out.feats.shape == (8, 2) + assert set(map(tuple, c[:, 1:].tolist())) == { + (2 + a, 4 + b, 6 + d) for a in (0, 1) for b in (0, 1) for d in (0, 1) + } + assert np.abs(np.asarray(out.feats) - feats).max() < 1e-6 + return 0.0 + + +if __name__ == "__main__": + tests = [ + ("subm k=3 vs reference", lambda: test_subm_conv_matches_reference(3)), + ("subm k=5 vs reference", lambda: test_subm_conv_matches_reference(5, n=140)), + ("k=1 is pointwise", test_kernel1_is_pointwise), + ("isolated voxel", test_isolated_voxel_sees_only_itself), + ("batch isolation", test_batches_do_not_leak), + ("indice map hit rate", test_indice_map_hit_rate), + ("subdivide", test_subdivide), + ] + failed = 0 + for name, fn in tests: + try: + err = fn() + print(f" PASS {name:28s} (max err {err:.2e})") + except AssertionError as e: + print(f" FAIL {name:28s} {e}") + failed += 1 + except Exception as e: # noqa: BLE001 + print(f" ERROR {name:28s} {type(e).__name__}: {e}") + failed += 1 + print(f"\n{len(tests)-failed}/{len(tests)} passed") + sys.exit(1 if failed else 0) diff --git a/trellis_sparse_mlx/__init__.py b/trellis_sparse_mlx/__init__.py new file mode 100644 index 0000000..9abb36c --- /dev/null +++ b/trellis_sparse_mlx/__init__.py @@ -0,0 +1,38 @@ +"""MLX implementation of the TRELLIS-lineage sparse module. + +Microsoft's TRELLIS.2 sparse stack has been inherited, near-verbatim, by a growing +family of 3D generation models — LATO.2 and TencentARC's Pixal3D among them. All of +them hard-require spconv or torchsparse, neither of which has a Metal build, and that +single dependency is what keeps the whole lineage off Apple Silicon. + +The blocker is one operation: submanifold 3x3x3 convolution. Every SparseConv3d in +these models is constructed `stride=1, padding=None`, which spconv dispatches to +SubMConv3d. Implement that in MLX and the rest is ordinary linear/norm/attention work. + +Packaged separately from any one model so each port depends on a tested core rather +than vendoring its own copy. +""" + +from .conv import SubMConv3d, build_indice_map +from .tensor import SparseTensor, downsample, subdivide +from .ops import ( + LayerNorm32, + SparseFeedForwardNet, + SparseGELU, + SparseGroupNorm32, + SparseLinear, + SparseMultiHeadAttention, + SparseResBlock, + SparseSiLU, + SparseTransformerBlock, + SparseTransformerCrossBlock, +) + +__all__ = [ + "SparseTensor", "subdivide", "downsample", + "SubMConv3d", "build_indice_map", + "SparseLinear", "LayerNorm32", "SparseGroupNorm32", "SparseSiLU", "SparseGELU", + "SparseResBlock", "SparseFeedForwardNet", "SparseMultiHeadAttention", + "SparseTransformerBlock", "SparseTransformerCrossBlock", +] +__version__ = "0.1.0" diff --git a/trellis_sparse_mlx/conv.py b/trellis_sparse_mlx/conv.py new file mode 100644 index 0000000..224404a --- /dev/null +++ b/trellis_sparse_mlx/conv.py @@ -0,0 +1,181 @@ +"""Submanifold sparse 3D convolution in pure MLX. + +This is the ONLY genuinely CUDA-locked operation in LATO.2's inference path. Upstream +routes it to `spconv.SubMConv3d` (or torchsparse); neither has a Metal build, which is +what has kept LATO.2 — and TRELLIS before it — off Apple Silicon. + +Every SparseConv3d in the LATO.2 model code is constructed with the defaults +`stride=1, padding=None`, which upstream dispatches to SubMConv3d. So only the +submanifold case is needed, at kernel sizes 3 and 1. + +Submanifold semantics: the output occupies EXACTLY the input coordinates (no dilation of +the occupied set). For each output voxel c: + + out[c] = bias + sum over kernel offsets d of W[d] @ feats[c + d] (c+d occupied) + +Absent neighbours contribute nothing. Implemented by gathering into a feature matrix with +one appended zero row, so "missing" is index N and needs no masking in the hot loop. + +Building the indice map is the expensive part and depends only on the coordinate set, so +it is cached on the SparseTensor under `indice_key` — the same trick spconv uses, and the +reason upstream threads `indice_key=f"res_{resolution}"` through the ResBlocks. +""" + +from __future__ import annotations + +from typing import Optional + +import mlx.core as mx +import mlx.nn as nn +import numpy as np + +from .tensor import SparseTensor + +_MISSING = -1 + + +def _kernel_offsets(k: int) -> np.ndarray: + """Kernel offsets in C order over (dz, dy, dx), centred — matches spconv's ordering.""" + r = np.arange(k) - (k // 2) + return np.stack(np.meshgrid(r, r, r, indexing="ij"), axis=-1).reshape(-1, 3) + + +def build_indice_map(coords: mx.array, kernel_size: int) -> np.ndarray: + """[K^3, N] int32 — for each offset, the row of the neighbour, or -1 if unoccupied. + + Uses a sorted-key binary search rather than a Python dict: at k=3 this is 27 lookups + per voxel, and a per-voxel dict lookup would dominate runtime for any real mesh. + """ + c = np.asarray(coords, dtype=np.int64) + n = c.shape[0] + if n == 0: + return np.full((kernel_size**3, 0), _MISSING, dtype=np.int32) + + offsets = _kernel_offsets(kernel_size) + pad = kernel_size // 2 + + # Encode (batch,z,y,x) into one int64. Shift by `pad` so that neighbour coordinates + # of -1 stay non-negative and cannot alias onto a real cell at the opposite edge. + lo = c.min(axis=0) - pad + ext = (c.max(axis=0) + pad) - lo + 1 + strides = np.array( + [ext[1] * ext[2] * ext[3], ext[2] * ext[3], ext[3], 1], dtype=np.int64 + ) + + def encode(arr: np.ndarray) -> np.ndarray: + return ((arr - lo) * strides).sum(axis=1) + + keys = encode(c) + order = np.argsort(keys, kind="stable") + sorted_keys = keys[order] + + imap = np.empty((offsets.shape[0], n), dtype=np.int32) + for i, d in enumerate(offsets): + probe = c.copy() + probe[:, 1:] += d # batch index (column 0) never shifts + pk = encode(probe) + pos = np.searchsorted(sorted_keys, pk) + pos_clipped = np.clip(pos, 0, n - 1) + hit = sorted_keys[pos_clipped] == pk + imap[i] = np.where(hit, order[pos_clipped], _MISSING).astype(np.int32) + return imap + + +class SubMConv3d(nn.Module): + """Submanifold sparse conv. Weight layout [K^3, in_channels, out_channels].""" + + def __init__( + self, + in_channels: int, + out_channels: int, + kernel_size: int = 3, + bias: bool = True, + indice_key: Optional[str] = None, + ): + super().__init__() + if kernel_size % 2 != 1: + raise ValueError(f"kernel_size must be odd, got {kernel_size}") + self.in_channels = in_channels + self.out_channels = out_channels + self.kernel_size = kernel_size + self.indice_key = indice_key + + scale = (in_channels * kernel_size**3) ** -0.5 + self.weight = mx.random.uniform( + -scale, scale, (kernel_size**3, in_channels, out_channels) + ) + if bias: + self.bias = mx.zeros((out_channels,)) + + def _gather_index(self, x: SparseTensor, n: int) -> mx.array: + """Cached [N, K^3] gather index, already on-device. + + Caching only the raw indice map is not enough: rebuilding the "missing -> N" + substitution and re-uploading the index cost more than the convolution itself. + At 128^3/128ch that overhead was ~26ms against ~13ms of actual GPU work, i.e. + two thirds of the measured time was CPU-side bookkeeping repeated every layer. + The transposed, sentinel-substituted device array depends only on the + coordinate set, so it is cached whole. + """ + key = f"gidx_k{self.kernel_size}_{self.indice_key}" + cached = x.cache_get(key) + if cached is not None: + return cached + imap = build_indice_map(x.coords, self.kernel_size) + idx_t = mx.array(np.where(imap == _MISSING, n, imap).T) # [N, K^3] + x.cache_put(key, idx_t) + return idx_t + + def __call__(self, x: SparseTensor) -> SparseTensor: + n = x.feats.shape[0] + + # k=1 touches only the centre voxel, so it is exactly a per-voxel linear — + # skip the indice map entirely. + if self.kernel_size == 1: + out = x.feats @ self.weight[0] + if hasattr(self, "bias"): + out = out + self.bias + return x.replace(out) + + # One appended zero row: absent neighbours index it and contribute nothing, + # which avoids a per-offset boolean mask. + feats_pad = mx.concatenate( + [x.feats, mx.zeros((1, self.in_channels), dtype=x.feats.dtype)], axis=0 + ) + idx_t = self._gather_index(x, n) # [N, K^3], cached on device + k3 = self.kernel_size**3 + + # Fuse the K^3 taps into ONE gather + ONE matmul. + # + # The obvious implementation loops over the K^3 offsets accumulating + # `gather(i) @ W[i]`, but that issues 2*K^3 tiny GPU dispatches with Python + # between them, and each is far too small to fill the machine. Fleet + # benchmarking made this unmistakable: the 80-core M3 Ultra came in SLOWER + # than a 38-core M2 Max (81.6ms vs 58.3ms at 128^3/128ch), i.e. throughput was + # anti-correlated with core count — the signature of launch-latency binding + # rather than compute binding. + # + # Concatenating neighbours along the channel axis turns the whole thing into a + # single [N, K^3*Cin] x [K^3*Cin, Cout] matmul, which is one dispatch big + # enough to actually occupy the GPU. + # + # That buffer is N*K^3*Cin floats, so it is chunked over rows to keep peak + # memory bounded (~256MB/chunk) — at 128^3 and 128ch the unchunked form alone + # would be ~2.9GB, which is fine on a Studio and not fine on an 8GB mini. + w_flat = self.weight.reshape(k3 * self.in_channels, self.out_channels).astype( + x.feats.dtype + ) + bytes_per_row = k3 * self.in_channels * 4 + chunk = max(1, min(n, (256 << 20) // max(bytes_per_row, 1))) + + outs = [] + for start in range(0, n, chunk): + stop = min(start + chunk, n) + g = mx.take(feats_pad, idx_t[start:stop].reshape(-1), axis=0) + g = g.reshape(stop - start, k3 * self.in_channels) + outs.append(g @ w_flat) + out = outs[0] if len(outs) == 1 else mx.concatenate(outs, axis=0) + + if hasattr(self, "bias"): + out = out + self.bias.astype(x.feats.dtype) + return x.replace(out) diff --git a/trellis_sparse_mlx/ops.py b/trellis_sparse_mlx/ops.py new file mode 100644 index 0000000..44fc92f --- /dev/null +++ b/trellis_sparse_mlx/ops.py @@ -0,0 +1,256 @@ +"""The rest of the sparse layer set, in MLX. + +Nothing here is CUDA-locked upstream — these are ordinary linear/norm/attention layers +that merely take a SparseTensor instead of a dense one. They are reimplemented rather +than adapted because upstream's versions inherit from torch modules. + +Two normalisation shapes are easy to conflate, and upstream uses both: + + LayerNorm32 applied to `x.feats` directly -> per-voxel over channels. + SparseGroupNorm32 reshapes [N_b, C] -> [1, C, N_b] per batch item, so statistics + are over (channels-in-group x voxels) WITHIN one batch item. + Getting this wrong is silent: shapes match either way. + +Attention runs in `attn_mode="full"`, which upstream defines as full attention *within* +each batch item (never across). Batch rows are contiguous, so each item is one slice. +""" + +from __future__ import annotations + +from typing import Optional + +import mlx.core as mx +import mlx.nn as nn + +from .conv import SubMConv3d +from .tensor import SparseTensor + +# ---------------------------------------------------------------- primitives + + +class SparseLinear(nn.Module): + def __init__(self, in_features: int, out_features: int, bias: bool = True): + super().__init__() + self.linear = nn.Linear(in_features, out_features, bias=bias) + + def __call__(self, x: SparseTensor) -> SparseTensor: + return x.replace(self.linear(x.feats)) + + +class LayerNorm32(nn.Module): + """Per-voxel LayerNorm over channels; computed in fp32 as upstream does.""" + + def __init__(self, dim: int, affine: bool = False, eps: float = 1e-6): + super().__init__() + self.eps = eps + self.affine = affine + if affine: + self.weight = mx.ones((dim,)) + self.bias = mx.zeros((dim,)) + + def __call__(self, feats: mx.array) -> mx.array: + dt = feats.dtype + f = feats.astype(mx.float32) + mu = mx.mean(f, axis=-1, keepdims=True) + var = mx.var(f, axis=-1, keepdims=True) + f = (f - mu) * mx.rsqrt(var + self.eps) + if self.affine: + f = f * self.weight + self.bias + return f.astype(dt) + + +class SparseGroupNorm32(nn.Module): + """GroupNorm over (channels-in-group x voxels), per batch item. fp32 internally.""" + + def __init__(self, num_groups: int, num_channels: int, eps: float = 1e-5): + super().__init__() + if num_channels % num_groups != 0: + raise ValueError(f"{num_channels} channels not divisible by {num_groups}") + self.num_groups = num_groups + self.num_channels = num_channels + self.eps = eps + self.weight = mx.ones((num_channels,)) + self.bias = mx.zeros((num_channels,)) + + def __call__(self, x: SparseTensor) -> SparseTensor: + dt = x.feats.dtype + g, c = self.num_groups, self.num_channels + parts = [] + for sl in x.layout: + f = x.feats[sl].astype(mx.float32) # [n_b, C] + n_b = f.shape[0] + if n_b == 0: + parts.append(f) + continue + # -> [G, (C/G)*n_b] so mean/var cover channels *and* voxels in the group + grouped = f.T.reshape(g, (c // g) * n_b) + mu = mx.mean(grouped, axis=1, keepdims=True) + var = mx.var(grouped, axis=1, keepdims=True) + grouped = (grouped - mu) * mx.rsqrt(var + self.eps) + f = grouped.reshape(c, n_b).T + parts.append(f * self.weight + self.bias) + out = parts[0] if len(parts) == 1 else mx.concatenate(parts, axis=0) + return x.replace(out.astype(dt)) + + +class SparseSiLU(nn.Module): + def __call__(self, x: SparseTensor) -> SparseTensor: + return x.replace(nn.silu(x.feats)) + + +class SparseGELU(nn.Module): + def __call__(self, x: SparseTensor) -> SparseTensor: + return x.replace(nn.gelu(x.feats)) + + +# ---------------------------------------------------------------- blocks + + +class SparseResBlock(nn.Module): + """norm1(affine) -> silu -> conv1 -> norm2(no affine) -> silu -> conv2 + skip.""" + + def __init__(self, channels: int, out_channels: Optional[int] = None): + super().__init__() + self.channels = channels + self.out_channels = out_channels or channels + self.norm1 = LayerNorm32(channels, affine=True, eps=1e-6) + self.norm2 = LayerNorm32(self.out_channels, affine=False, eps=1e-6) + self.conv1 = SubMConv3d(channels, self.out_channels, 3) + self.conv2 = SubMConv3d(self.out_channels, self.out_channels, 3) + self.skip_connection = ( + SparseLinear(channels, self.out_channels) + if channels != self.out_channels + else None + ) + + def __call__(self, x: SparseTensor) -> SparseTensor: + h = x.replace(self.norm1(x.feats)) + h = h.replace(nn.silu(h.feats)) + h = self.conv1(h) + h = h.replace(self.norm2(h.feats)) + h = h.replace(nn.silu(h.feats)) + h = self.conv2(h) + skip = self.skip_connection(x).feats if self.skip_connection else x.feats + return h.replace(h.feats + skip) + + +class SparseFeedForwardNet(nn.Module): + """Upstream is nn.Sequential(Linear, GELU, Linear), so its checkpoint keys are + `mlp.mlp.0` and `mlp.mlp.2` — index 1 is the activation and carries no weights. + Named `mlp_0`/`mlp_2` here because a Python list with a None hole does not survive + MLX's parameter tree; the loader remaps the dotted indices onto these.""" + + def __init__(self, channels: int, mlp_ratio: float = 4.0): + super().__init__() + hidden = int(channels * mlp_ratio) + self.mlp_0 = nn.Linear(channels, hidden) + self.mlp_2 = nn.Linear(hidden, channels) + + def __call__(self, x: SparseTensor) -> SparseTensor: + return x.replace(self.mlp_2(nn.gelu_approx(self.mlp_0(x.feats)))) + + +def _sdpa_per_batch( + q: mx.array, k: mx.array, v: mx.array, layout_q, layout_kv, heads: int, scale: float +) -> mx.array: + """Full attention inside each batch item. q/k/v are [N, H, D] flattened over batch.""" + outs = [] + for sq, skv in zip(layout_q, layout_kv): + qi = q[sq].transpose(1, 0, 2)[None] # [1, H, n, D] + ki = k[skv].transpose(1, 0, 2)[None] + vi = v[skv].transpose(1, 0, 2)[None] + o = mx.fast.scaled_dot_product_attention(qi, ki, vi, scale=scale) + outs.append(o[0].transpose(1, 0, 2)) # [n, H, D] + return outs[0] if len(outs) == 1 else mx.concatenate(outs, axis=0) + + +class SparseMultiHeadAttention(nn.Module): + """attn_mode='full' only — the sole mode LATO.2's model code instantiates.""" + + def __init__( + self, + channels: int, + num_heads: int, + ctx_channels: Optional[int] = None, + attn_type: str = "self", + qkv_bias: bool = True, + ): + super().__init__() + if channels % num_heads != 0: + raise ValueError(f"{channels} channels not divisible by {num_heads} heads") + self.channels = channels + self.num_heads = num_heads + self.head_dim = channels // num_heads + self.scale = self.head_dim**-0.5 + self._type = attn_type + self.ctx_channels = ctx_channels if ctx_channels is not None else channels + if attn_type == "self": + self.to_qkv = nn.Linear(channels, channels * 3, bias=qkv_bias) + else: + self.to_q = nn.Linear(channels, channels, bias=qkv_bias) + self.to_kv = nn.Linear(self.ctx_channels, channels * 2, bias=qkv_bias) + self.to_out = nn.Linear(channels, channels) + + def __call__( + self, x: SparseTensor, context: Optional[SparseTensor] = None + ) -> SparseTensor: + n = x.feats.shape[0] + h, d = self.num_heads, self.head_dim + if self._type == "self": + qkv = self.to_qkv(x.feats).reshape(n, 3, h, d) + q, k, v = qkv[:, 0], qkv[:, 1], qkv[:, 2] + lq = lkv = x.layout + else: + if context is None: + raise ValueError("cross-attention needs a context") + q = self.to_q(x.feats).reshape(n, h, d) + m = context.feats.shape[0] + kv = self.to_kv(context.feats).reshape(m, 2, h, d) + k, v = kv[:, 0], kv[:, 1] + lq, lkv = x.layout, context.layout + o = _sdpa_per_batch(q, k, v, lq, lkv, h, self.scale) + return x.replace(self.to_out(o.reshape(n, self.channels))) + + +class SparseTransformerBlock(nn.Module): + """Pre-norm self-attention + FFN. Norms are non-affine (ln_affine=False upstream).""" + + def __init__(self, channels: int, num_heads: int, mlp_ratio: float = 4.0): + super().__init__() + self.norm1 = LayerNorm32(channels, affine=False, eps=1e-6) + self.norm2 = LayerNorm32(channels, affine=False, eps=1e-6) + self.attn = SparseMultiHeadAttention(channels, num_heads) + self.mlp = SparseFeedForwardNet(channels, mlp_ratio) + + def __call__(self, x: SparseTensor) -> SparseTensor: + h = self.attn(x.replace(self.norm1(x.feats))) + x = x.replace(x.feats + h.feats) + h = self.mlp(x.replace(self.norm2(x.feats))) + return x.replace(x.feats + h.feats) + + +class SparseTransformerCrossBlock(nn.Module): + """Pre-norm self-attn -> cross-attn -> FFN.""" + + def __init__( + self, channels: int, ctx_channels: int, num_heads: int, mlp_ratio: float = 4.0 + ): + super().__init__() + self.norm1 = LayerNorm32(channels, affine=False, eps=1e-6) + self.norm2 = LayerNorm32(channels, affine=False, eps=1e-6) + self.norm3 = LayerNorm32(channels, affine=False, eps=1e-6) + self.context_norm = LayerNorm32(ctx_channels, affine=False, eps=1e-6) + self.self_attn = SparseMultiHeadAttention(channels, num_heads) + self.cross_attn = SparseMultiHeadAttention( + channels, num_heads, ctx_channels=ctx_channels, attn_type="cross" + ) + self.mlp = SparseFeedForwardNet(channels, mlp_ratio) + + def __call__(self, x: SparseTensor, context: SparseTensor) -> SparseTensor: + h = self.self_attn(x.replace(self.norm1(x.feats))) + x = x.replace(x.feats + h.feats) + ctx = context.replace(self.context_norm(context.feats)) + h = self.cross_attn(x.replace(self.norm2(x.feats)), ctx) + x = x.replace(x.feats + h.feats) + h = self.mlp(x.replace(self.norm3(x.feats))) + return x.replace(x.feats + h.feats) diff --git a/trellis_sparse_mlx/tensor.py b/trellis_sparse_mlx/tensor.py new file mode 100644 index 0000000..6261a01 --- /dev/null +++ b/trellis_sparse_mlx/tensor.py @@ -0,0 +1,159 @@ +"""SparseTensor for MLX — a coords/feats pair, mirroring LATO.2's upstream container. + +Upstream wraps either `spconv.SparseConvTensor` or `torchsparse.SparseTensor`; both are +CUDA-only, which is what blocks LATO.2 on Apple Silicon. Nothing about the *data* needs +CUDA — it is just coordinates plus features — so this is a plain MLX reimplementation. + +Layout matches upstream exactly so weight conversion stays a straight mapping: + coords : int32 [N, 4] -> (batch, z, y, x) + feats : float [N, C] +and rows belonging to one batch item are contiguous (upstream asserts this). +""" + +from __future__ import annotations + +from typing import List, Optional, Tuple + +import mlx.core as mx +import numpy as np + + +class SparseTensor: + """N non-empty voxels, each with a coordinate and a feature vector.""" + + __slots__ = ("feats", "coords", "_scale", "_spatial_cache", "_layout") + + def __init__( + self, + feats: mx.array, + coords: mx.array, + scale: Tuple[int, int, int] = (1, 1, 1), + spatial_cache: Optional[dict] = None, + layout: Optional[List[slice]] = None, + ): + if feats.shape[0] != coords.shape[0]: + raise ValueError( + f"feats/coords length mismatch: {feats.shape[0]} vs {coords.shape[0]}" + ) + if coords.ndim != 2 or coords.shape[1] != 4: + raise ValueError(f"coords must be [N, 4] (batch,z,y,x), got {coords.shape}") + self.feats = feats + self.coords = coords + self._scale = tuple(scale) + # Indice maps are expensive to build and identical for every conv that shares a + # coordinate set — upstream exploits this via spconv's `indice_key`. Same idea. + self._spatial_cache = spatial_cache if spatial_cache is not None else {} + self._layout = layout + + # -- basics --------------------------------------------------------------- + @property + def shape(self) -> Tuple[int, int]: + return (self.batch_size, self.feats.shape[1]) + + @property + def batch_size(self) -> int: + if self.coords.shape[0] == 0: + return 0 + return int(mx.max(self.coords[:, 0]).item()) + 1 + + @property + def layout(self) -> List[slice]: + """One slice per batch item. Relies on batch-contiguity, as upstream does.""" + if self._layout is None: + b = np.asarray(self.coords[:, 0], dtype=np.int64) + counts = np.bincount(b, minlength=self.batch_size) + offs = np.cumsum(counts) + self._layout = [ + slice(int(offs[i] - counts[i]), int(offs[i])) for i in range(len(counts)) + ] + return self._layout + + def replace(self, feats: mx.array) -> "SparseTensor": + """New tensor, same coordinates — so the indice-map cache stays valid.""" + return SparseTensor( + feats, + self.coords, + scale=self._scale, + spatial_cache=self._spatial_cache, + layout=self._layout, + ) + + # -- cache ---------------------------------------------------------------- + def cache_get(self, key: str): + return self._spatial_cache.get(key) + + def cache_put(self, key: str, value) -> None: + self._spatial_cache[key] = value + + def __repr__(self) -> str: + return ( + f"SparseTensor(N={self.coords.shape[0]}, C={self.feats.shape[1]}, " + f"batch={self.batch_size}, scale={self._scale})" + ) + + +def subdivide(x: SparseTensor) -> SparseTensor: + """Upsample ×2 by splitting each voxel into its 8 children (nearest-neighbour). + + Mirrors upstream `SparseSubdivide`: coords are doubled then offset by the unit + cube, features are replicated. Child order is the C-order of nonzero(ones(2,2,2)), + matching upstream's `torch.nonzero`, so replicated features line up identically. + """ + n = x.coords.shape[0] + offsets = np.stack(np.meshgrid(*[np.arange(2)] * 3, indexing="ij"), -1).reshape(-1, 3) + offsets = np.concatenate([np.zeros((8, 1), dtype=offsets.dtype), offsets], axis=1) + + base = np.asarray(x.coords, dtype=np.int32).copy() + base[:, 1:] *= 2 + new_coords = (base[:, None, :] + offsets[None, :, :]).reshape(n * 8, 4) + + new_feats = mx.repeat(x.feats, 8, axis=0) + return SparseTensor( + new_feats, + mx.array(new_coords, dtype=mx.int32), + scale=tuple(s * 2 for s in x._scale), + ) + + +def downsample(x: "SparseTensor", factor: int = 2) -> "SparseTensor": + """Downsample by `factor`, reducing colliding voxels with MAX. + + Upstream's docstring says "average pooling" but the implementation passes + reduce="amax" (the `reduce='mean'` line is commented out). Following the code, + not the docstring — mean vs max here is numerically silent in shape and would + quietly change every downsampled feature. + + Output coordinates come out sorted by the same packed code upstream sorts on, so + rows stay batch-contiguous as SparseTensor requires. + """ + import numpy as _np + + c = _np.asarray(x.coords, dtype=_np.int64).copy() + c[:, 1:] //= factor + + maxs = c[:, 1:].max(axis=0) + 1 + # OFFSET = reversed cumprod, matching upstream's packing + off = _np.array( + [maxs[0] * maxs[1] * maxs[2], maxs[1] * maxs[2], maxs[2], 1], dtype=_np.int64 + ) + code = (c * off).sum(axis=1) + + uniq, inv = _np.unique(code, return_inverse=True) + feats = _np.asarray(x.feats) + out = _np.full((uniq.shape[0], feats.shape[1]), -_np.inf, dtype=feats.dtype) + _np.maximum.at(out, inv, feats) + + new_coords = _np.stack( + [ + uniq // off[0], + (uniq // off[1]) % maxs[0], + (uniq // off[2]) % maxs[1], + uniq % maxs[2], + ], + axis=-1, + ).astype(_np.int32) + return SparseTensor( + mx.array(out), + mx.array(new_coords), + scale=tuple(s * factor for s in x._scale), + )