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× |


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.

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.

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.

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.

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
- Library flash attention: wrong output and slower. Flash streams keys that do not fit on chip; here they do.
- A fused q/k preparation kernel: would not compile — it read 64-lane slices, and TPU blocks must be multiples of 128 lanes or the whole array.
- exp in bf16: no change in either kernel.
- int8 matmuls: 2 % faster (the transformer gained 8 ms, the decoder lost 6 — quantising activations per token is vector work, which this chip is short of) and it failed the gate: MERT 0.9866 against a 0.9893 floor, spectral distance 1.11 → 1.36. Removed.

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
- Build the quality gate first, from the reference disagreeing with itself.
- Attribute device time to source lines before touching code.
- A/B in one process; toolchains move numbers.
- Measure the thing you believe — here, that fewer FLOPs and cheaper arithmetic would help. Both were wrong.
- 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)
- Profile, attribute, then act. An XProf trace of the
TPU gives every device op with its duration. Each op is attributed to
the innermost Python source line whose "Source code" event covers it, so
"the decoder is slow" becomes "
same_l.py:127spends 51 ms incopy". The baseline split was derived this way, never guessed. - Roofline first. For each stage, FLOPs ÷ measured GEMM rate gives a floor. The baseline decoder does ~18.5 TFLOP in 230 ms; at the chip's measured 644 TFLOP/s that is a ~29 ms floor, 8× away. When a stage is that far from its floor, the time is not in the math — it is in memory traffic or overhead.
- Same-session A/B only. Toolchains move numbers (libtpu 0.0.21 → 0.0.23 changes the compiler), so every candidate is timed next to the XLA path in the same session and process. The frozen baseline is re-measured on the new toolchain for the progress table.
- Quality is judged against the official output, not against ourselves. The gate (metrics.md §5) compares a candidate's 32 renders to the official fp16 renders. Comparing to our own previous output would fail any change that merely reorders a floating-point sum.
- Batch the TPU. One Colab session runs several candidates as one detached chain that ships every result to HF before teardown; gates run only when a candidate is a keeper.
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):
- Grid
(batch, row block i, head pair h). Each program computes one(rows, 128)output tile.his the innermost loop, so data that depends only oni(the RoPE tables) is DMA'd once peri: Pallas skips a copy when a block's index does not change between consecutive programs. - No transposes.
to_qkv's output is[q | k | v | q_diff | k_diff]along the channel axis, each 12 tiles of 128 lanes. ABlockSpecindex map picks "section s, tile h" directly, and the output tile is written straight into the merged(tokens, heads·64)layout. The DMA engine does the layout work. - No windows. A query at position p sees keys in [p−17,
p+17]. If a row block has ≥ 17 rows, every key a block needs lies in the
previous, current or next block. So K, V and K_diff are each passed
three times with index maps
i−1,i,i+1(clamped at the ends), and the logits come in three(rows, rows)pieces joined by one softmax (a shared max, then exp, then one sum). The band mask is computed from nominal positions, so a clamped duplicate block and the zero padding past the last token are masked out instead of treated as keys. - Two heads per tile. head_dim is 64, a tile is 128 lanes. To get head A's scores, head B's lanes of q are zeroed: the 128-lane dot product then sums exactly A's 64 dims. The output takes A's lanes from A's result and B's from B's. This doubles the (tiny) attention matmuls but every op stays a plain 2-D tile op — the Mosaic lowering rule flicker learned the hard way (lane slicing became gathers there).
- RoPE by lane rotation. rotate-half pairs dim d with d±16.
pltpu.rollrotates lanes natively; a select picks the right direction per lane, and the sign is folded into a precomputed "signed sin" table, so rotary is two multiplies and an add. Non-rotary dims get cos=1, sin=0 and pass through. - Numerics, cast for cast. DyT and RoPE in fp32 (the fp32 DyT
island from the baseline fix); q/k rounded to bf16 for
QK^Twith fp32 accumulation — exactly what XLA:TPU already did to the fp32 dot atdefaultprecision; fp32 softmax; normalised weights to bf16 forP·V; subtraction in fp32; one cast.
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.
- Halos instead of neighbour blocks. A second, smaller
BlockSpec brings in only the 32 rows before and after the block (32 ≥
17, a multiple of 8, divides 256). Index maps are in units of the block
size, so "the 32 rows before block i" is halo index
i·8 − 1. The key buffer becomes 320 contiguous positionsi·256 − 32 + r, prepped once. - Query slices. Queries go in 64-row slices, each against the
128 keys that can reach it — a static, aligned slice of the 320. Keys
per query drop 768 → 128 (6×), and so do
expandP·Vwork. - Ragged grid, no pad. The grid is
ceil(S / 256); Pallas reads undefined rows past the end of the last block and drops their writes. Garbage keys are masked by position, but V needs care: a masked weight is exactly 0, and0 × NaN = NaN, so V rows outside[0, S)are zeroed explicitly.
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.
- DiT attention: 94 ms measured vs ~10 ms fused floor — 84 ms,
the biggest item. Not a FLOP problem: the XLA path writes and
re-reads a
(24, 1422, 1422)score tensor per branch per block per step (~19 GB of HBM per step); measured QKᵀ / softmax / PV each sit near their as-written floor. Flash attention is the lever (D1). - Decoder band kernel (E1): 41 ms vs 3.7 ms. Memory-bound in the model, but the kernel does 7× the needed work (768 keys for a 35-key band, doubled by two heads per lane tile) and preps K three times. E2 addresses the first two.
- GEMMs have no bf16 headroom. Decoder linears run at 79–81 % of the measured GEMM rate, the DiT's at or above it. The only lever left is int8 (v6e: 1836 TOPS, 2×). Flicker's SmoothQuant W8A8 on the same chip bought 16 % of its DiT at ~2× its LPIPS floor — so for SA3 int8 is an opt-in tier, tried last.
- No host gaps. The 8-step sampler is one contiguous
149.5 ms
whileon device; dispatch is not a bottleneck, so CUDA-graph-style tricks have nothing to win on TPU.
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:
- Padding. The kernel needs the sequence to be a multiple of
its tile (1422 → 1536). Padded keys must not enter the softmax, so real
tokens get segment id 1 and padding 0; the kernel masks cross-segment
pairs. This is not the same as the model's own padding, which the
reference handles by zeroing V for padded latents while leaving those
keys in the softmax — that
v_maskis kept exactly as it was. - Numerics. XLA rounded the logits to bf16 before its fp32 softmax; flash accumulates them in fp32 and never rounds. That is more accurate, but it is a real change, so the gate decides.
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)
- D3 fused q/k prep would not lower: "the last
two dimensions of your block shape [must be] divisible by 8 and 128
respectively, or equal to the respective dimensions of the overall
array." It read one 64-lane head at a time out of the 7680-lane
projection — 64 is neither a multiple of 128 nor the array width. (The
output
(B, H, L, 64)was legal, because there 64 is the full width.) The rule is a property of the TPU's (8, 128) register tiles and the DMA engine, not a Pallas quirk. - D1 library flash attention ran but was wrong (5.9 dB against XLA) and slower (267 vs 212 ms). Wrong: most likely the segment-id padding path with head_dim 64 on jax 0.7.2's kernel; not pursued, because slower mattered more. Slower: flash is built for sequences that do not fit on chip — it streams K/V tiles with an online softmax. Here the whole K/V of one head is 1422 × 64 bf16 ≈ 180 KB, so the streaming machinery is pure overhead, and head_dim 64 half-fills the contraction.
- D2 made the XLA DiT slower (159 → 212 ms) — not the
hoist itself but where it went:
prepare_staticwas eager, so the hoisted cross K/V ran as ~100 separately dispatched ops per clip. Fix:prepare_staticis now an eager value check plus one jitted_static(params as arguments).
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
- Decoder (decode only, same process): 256×128 84.8, 512×128 83.0, 256×256 73.8, 512×256 69.4 ms (XLA 228.1). The trend reverses E2's premise completely: the fastest tiling is the one with the largest per-slice work. E2's r3 lane (256×128) passed the gate at baseline quality.
- DiT: XLA with the jitted static prep 152.7 ms (vs 158.9 at E1: D2's hoist now pays, ~6 ms). D4 fused: 144.1 ms — only 6 % better, latents 12.3 dB from XLA (fp32 vs bf16 scores through 8 chaotic steps; the gate decides, not this 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:
- fold
1/√64 = 2⁻³into q before the MXU (a power of two — exact in bf16); - normalise after
P·V, so the division runs on 128 output lanes instead of 1422 score lanes; - (session 6)
expin 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
- Trims (scale folded into q, normalise after P·V), same process: decoder 69.4 → 63.3 ms at 512×256, 60.3 ms at 1024×256 (2048 rows: same); DiT fused 144.1 → 133.5 ms (XLA 152.7).
- r5 end to end (decoder 1024×256 + fused DiT): 5 s 30.5, 30 s 63.5,
120 s 197.8, 380 s 903.7 ms. At 120 s
this is below TensorRT's fp8 tier on an RTX PRO 6000 (205.2 ms). But at
5 s and 30 s the fused DiT is slower than XLA (24.5 vs 21.4,
40.0 vs 34.2 ms): at 184 / 452 tokens there is too little work per
program to amortise the kernel's fixed cost. Since the sequence length
is a static shape, the DiT now picks per length:
DiTConfig.fused_min_tokens(1024 until the crossover is measured). - Null result:
expin bf16 did nothing (decoder 61.3 vs 60.3 ms, DiT 135.2 vs 133.4). Most likely Mosaic evaluates the transcendental in fp32 regardless and the casts only add work. The knob was removed. Worth writing down: "v6e has bf16 VPU" is true, but it did not make this softmax faster.
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 %.