pixal3d_mrp_mlx/PROFILE.md
m3ultra 202ddde31d Profile: neither planned optimisation is worth doing
The plan was to port the fused Metal spconv kernel and the 15 mx.compile sites down
from trellis-2-mrp-mlx into the shared core. Measured on the full textured cascade,
both are dead ends. PROFILE.md has the numbers; the short version:

BY BACKEND (204.3s total, peak 37.1GB)
  DiT transformer flows   ~173s   ~85%
  sparse conv (decoders)   ~26s   ~13%
  torch/MPS (DINOv3+NAF)   4.1s     2%
  o_voxel native           0.8s   0.4%

The Metal kernel can only touch the 13%. Verified from the checkpoints: the four flow
models contain ZERO 5-D tensors, so they never call sparse conv at all - only shape_dec
(40), tex_dec (40) and ss_dec (20, dense) do. A free 2x on ALL sparse conv returns 6%
of runtime. The kernel is still the right port if peak memory ever matters, which is
its stated prize, but it is not a speed fix.

mx.compile is no better, because the DiT loop is compute-bound rather than
dispatch-bound - the opposite of the launch-latency problem that motivated the fused
gather-matmul in the shared core:

  scaling      25% tokens -> 0.15x,  50% -> 0.36x,  100% -> 1.00x  (super-linear)
  per block    attention 1062 GFLOP 58.8% (O(n^2)), MLP 496, proj 248
               12.7 TFLOP/s achieved at 142 ms/block
  direct test  eager 23.9 ms vs mx.compile 23.7 ms  -> 0.8%, i.e. nothing

Attention already routes through mx.fast.scaled_dot_product_attention. mx.compile also
wants pure array-in/array-out functions while the sparse path threads SparseTensor
objects with Python-side layout, so it would mean restructuring for a measured ~0%.

The real lever is TOKEN COUNT: attention is O(n^2), half the tokens ran 2.8x faster,
and refine_coords already backs the grid off past max_num_tokens. That knob trades
resolution for time honestly.

m4pro cannot run this: peak 37.1GB against 24GB of RAM, and the weights alone are 24GB
before activations. Capacity, not tuning.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-03 19:19:17 +10:00

107 lines
4.3 KiB
Markdown
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

# Profile — and why the two planned optimisations are not worth doing
Measured on m3ultra (M3 Ultra 256GB), `0_img.png`, full textured cascade, MLX 0.32.
The plan going in was to port two things down from `trellis-2-mrp-mlx` into the shared
sparse core: the **fused Metal spconv kernel** and the **15 `mx.compile` sites**. The
profile says neither pays. Recording the numbers so nobody re-opens it on a hunch.
## Stage breakdown
| stage | seconds | peak GB | backend |
|---|---|---|---|
| ss_flow + ss_dec | 18.0 | 6.7 | DiT + dense Conv3d |
| cond_512 (DINOv3) | 1.6 | 6.7 | torch/MPS |
| LR slat_flow | 12.5 | 12.4 | DiT |
| refine (partial decode) | 3.4 | 15.5 | sparse conv |
| cond_1024 + NAF | 0.9 | 15.5 | torch/MPS |
| **HR slat_flow** | **90.9** | 23.6 | DiT |
| shape_dec | 11.6 | 28.4 | sparse conv |
| fdg_to_mesh | 0.8 | 28.4 | o_voxel native |
| cond_tex + NAF | 1.6 | 28.4 | torch/MPS |
| **tex stage** | **63.0** | **37.1** | 51.3 DiT + 11.1 sparse |
| **TOTAL** | **204.3** | **37.1** | |
Splitting the tex stage by hand (it is the only mixed one): **51.3s flow / 11.1s
decoder**, so sparse conv is 17.8% of it.
## By backend
| backend | seconds | share |
|---|---|---|
| **DiT (transformer flows)** | **~173** | **~85%** |
| sparse conv (decoders) | ~26 | ~13% |
| torch/MPS (DINOv3 + NAF) | 4.1 | 2% |
| o_voxel native | 0.8 | 0.4% |
**The four flow models contain ZERO sparse-conv tensors** — verified by counting 5-D
tensors in the checkpoints (700/700 params each, none 5-D). Only `shape_dec` (40),
`tex_dec` (40) and `ss_dec` (20, dense Conv3d) use them. So the two candidate ports
land on completely disjoint parts of the pipeline, and one of those parts is 13%.
## Why the Metal spconv kernel is not the win
It can only touch the ~13% in the decoders. Its own docstring says the prize is **peak
memory, not time** — it was written because 128GB gated the M1 Ultra. Even a free 2x
on all sparse conv would return ~13s of 204s (6%).
It remains the right thing to port **if peak memory ever matters** — 37.1GB is fine on
a 256GB Ultra and impossible on a 24GB M4 Pro (see below).
## Why `mx.compile` is not the win either
The DiT loop is **compute-bound, not dispatch-bound**. Two measurements:
**Scaling is super-linear in tokens** — doubling tokens costs 2.42.8x, which is the
O(n²) attention, not launch overhead:
```
25% tokens ( 3,286) 646 ms 0.15x
50% tokens ( 6,573) 1,532 ms 0.36x
100% tokens (13,147) 4,261 ms 1.00x
```
**Per-block FLOPs confirm attention dominates:**
```
self-attention 1062 GFLOP 58.8% <- O(n^2)
MLP 496 GFLOP 27.5%
qkv+out proj 248 GFLOP 13.7%
total 1806 GFLOP
at 142 ms/block -> 12.7 TFLOP/s achieved
```
12.7 TFLOP/s is a respectable fraction of what an M3 Ultra delivers on bf16 weights
with fp32 accumulation, and attention already routes through
`mx.fast.scaled_dot_product_attention` — the fast path.
**Direct test on the compilable part** (the norm+MLP chain, pure arrays):
```
eager 23.9 ms
mx.compile 23.7 ms <- 0.8%, i.e. nothing
```
There is also a structural obstacle: `mx.compile` wants pure array-in/array-out
functions, while the sparse path threads `SparseTensor` objects carrying Python-side
layout. Compiling the blocks would mean restructuring that, for a measured ~0%.
## What WOULD move the needle
1. **Fewer tokens.** Attention is O(n²), so token count is the dominant lever — 50%
of the tokens ran 2.8x faster. `refine_coords` already backs the grid off when the
count exceeds `max_num_tokens` (49152); lowering that knob is the real speed
control, and it trades resolution for time honestly.
2. **Quantization**, for memory and possibly bandwidth. Not attempted. Treat with
suspicion: a published field report on ARDY found INT8 preserved embedding cosine
at 0.992+ while the generated output diverged badly, so it would need output-level
validation (the silhouette gate is well suited to that), not a similarity metric.
3. **The Metal spconv kernel — for memory only**, if this ever needs to run somewhere
smaller than a Studio.
## m4pro cannot run this
Peak is **37.1 GB**. The m4pro is 24 GB, so the full cascade will not fit, and the
weights alone are 24 GB before activations. Geometry-only peaks lower (~28 GB) and
still does not fit. Not a tuning problem — a capacity one.