lab

401 → 163 ms: Stable Audio 3 medium on one TPU v6e chip

Stable Audio 3 medium, a JAX port, bf16, a 120 s clip, batch 1, 8 sampling steps, text encoder and audio decoder included, on one Colab TPU v6e-1. From a frozen baseline of 401.4 ms to 163.4 ms — 2.46×, 1.26× faster than Stability's own TensorRT fp8 engine on an RTX PRO 6000 (205.2 ms). A full 380 s song: 1766 → 674 ms, also ahead of TensorRT fp8 (693 ms). Every kept change passed a quality gate against the official PyTorch output. Numbers are warm p50s around jax.block_until_ready(), each A/B in one process.

1. Result

stage 120 s clip diffusion transformer decoder vs baseline
frozen baseline 401.4 ms 160.4 230.0 —
fused decoder attention 256.0 ms 158.9 90.0 1.57×
decoder tiling + cached cross-attention keys 228.8 ms 153.8 70.8 1.75×
fused transformer attention 198.5 ms 135.5 61.9 2.03×
keys prepared once per head, both heads as one operand 163.4 ms 99.6 60.9 2.46×

Latency per stage on one v6e-1; dashed: engines on an RTX PRO 6000.

Where the time went, baseline → now.

For scale, on an RTX PRO 6000 (different hardware — an implementation reference): TensorRT fp16 274.2 ms, TensorRT fp8 205.2 ms, official PyTorch 474.7 ms, diffusers 981 ms. TensorRT renders 5 % fewer latents for the same 120 s.

Latency against clip length.

2. A gate before any speed

Tokens are not images: audio quality needs its own instrument. 16 prompts × 2 seeds, 120 s each, with all four random inputs injected identically into both implementations (initial latent, the sampler's per-step re-noise, the decoder's dither and mask noise) — seeds cannot pair runs across PyTorch's and JAX's random generators. Scored with multi-resolution STFT and log-mel distances and MERT embeddings. Thresholds were set from the official code disagreeing with itself (swapping its attention kernel), never picked.

The gate paid for itself before optimisation began. The port scored halfway to "different music". Fed identical inputs in fp32, one transformer forward matched PyTorch to 88.9 dB and the decoder to 99.0 dB — the port was right; precision was not. Op-by-op bisection found it: the decoder's DyT normalisation in bf16. Its trained gains are large and its softmax is nearly an argmax, so four bf16 roundings per call flip which key wins; fp32 DyT alone took decoded audio from 17 to 36 dB SNR.

Quality per stage against the official output.

3. The decoder was moving data, not computing

The baseline decoder spent 230 ms per clip, of which its matmuls needed ~38. The rest was glue: sliding windows built by padding and concatenating, rotary embedding by slice-and-negate, head transposes, a 17 × 51 softmax padded to the TPU's 8 × 128 tiles — each op a full round trip to HBM, 23 086 tokens × 12 blocks.

One Pallas kernel replaced it (230 → 90 ms). It reads the attention projection once, does normalisation, rotary embedding, both differential attention branches and their subtraction in on-chip memory, and writes one result. No windows: a query sees ±17 keys, so with ≥ 17-row blocks every key lies in the previous, current or next block — three block reads, masked by position. No transposes: the block index maps pick heads straight out of the projection. Two 64-dim heads share one 128-lane tile by masking lanes, keeping every op a plain 2-D tile op.

4. Fewer FLOPs was slower

The first kernel multiplied 768 keys per query to use 35. Cutting that to 96 made it 2× slower.

Decoder time against kernel tiling.

A TPU core pairs a 256 × 256 matrix unit with a vector unit that pays per op. Small slices meant many small ops — mask, exp, max, sum, divide — each too little work to amortise. The fastest tiling was the one with the largest slices (90 → 62 ms), plus two trims that remove whole passes over the scores: fold 1/√64 into the query before the matmul (a power of two — exact), and divide by the softmax sum after multiplying by the values, on 128 lanes instead of every score.

5. The transformer: one hoist, one kernel, one cliff

Cached cross-attention keys. The text conditioning is fixed while sampling, yet its keys and values were recomputed in 24 blocks × 8 steps. Computing them once per clip is loop-invariant code motion — bit-identical. It first made things slower: the hoisted work sat in an eager function, ~100 separately dispatched ops. One jit fixed it.

A fused attention kernel — the decoder idea without the band. A head's 1422 keys fit on chip, so one program holds the whole score row: normalise, rotate, both softmaxes, subtract. It keeps scores in fp32 where the compiler rounded them to bf16, and passed the gate closer to the official output than the baseline.

Transformer time against sequence length.

At first it was not faster than the compiler in general: ~20 % slower up to 90 s, winning only because the compiler's attention falls off a cliff between 1162 and 1422 tokens (1.83× time for 1.22× tokens).

The fix was hiding in the grid order. Each kernel program handles one head pair and one block of 256 queries; the query block was the innermost loop. Keys and values do not depend on the query block — yet every program re-normalised and re-rotated all 1422 keys, for both branches, six times per head pair. By op count that was as much vector work as the softmax itself. Preparing them once per head pair, in the first program, into on-chip scratch that the next five read, cut the transformer from 133.5 to 109.3 ms. Running both heads of a tile as one stacked operand, and using a single 1424-row query block at 120 s, took it to 99.0 ms. The kernel now wins from 838 tokens and ties at 452; the sequence length is static, so the choice is made per clip length at compile time.

6. What the chip decided

This is a runtime for one device, and most of its decisions would be different, or wrong, on a GPU.

v6e fact what it decided
Registers are (8 × 128) tiles; a kernel block's last two dims must be multiples of (8, 128) or the whole array 64-dim heads cannot be addressed alone: two share a 128-lane tile and are separated by masking lanes. The fused q/k preparation kernel that read 64-lane slices would not compile. A block the full length of the sequence (1422) needs no padding.
Matrix unit: a 256 × 256 systolic array per core Small tiles starve it. Slicing the decoder's attention to cut wasted FLOPs made it 2× slower; the fastest tilings were the largest.
Vector unit and matrix unit are separate, and the vector unit pays per operation Once memory traffic was gone, cost tracked the number of vector ops, not FLOPs. Hence: fold the softmax scale into q, normalise after P·V, run two heads as one operand, compute K once. bf16 exp changed nothing.
128 MiB of on-chip VMEM per core (a GPU SM has ~0.2 MiB) A head's full K and V (1422 × 64) fit, so the transformer computes each softmax row in one pass. Flash attention's streaming, built for memory that does not fit, was pure overhead — and slower.
One TensorCore per chip; a kernel grid runs in order K, K_diff and V can be prepared by the first program of a head pair and read from scratch by the rest — the change worth 24 ms. On a GPU those programs run concurrently and cannot share scratch this way.
XLA compiles one program per shape, sampler loop included No host dispatch: the 8-step sampler is one contiguous device loop, where PyTorch launched ~30 000 kernels per request. And the choice of attention kernel is made per length at compile time, for free.
XLA's own attention falls off a cliff between 1162 and 1422 tokens (1.83× the time for 1.22× the tokens) The fused kernel first won only past that cliff; after the K cache it ties at 452 tokens and wins from 838.
An fp32 matmul at default precision is a bf16 matmul on the MXU; highest is 3-pass The decoder's "fp32" scores were already bf16 inputs; the kernel reproduces that on purpose. highest cost +30 % in the decoder and bought no quality, so it is off.
int8 runs at 2× bf16 on the MXU, but quantising activations is vector work int8 gained 2 % and failed the gate.
Pallas kernels ship as Mosaic IR, versioned against libtpu Colab's image (libtpu 0.0.21.1) could not load a jax 0.7.2 kernel; every session pins 0.0.23.

7. What did not work

bf16 against int8 per stage.

8. Where the time goes now

At 120 s: transformer 99.6 ms, decoder 60.9, conditioning 1.7. The bf16 matmuls run at 77–100 % of the chip's measured rate and are now the larger share of both stages; what remains in the kernels is vector work on the attention scores. At batch 1, SA3 on v6e was never limited by its matrix unit: first by HBM round trips between ops, then by repeated vector work.

9. In context

One chip, one clip at a time: ~22 000 two-minute clips or ~5 300 full 380 s songs per hour (≈ 734 hours of audio per wall-clock hour). At Google Cloud's on-demand 2.70perv6echip − hourthatis 0.0005 per full song; Stability's API charges $0.26 per generation.

implementation hardware 120 s clip
this engine 1× TPU v6e 163 ms
TensorRT fp16 (Stability's README) H100 ~150 ms, 5 % fewer latents
TensorRT fp8 (measured) RTX PRO 6000 205 ms
TensorRT fp16 (measured) RTX PRO 6000 274 ms
official PyTorch (Stability's paper) H200 780 ms

Per latent, Stability's H100 TensorRT figure and this engine are a tie (0.116 vs 0.120 ms); its fp8 tier on an H100 or newer is likely faster, and unpublished. No hosted API publishes a latency. For a different kind of model: YuE2-3B, which writes lyrics and vocals autoregressively, renders a 3.6-minute song in ~71 s on an RTX 4090 and ~373 songs per hour on an H800 at 32 concurrent requests.

10. Method, transferable

  1. Build the quality gate first, from the reference disagreeing with itself.
  2. Attribute device time to source lines before touching code.
  3. A/B in one process; toolchains move numbers.
  4. Measure the thing you believe — here, that fewer FLOPs and cheaper arithmetic would help. Both were wrong.
  5. Read the kernel's grid order: the largest single kernel win was work repeated in every program.

Appendix: optimization log (teaching notes)

A teaching log written alongside the work, meant to be mined for the article. Each entry: what the trace said, the idea, how it is built, what could go wrong, what was measured. Numbers live in metrics.md (frozen baseline) and ../metrics-progress.md (one row per stage).

0. Method (applies to every entry)

E1 — fused Pallas attention for the SAME-L decoder

What the trace said. Decoder 227 ms on device, of which the four big matmuls per block took only ~38 ms (≈490 TFLOP/s, 76 % of the chip). The other ~190 ms:

source what it is ms
_windows (same_l.py:127) building K/V windows: pad, 3 shifted slices, concatenate 51
_sdpa (:143-146) (…,17,51) logits, mask, softmax, P·V, differential subtract 71
_rope (:98-99) slice / negate / concatenate for rotate-half 29
_dyt (:68) q/k DyT in fp32 + the norm sites 27
transposes, split, merge head layout changes 10

Why that is slow on a TPU. A TPU core works on tiles of 8 rows × 128 lanes. The band attention's shapes are 17 queries × 51 keys × 64 head dims: 17 pads to 24 rows, 51 to 128 lanes, 64 to 128 lanes, so most of every tile is padding. Worse, each step (windowing, logits, mask, softmax, weighted sum) is a separate XLA op that writes its whole result to HBM and the next op reads it back. At 23 086 tokens × 1536 channels × 12 blocks those round trips — not arithmetic — are the 190 ms.

The idea: one kernel, one read, one write. Read the to_qkv projection once, do DyT → RoPE → band attention for both differential branches → subtract inside the chip's on-core memory (VMEM, ~128 MiB on v6e), and write only the final (tokens, 1536) result. A hand-written kernel is the only way to say "keep this in VMEM" — XLA's fuser will not fuse across the attention matmuls.

How it is built (sa3jax/band_attention.py, Pallas, the JAX kernel language for TPU/GPU):

Correctness before TPU time. Pallas has an interpret mode that runs the real kernel plumbing (block specs, index maps, grid) on CPU. Against the XLA path on a small config: 51.3 dB, identical for row tiles 24/32/64 — the bf16 output floor, and tile-independence proves the band edges and padding masks.

What went wrong first, and why. The first TPU run failed: "Failed to deserialize the Mosaic module: expected ≤ 7 but got 8". Colab's image pairs jax 0.7.2 with libtpu 0.0.21.1, but jax 0.7.2 declares libtpu 0.0.23. XLA-only programs tolerate the older runtime; Pallas kernels are serialised in Mosaic IR, whose version the runtime must understand. Pinning libtpu 0.0.23 for the session fixed it (flicker hit the same wall). Consequence: the XLA baseline is re-measured on 0.0.23 in the same session.

Measured (v6e-1, 120 s, decode only, same process, libtpu 0.0.23). XLA 229.4 ms → Pallas 153.0 (64 rows) / 109.4 (128) / 88.4 (256) / 98.5 (512). Bigger row tiles mean fewer grid steps (each costs ~0.35 µs of fixed overhead and a DMA setup) but more wasted attention work, since a query only needs 35 of the 3×rows keys it multiplies against; 256 is the sweet spot. Pallas vs XLA audio: 34.2 dB — not bit-identical, because summation order and rounding points inside the attention differ, and the decoder amplifies small differences through 12 blocks of large-gain DyT and near-argmax softmax (the same mechanism that made bf16 DyT cost 20 dB). Whether 34 dB is acceptable is the gate's call, against the official output. Gate and end-to-end numbers: ../metrics-progress.md.

End to end (same session, libtpu 0.0.23). 120 s: 400.9 → 256.0 ms (1.57×; DiT 158.9, decode 90.0). 380 s: 1792.5 → 1363.7 ms. Gate A: PASS on all six rows, each at or slightly better than the frozen baseline (MR-STFT median 1.113 vs 1.114, MERT-v2 worst pair 0.936 vs 0.933). So the 34 dB audio difference from the XLA decode is the same kind of drift the port already carries against the official fp16 output, not a quality loss: the kernel moved sideways, not away from the reference. This is why the gate measures distance to the official output rather than to our own previous version.

Where it leaves us. Decode is 90 ms against a ~29 ms matmul floor; the DiT (159 ms) is now the largest stage, and 63 % of it is its own unfused attention — the next target (D1–D3).

E2 — the decoder kernel, second pass: stop multiplying keys nobody needs

What the E1 trace said. Decode 87.6 ms: kernel 41.3 ms (3.4 ms/block against a ~0.5–1 ms floor), matmuls 37.5 ms (~494 TFLOP/s, 77 % of the chip's measured GEMM rate — near its bf16 ceiling), and a 7.4 ms pad that copied the whole 354 MB projection so S would divide the row tile.

Why the E1 kernel wasted work. Each program took 256 queries against 3 × 256 = 768 keys, but a query only reaches 35 of them, so ~95 % of every logit, exp and P·V product was on masked entries. It also ran DyT + RoPE on the previous and next K blocks for every program, i.e. K prep three times over.

The changes.

Correctness: interpret mode on CPU, S = 391 (ragged), four (rows, halo, sub) tilings, all 51.3 dB vs XLA — identical to E1. TPU numbers: next session (swept rows:sub ∈ 256:64, 512:64, 256:32, 256:128).

Roofline after E1 — where the remaining time is (docs/roofline-ops.md, scripts/roofline_ops.py)

Per-op floors = max(FLOPs / peak, bytes / bandwidth) at the chip's measured 643.7 TFLOP/s and ~1400 GB/s (datasheet 918 / 1640 in the table too); attention is charged its fused traffic (logits never leave VMEM). Whole clip: 54.8 TFLOP, floor ~128 ms, measured 239 ms on device → ~139 ms of headroom.

D1–D3 — the DiT: stop recomputing, stop materialising, stop round-tripping

The DiT runs 24 blocks × 8 sampler steps = 192 block evaluations per clip, over 64 memory tokens + 1358 latents = 1422 tokens, with differential attention (two softmaxes per attention) and a cross-attention to 257 constant conditioning tokens. Three separate wastes, three changes, each a static config switch so the TPU A/B compares them in one process.

D2 — hoist cross-attention K/V (always on; bit-identical). to_kv(ctx) and the K RMSNorm ran in every block of every step, though ctx never changes during sampling. prepare_static now computes, per block, the normalised (k, k_diff, v) in heads layout once per clip; the block only projects its query. 192 → 24 evaluations of a 257 × 1536 → 4608 projection: ~2.7 GB of weight reads and ~0.7 TFLOP gone. This is loop-invariant code motion — the same arithmetic, moved out of the loop, so the result is bit-identical (verified: max |Δ| = 0 in bf16 and fp32). The roofline called it the one free change.

D1 — flash attention for self-attention (DiTConfig.attention = "flash"). The XLA path computes QKᵀ into a (24, 1422, 1422) tensor per branch, writes it to HBM, reads it back for an fp32 softmax, writes the weights, reads them for P·V. That round trip is ~19 GB per step; the roofline puts it at ~84 ms of the DiT's 150. Flash attention (Pallas' TPU kernel) tiles Q and K/V, keeps a running max and running sum per query row (the "online softmax"), rescales the partial output as each new K/V tile arrives, and never writes the score matrix at all. Two details matter here:

D3 — fused q/k prep (DiTConfig.qk_prep = "pallas"). Between to_qkv and attention, XLA ran split → 5 head transposes → 4 RMSNorms → 4 RoPEs → the V mask, each an HBM round trip at 1422 × 7680 × 192. The kernel (sa3jax/qk_prep.py, adapted from flicker's D3c) reads the projection once and writes q, k, v, q_diff, k_diff already in the (B, H, L, D) layout flash consumes: the output index map (b, h, i, 0) makes the DMA engine perform the head transpose for free. One head per 64-lane tile (not two per 128) — because only a one-head tile can be written straight into (B, H, L, D); lane-slicing a two-head tile would need a gather. Measured on CPU (interpret mode): within 0.01 bf16 ulp of the XLA chain.

Risks the TPU decides. head_dim 64 in the flash kernel on jax 0.7.2 (supported in its source, untested here); D3's and D1's pads are separate copies (a possible follow-up is one shared pad); flash's fp32 logits.

E2 result — fewer FLOPs was slower (session 3, same process, decode only, 120 s)

tiling (rows × query slice) keys per query decode ms
E1 (256, no slicing) 768 88.4 (session 2)
E2 256 × 32 96 186.1
E2 256 × 64 128 120.1
E2 512 × 64 128 117.0
E2 256 × 128 192 85.0

Cutting the attention work 6× made the kernel up to 2× slower. The lesson is about what a TPU core is: the MXU is a 256 × 256 systolic array that wants large, square-ish operands, and every op on the vector unit (masking, exp, max, sum, select) pays per-op overhead that small tiles cannot amortise. A 64 × 128 slice is a sliver of MXU work wrapped in five or six vector ops on a tiny tile; the E1 version did "wasted" multiplies at full MXU efficiency. The roofline agent's reading was right: the kernel was not FLOP-bound, so removing FLOPs could not help. The halo and no-pad changes still pay (256 × 128 beats E1), and the remaining lever is the other direction — larger slices (the 320-key buffer in one piece).

D1 and D3 on TPU — both failed, for reasons worth knowing (branch opt/d1-d3-attempt)

D4 — a fused dense attention kernel for the DiT (the E1 idea, without the band)

Because a head's entire K/V fits in VMEM, there is no need for flash's online softmax: one kernel program per (head pair, 256-query block) holds all 1422 keys and computes the full softmax row at once. It reads the to_qkv projection once — q/q_diff as 256-row blocks, k/v/k_diff as whole-sequence blocks (legal: a block dim equal to the array dim needs no 8-multiple, so no padding and no key masking) — and with the query block the innermost grid axis, each pair's K/V is DMA'd once and reused across all six query blocks. Inside: per-head RMSNorm in a two-head tile (two masked lane sums), RoPE by lane roll, V times the padding mask, both differential softmaxes with scores in VMEM, subtraction in fp32, merged output. Its lane RoPE tables and lane mask are built once per clip in the jitted static prep. CPU check (interpret): 103 dB against XLA with fp32 logits (i.e. exact), 43.8 dB against XLA as written — the difference is XLA rounding the scores to bf16, which torch's kernels (the reference) do not do.

Session 4 — bigger slices, the jitted static prep, and D4's first number

Why D4 barely helped — and what that says about the TPU. Per call the fused kernel spends ~0.45 ms; its matmuls need ~0.04 ms. The rest is the vector unit walking the (rows × 1422) score matrix per head per branch: scale, max, subtract, exp, sum, divide, cast — ~97 M elements × 6-7 ops per call on a unit that does ~1 K lanes per cycle. XLA's softmax does the same element work, so fusing it only removed the HBM traffic, which on this chip was not the binding cost. Attention here is VPU-bound, not MXU- or HBM-bound, and the lever is fewer vector ops per score:

  1. fold 1/√64 = 2⁻³ into q before the MXU (a power of two — exact in bf16);
  2. normalise after P·V, so the division runs on 128 output lanes instead of 1422 score lanes;
  3. (session 6) exp in bf16: v6e's VPU does bf16 natively at ~2× the fp32 rate, and the weights are rounded to bf16 for the MXU anyway (XLA's reference path already rounded the scores to bf16). 1 and 2 are in both kernels now (CPU: band 55.3 dB vs XLA — slightly better than before, because the normalised value is rounded once instead of every weight).

Sessions 5–6 — the trims, the fused DiT end to end, and a null result

Session 7 — where the fused DiT kernel actually wins

DiT 8-step ms, same process (XLA / fused): 30 s 452 tok 32.3 / 40.6 · 45 s 645 tok 45.5 / 54.1 · 60 s 838 tok 61.2 / 72.4 · 90 s 1162 tok 83.3 / 99.4 · 120 s 1422 tok 152.7 / 133.5 · 380 s 4160 tok 1062 / 710. The kernel is ~20 % slower than XLA wherever XLA scales smoothly; it wins only because XLA falls off a cliff between 1162 and 1422 tokens (1.83× time for 1.22× tokens — consistent with its materialised score tensors no longer staying on chip). So fused_min_tokens = 1400, and the kernel has real headroom left: its vector work per score is the target, not HBM traffic.

Q1 — int8 W8A8 (opt-in tier, sa3jax/quant.py): the 2× MXU does not reach the clock

Same session, 120 s: DiT 135.5 → 127.5 ms (−8), decoder 61.9 → 67.9 (+6), request 198.5 → 194.2 ms (−2.2 %); 380 s 905.6 → 902.8. The int8 dot runs at twice the bf16 rate, but dynamic per-token activation quantisation (absmax reduce, divide, round, clip before every GEMM) is vector-unit work — the resource this chip is short of in every other experiment too — and in the decoder it costs more than the MXU saves. At batch 1, SA3 on v6e is VPU-bound around its matmuls, not MXU-bound. Static (calibrated) activation scales would remove the reduce but not the per-element scale/round; flicker's SmoothQuant bought 16 % of its DiT at ~2× its quality floor.

Q1 gate: FAIL — MERT-v2 median 0.9866 (< 0.9893), worst 0.9261, MERT-v1 0.9866; MR-STFT 1.357 (baseline 1.114). int8 is off.

V1 — vector work inside the kernels (int8 removed)

int8 is gone from the code (3df1071): 2 % faster and it failed the gate. What is left, per the r5 trace at 120 s (device ms): DiT 131.9 = fused self-attention kernel ~77–82 + GEMMs 33.5 (at ceiling) + cross-attention (still XLA) 11.6 + small ops; decoder 59.6 = band kernel 19.9 + GEMMs 38.2 (at ceiling). The DiT kernel is the target: its useful matmuls need ~7.4 ms at peak, yet it spends ~5.6 µs in each of 13 824 programs per clip.

Waste found by reading the kernel, not the trace: K was re-prepared per query block. The grid is (batch, head pair, query block) with the query block innermost. K, K_diff and V do not depend on the query block — but RMSNorm + RoPE of all 1422 keys (both branches) and the V mask ran in every program, six times per head pair. By op count that is about as much vector work as the score softmax itself. Fix: prepare them once per head pair, at query block 0, into VMEM scratch (scratch_shapes), and declare the query-block axis "arbitrary" so it runs in order on one core (block 0 must fill the scratch before the others read it). Bit-identical on CPU.

Fewer, bigger ops: stack the two heads. A 128-lane tile holds two 64-dim heads, so each branch did two matmuls, two softmaxes and two P·V products on (rows × L). Stacking the two masked query copies as one (2·rows × 128) operand gives one of each on (2·rows × L): the same element work in half the instructions — the lesson of E2, where per-op overhead, not arithmetic, set the cost. Knob stack, both kernels; CPU: DiT bit-identical, decoder within one bf16 ulp.

Any query-block size. The RoPE tables are now exactly L rows and the last block is ragged, so the sweep can try 256 / 512 / 712 (two blocks) / 1424 (one block, which also makes the K cache moot). A 100 MiB VMEM limit (v6e has 128 MiB) leaves room for the (2·rows × 1536) fp32 score matrix of the big blocks.

Session 8 — the K cache was the big one

DiT 8-step ms at 120 s (1422 tokens), same process: compiler 153.2 · previous kernel 133.5 · K prep cached, 256-row blocks 109.3 · + stacked heads 105.1 · 512-row 103.3 · 712-row stacked 100.7 · one 1424-row block, stacked 99.0. Caching K alone was worth 24 ms — confirming that re-preparing it per query block was about half the kernel's vector work. Stacking bought a further 4 %, the single block (no cache needed, one softmax per branch) 6 % more.

The new kernel also moved the crossover: 90 s (1162 tokens) 83.3 compiler vs 82.7 fused (was 99.4); 60 s (838 tokens) 61.0 vs 59.3. The best block size depends on length — at 1162 tokens 256-row beats 712-row (82.7 vs 87.8) — so the default is one block for 1400–1536 tokens and 256 rows otherwise, and the fused path now engages from 800 tokens. Decoder: stacking 60.5 → 59.4 ms (its kernel is only ~20 of the 60; the rest is GEMMs at their ceiling). Stacking is now on by default in both kernels.

V1 end to end (r8, committed defaults). 5 s 27.8 · 30 s 53.0 · 60 s 96.0 · 120 s 163.4 (DiT 99.6, decoder 60.9) · 380 s 673.5 ms — now also ahead of TensorRT fp8 on the RTX PRO 6000 at the longest length (692.5), which the previous version (905.6) was not. Gate: pass, best so far — MR-STFT 1.087, MERT-v2 0.9935 (worst 0.9367), MERT-v1 0.9917. A later 30 s check: compiler 32.6 vs kernel 32.2 ms — a tie, so the 800-token threshold could drop to ~450 for ~1 %.