lab

Roofline and communication cost of klein-4B and Stable Audio 3 medium on TPU v6e

Two inference workloads, one question per module: how many FLOPs, how many bytes, is it compute- or memory-bound, and how much of the wall time is a collective. klein-4B is FLUX.2's text-to-image DiT at 512² (1024 image tokens) and 1360×768 (4080), four Euler steps, with a 27-layer Qwen3-4B text encoder and a conv VAE decoder. Stable Audio 3 medium is a 24-block differential-attention DiT over 256 latent channels, eight steps, with a SAME-L codec and a 12-layer T5Gemma encoder.

Every number is tagged. [M] measured — the source is named in the row. [D] derived — the formula is in tools/roofline.py, whose tables below come from python3 tools/roofline.py all --chip v6e and which is the only thing needed to re-derive them. [S] sourced from a document, with a link. The klein measurement harness is tools/modal_roofline.py; its raw output is out/roofline/klein_l4_512x512.json.

1. What is not measured

The brief assumed hummingbird had already measured Stable Audio 3 on Kaggle TPU. It has not. out/kaggle-tpu-run/ is empty, both out/kaggle-local*/tpu-report.json carry "timing": {}, and the port says so itself: "No TPU number exists yet. Every timing in the docs is CPU, where bf16 is ~2× slower than fp32 because XLA emulates it — those numbers say nothing about TPU" (hummingbird/docs/port-status.md:161-162). medium has never been loaded on any hardware — its 9.8 GB checkpoint is TPU-only for that project (hummingbird/docs/port-status.md:153-155).

So there is no SA3 TPU measurement to compare against and every SA3 wall time here is [D]. What exists is a real measurement of SA3 small-music — the sibling preset — in PyTorch on an RTX 3050 6 GB. Section 7 uses it: hardware-wrong, model-right, enough to test the model's shape, not its scale.

Update, 2026-09-21 — the premise of this section has changed. medium has now run: on a Colab TPU v6e-1, jax 0.7.2, bf16, 8 steps, 120 s of audio. The measurement closes the loop this report could only model, and both of its SA3 predictions survive. The 24-block DiT runs at 217 TFLOP/s — 57 % of the chip's measured bf16 GEMM rate, 1.75× its modelled floor — and the SAME-L codec at 79 TFLOP/s — 21 % of that rate, 4.83× its floor, so §6's memory-bound classification and §7's "×2.5 for GEMM-dominated, ×3–6 for convolutional or small-shape" calibration rule are both confirmed rather than revised. The whole clip lands at 127 TFLOP/s (33 % of the measured GEMM rate, 2.99× the floor).

Consolidated table and method: hummingbird/scripts/roofline_medium.py, which imports this report's own sa3_dit / sa3_codec / t5_trunk so the two cannot drift apart. Full report: stable-audio-3-medium-tpu-vs-l4.html. Every other SA3 figure in this document is still derived; the measured numbers live there.

2. Hardware

ICI is quoted per direction, which is half of what the TPU docs print: the v6e page's "800 GBps bidirectional" sums both directions, so a ring collective that sends and receives at once gets 400 GB/s. [S]

spec v6e source
bf16 TFLOP/s per chip 918 v6e
int8 TOPS per chip 1836 v6e
fp8 TFLOP/s per chip 918, i.e. the bf16 rate TPU7x table; native-vs-emulated unconfirmed
HBM GiB per chip 32 v6e
HBM GB/s per chip 1638 JAX TPU hardware reference
ICI GB/s per chip per direction 400 derived from the bidirectional quote; not stated directly
VMEM / SMEM MiB per TensorCore 128 / 1 JAX TPU hardware reference
MXU per TensorCore 256×256 ×2 v6e
TensorCores per chip 1 v6e
host DRAM GiB, chips per host 1536, 8 v6e
USD per chip-hour, on demand 2.70 TPU pricing
ridge point, bf16 FLOP/byte 560 [D] 918e12/1638e9
TFLOP/s at 30 % MFU 275.4 [D] 0.30 × peak

One correction to flicker/docs/project.md §8.3, whose hardware paragraph predates this port:

v6e's fp8 number is not a speedup. Google publishes 918 fp8 TFLOP/s for v6e — exactly the bf16 figure, and the JAX runtime reports fp8_ops_per_second = 920. That is the signature of bf16-rate emulation; vLLM's TPU stack lists fp8 hardware acceleration for v7 only. This analysis models v6e as bf16-only for dense compute. The brief's "v6e does [have fp8]" is half right: the number exists, the throughput does not.

The load-bearing v6e fact is its 560 FLOP/byte ridge point — 918 TFLOP/s against 1638 GB/s. Anything below that arithmetic intensity is memory-bound, and seven rows across the two models are, bolded in §4.1 and §6: their fused intensity looks like a GEMM's, but their unfused intensity lands under the ridge. v6e therefore needs fusion, and it needs it where the traffic is, not where the FLOPs are.

3. Method

A multiply-accumulate is 2 FLOPs. A GEMM (m,k) @ (k,n) is 2mkn FLOPs, kn weight elements and m·k + m·n activation elements. Attention scores cost 4·head_dim + 5 FLOPs each — QK^T 2D, softmax 5, PV 2D — over heads × queries × keys of them. Elementwise: rmsnorm 4 FLOP/element, layernorm 6, DyT 3, rope 4 per rotated element, silu_gate 5 per output, modulate 2, residual 1. All bf16 (2 B). [D]

Classification uses the unfused column, so a module that is compute-bound only if XLA fuses everything is called memory-bound. Each row reports two byte counts, and the gap between them is the argument about fusion. act MB is the minimum a fused kernel needs: weights read once, the module's input and output touched once. unfused MB adds every interior tensor written and read again. Real XLA fuses elementwise chains into their producer, so the truth is between the columns; unfused is the pessimistic bound every time prediction uses, so the model is not flattered. It has no launch overhead and no tail effect; section 7 prices that omission.

4. klein-4B

4.1 The DiT is a GEMM problem, and the parameter-count shortcut overcounts it by 26 %

module, one forward, 512² GFLOP w MB act MB unfused MB AI AI unf v6e
embeddings (in_proj + time_in) 30.20 68.4 9.4 90.8 388 332 memory
modulation ×3 (shared) 0.28 283.1 0.1 283.2 1.0 1.0 memory
final layer 0.87 38.5 6.3 70.3 19 12 memory
×5 double: norm+modulate 0.19 0.0 94.4 188.7 2.0 1.0 memory
×5 double: qkv 434.87 566.2 141.6 755.0 614 576 compute
×5 double: qk-norm+rope 0.38 0.0 94.4 377.5 4.0 1.0 memory
×5 double: attention 146.37 0.0 47.2 188.7 3102 776 compute
×5 double: proj 144.96 188.7 47.2 283.1 614 512 memory
×5 double: mlp 1304.95 1698.7 47.2 2500.9 747 522 memory
×20 single: norm+modulate 0.75 0.0 377.5 755.0 2.0 1.0 memory
×20 single: linear1 5219.14 3397.4 1887.4 6039.8 988 864 compute
×20 single: qk-norm+rope 1.51 0.0 377.5 1509.9 4.0 1.0 memory
×20 single: attention 585.48 0.0 188.7 755.0 3102 776 compute
×20 single: linear2 2320.70 1509.9 943.7 3586.1 946 647 compute

[D] 512², L=1024 image tokens, 512 text tokens, joint sequence 1536, hidden 3072, 24 heads × 128, mlp_hidden 9216 with a gated 3072→18432 up-projection. Shapes verified against the checkpoint header: single_blocks.0.linear1.weight is (27648, 3072), linear2 is (3072, 12288), img_mlp.0 is (18432, 3072), and the DiT totals 3,875,544,576 parameters — matching the repo's 3.8755 B, with 245.37 M per double block and 122.68 M per single block. [M]

The DiT is compute-bound even unfused. One forward is 10.191 TFLOP against 17.384 GB of unfused traffic: AI 848 fused, 586 unfused, both above v6e's 560 ridge.

single: linear1 and linear2 are 74 % of DiT FLOPs. Attention is 7.2 % at 512² (0.732 of 10.191 TFLOP) and 18.8 % at 1360×768 (6.541 of 34.761). That inverts the usual diffusion intuition, and makes splash attention a VMEM play here, not a FLOP play.

The repo's parameter-count shortcut overcounts klein by 26 %. docs/benchmarks.md:62 computes matmul FLOPs as 2·3.8755e9·seq, assuming every parameter multiplies every token. It does not: img_in sees 128 channels, txt_in sees 512 tokens not 1536, final_layer emits 128, the modulation matrices see one token, and the double block's down-projection consumes 9216 where its up-projection produced 18432. The correct matmul figure at 512² is 9.459 TFLOP/step against the shortcut's 11.906. [D] The attention term agrees to 1 % (25 · 4 · seq² · 3072 = 0.725 TFLOP; the 1 % is the softmax).

block TFLOP weights GB unfused GB
text encoder, 27 layers, S=512 2.909 TFLOP 5.450 GB 8.897 GB
VAE decoder, 512² 1.994 TFLOP 0.099 GB 2.366 GB
VAE decoder, 1360×768 8.355 TFLOP 0.099 GB 9.132 GB
DiT, 1360×768, 1 forward 34.761 TFLOP 7.751 GB 36.556 GB
whole pipeline, 512², 4 steps 45.666 TFLOP 36.553 GB 80.799 GB
whole pipeline, 1360×768, 4 steps 150.309 TFLOP 36.553 GB 164.253 GB
distinct klein weights (DiT + text + VAE) 13.30 GB

[D] Qwen3 weights verified against the checkpoint: 27 layers × 100,930,816 = 2.725 B plus a 389 M embedding, so the executed trunk is 3.114 B and 5.450 GB of bf16. [M] The 4.02 B in the file includes nine layers and a model.norm that never run. The text encoder is 6.4 % of 512² FLOPs but sits serially ahead of the DiT, and every one of its five rows is memory-bound on v6e (unfused AI 317-465 against a ridge of 560) — traffic, not FLOPs, is what it costs. Its causal attention is counted at the full S² the port actually computes.

The VAE decoder path is 49.6 M parameters, not the ~84 M in the brief — 84.05 M is the whole file including the 34.4 M encoder this port never loads. [M] But it scales badly: at 1360×768 it is 8.355 TFLOP, 5.6 % of the pipeline, because up2 and up3 are a quarter of the FLOPs each and the resolution keeps doubling. At 512² it is 1.994 TFLOP, so the repo's "~1-2 TFLOP" guess holds there and is 4× low at 768p.

4.2 Divisibility: the 4080-token claim is wrong, but the padding is real

The brief states klein's 1360×768 image sequence "is 4080 tokens, which is not divisible by 8". It is: 4080 = 8 × 510 = 16 × 255. The joint [txt | img] sequence is 4592 = 8 × 574. Ulysses over 8 chips is exact for both streams. [D]

divisor image mod joint mod image pads to extra image tokens extra attention FLOPs
8 0 0 4080 +0.000 % +0.000 %
16 0 0 4080 +0.000 % +0.000 %
32 16 16 4096 +0.392 % +0.698 %
128 112 112 4096 +0.392 % +0.698 %

What is misaligned is splash attention's block_kv % 128 == 0. 4080 = 2⁴·3·5·17, so it divides by 2, 4, 8 and 16 and by nothing above. [D] The 128-token pad costs 0.392 % more tokens and 0.698 % more attention FLOPs — 0.13 % of DiT FLOPs, a rounding error next to FSDP. The padding matters for a splash kernel's correctness, not for performance.

5. Collectives: FSDP is the trap, Ulysses is nearly free

Per-chip wire bytes: ring all-reduce 2(N-1)/N · S, all-gather (N-1)/N · S, all_to_all (N-1)/N² · S. Time is wire bytes over 400 GB/s (v6e, one direction). [D]

collective, klein 512², per image (4 steps) tensor MB count wire MB/chip @1 @4 @8 ms @1 @4 @8
TP double block (img x3) 6.29 15 0.00 141.56 165.15 0.000 0.354 0.413
TP double block (txt x3) 3.15 15 0.00 70.78 82.58 0.000 0.177 0.206
TP single block (x1) 9.44 20 0.00 283.12 330.30 0.000 0.708 0.826
FSDP all-gather (all 3.8755 B weights) 7751.09 1 0.00 5813.32 6782.20 0.000 14.533 16.956
Ulysses qkv+out all_to_all 37.75 25 0.00 176.95 103.22 0.000 0.442 0.258
collective, SA3 medium 120 s, per clip tensor MB count wire MB/chip @1 @4 @8 ms @1 @4 @8
FSDP all-gather (1.4532 B DiT weights) 2906.34 1 0.00 2179.76 2543.05 0.000 5.449 6.358
Ulysses qkv all_to_all 21.84 48 0.00 196.58 114.67 0.000 0.491 0.287
Ulysses attn-out all_to_all 4.37 48 0.00 39.32 22.93 0.000 0.098 0.057

FSDP-8 costs 3.6× klein's floor time. The 3.8755 B weights are all-gathered once per forward: 16.96 ms of pure collective against a 6.22 ms math floor at 8 chips. The pipeline's floor goes from 6.48 ms (Ulysses, 96.0 % of peak) to 23.17 ms (FSDP, 26.8 %). [D] The repo's docs/benchmarks.md:124 reaches the same conclusion: its 17 ms is this all-gather at the v6e per-direction rate.

Ulysses is 5.6× cheaper than tensor parallelism and 65× cheaper than FSDP-3 — 0.26 ms at 512² against 1.45 and 16.96 (at 1360×768 Ulysses is 0.77 ms, TP 4.32 ms, FSDP the same 16.96). [D] The mechanism: Ulysses moves 4·S·H bytes per attention layer through an all_to_all whose per-chip share falls as (N-1)/N², while TP moves S·H per GEMM through an all-reduce whose share falls only as 2(N-1)/N — and FSDP moves weights.

Depth sharding is the wrong axis at batch 1. Splitting 25 klein blocks over 8 chips leaves ~3 blocks per chip and a 9.44 MB boundary activation per hop — cheap on bandwidth, but only one chip has work at a time, so the schedule pays a 7/8 bubble. It is a throughput mechanism only if images are packed, never a latency one. Same for SA3's 24 blocks. [D]

6. Stable Audio 3 medium

At 120 s: latents 1358, DiT tokens 1422 (64 memory tokens + 1358), SAME-L sequence 23,086 (17 per latent frame), 8 steps, cross-attention context 257. [D] from sa3jax/pipeline.py:52-62.

module, 120 s clip GFLOP w MB act MB unfused MB AI AI unf v6e
×24 DiT: adaLN + pre-norms 0.84 0.0 209.7 838.7 4.0 1.0 memory
×24 DiT: self to_qkv 805.18 566.2 629.0 1195.3 674 674 compute
×24 DiT: self qk-norm+rope 1.05 0.0 419.4 1258.1 2.5 0.8 memory
×24 DiT: self attention (differential ×2) 607.98 0.0 104.8 419.4 5799 1450 compute
×24 DiT: cross to_q/to_kv 496.70 906.0 123.8 1353.1 482 367 memory
×24 DiT: cross attention (differential ×2) 109.88 0.0 104.8 247.6 1048 444 memory
×24 DiT: cross to_out + residual 161.04 113.2 209.7 322.9 499 499 memory
×24 DiT: local cond add 0.05 0.0 209.7 209.7 0.2 0.2 memory
×24 DiT: mlp (SwiGLU) 1933.48 1359.0 629.0 3665.5 973 527 memory
×12 codec: pre_norm + to_qkv 6537.28 283.1 5106.3 7091.5 1213 922 compute
×12 codec: qk DyT + rope ×4 8.51 0.0 3404.2 10212.5 2.5 0.8 memory
×12 codec: sliding-window attention 177.00 0.0 851.0 1702.1 208 104 memory
×12 codec: ff_norm + GLU + proj_out 11772.47 509.6 4255.2 16679.4 2471 706 compute
×12 T5Gemma: qkvo + rope 14.51 56.6 9.4 113.2 220 128 memory
×12 T5Gemma: mlp (GeGLU) 29.02 113.2 22.0 190.3 215 152 memory

[D] MEDIUM: embed_dim 1536, depth 24, 24 heads × 64, ff_inner 6144, differential=True (so to_qkv emits 5·dim and each attention runs two SDPAs sharing v), 64 memory tokens. The cross-attention to_kv recomputes from the 257 context tokens in every block of every step — 906 MB of weights, 24 times per step, 8 steps per clip, for a tensor that never changes.

The model reproduces hummingbird's measured FLOP split on different hardware. Here the codec is 36.0 % of clip FLOPs at 120 s and the DiT 64.0 %. The repo measured, in PyTorch on an RTX 3050 at 120 s, DiT 65.3 % / decoder 33.9 % / post-process 2.0 % / conditioning 0.5 % (hummingbird/docs/project-readme-hummingbird.md:113-114). [M] at the source. One analytical derivation and one profiler agreeing to 2 points on different hardware is the strongest evidence here that the FLOP model is shaped right.

The SAME-L codec is 99.0 % GEMM (18,345 of 18,531 GFLOP) and its sliding-window attention — the mechanism the whole SAME-L design exists for — is 0.96 %. But its traffic exceeds its FLOPs share: 35.9 GB of the clip's 112.8 GB, with qk DyT + rope ×4 alone at 10.2 GB and an arithmetic intensity of 0.8.

Four SA3 rows are memory-bound on v6e: cross to_q/to_kv (367 unfused), cross attention (444), cross to_out (499) and mlp (527) all sit under the 560 ridge. int8 helps them twice — halved weights, halved traffic — where a better kernel would not. The same applies to klein's embeddings (332), double: proj (512) and double: mlp (522).

The DiT has a VMEM problem at long clips. sa3jax/dit.py:116 materialises the score matrix, so at 4160 tokens the fp32 logits are 24 × 4160² × 4 B = 1.66 GB per differential branch, 3.3 GB for one attention, against 128 MiB of VMEM per TensorCore. At 120 s it is 194 MB per branch. [D] A Pallas flash kernel is not an optimisation here; it is the difference between fitting and not.

6.1 Clip length

clip latents DiT tokens codec tokens DiT 8 steps codec codec share text

[D] medium clamps at 4096 latents = 380.4 s, not at 120 s — the 1292-latent clamp in sa3jax/pipeline.py:56-58 is small-music's, whose sample_size is 5,292,032. Past 30 s the split is flat, because both stacks are linear in latent count once attention is a small share.

7. Measured against modelled

One L4 run, 512², 4 steps, warm-up and compile excluded ([M], tools/modal_roofline.py). The L4's own calibration in the same process: 80.6 TFLOP/s on a 4096³ bf16 GEMM — 67 % of its 121 TFLOP/s datasheet peak — and 226.2 GB/s streaming read, 75 % of 300 GB/s. [M] Two model columns: the pure roofline at datasheet peak, and the same formula at the machine's measured rates. The gap between the two ratios is the hardware's shortfall; what is left is the model's.

target measured ms floor @ peak ratio floor @ measured rate ratio achieved
klein text encode, S=512 318.9 29.7 0.093× 39.3 0.123× 11 % of measured GEMM rate
klein DiT, 1 forward, 512² 350.2 84.2 0.240× 126.5 0.361× 36 %
klein DiT, 4-step scanned denoise 1325.2 336.9 0.254× 506.0 0.382× 38 %
klein VAE decode, 512² 134.2 16.5 0.123× 24.8 0.184× 18 %
klein 1 double block (differenced) 13.5 3.4 0.248× 5.0 0.373× 37 %
klein 1 single block (differenced) 12.0 3.4 0.279× 5.0 0.419× 42 %
klein DiT, 1 forward, 512², v6e-1 58.3 11.1 0.190× 26.8 0.460× 46 %
klein DiT, 4-step scanned denoise, v6e-1 158.1 44.4 0.281× 107.2 0.678× 68 %
klein VAE decode, 512², v6e-1 14.4 2.2 0.151× 5.2 0.364× 36 %
klein text encode, S=512, v6e-1 104.8 5.4 0.052× 7.7 0.073× 7 %
sa3 small-music DiT, 1 forward, T=1292 173.7 78.4 0.452× — — RTX 3050, measured rates
sa3 small-music T5Gemma, S=256 11.0 3.0 0.277× — — RTX 3050, measured rates

[M] Sources: this report's L4 run for the L4 klein rows; hummingbird/out/shots/kernel-stats.md:47-54 and hummingbird/docs/project-readme-hummingbird.md:104 for SA3. SA3 uses small-music because that is the preset hummingbird profiled — medium has never run on hardware. The v6e-1 rows are flicker's single-chip TPU measurement (§4/§5 of the TPU report): @ peak is max(math, memory) at 918 TFLOP/s and 1638 GB/s, and @ measured rate raises only the math term at the calibrated 380.1 TFLOP/s — no streaming rate was calibrated on TPU, so the memory term keeps the datasheet bandwidth.

The model underestimates by 2.4-2.7× on the DiT and by 8.1× on the text encoder. Three mechanisms, only the first of which the brief anticipated:

  1. Prologue and non-GEMM work (2.4-2.7×). A per-block difference isolates it: one single block measures 12.0 ms against a 5.0 ms floor, and the blocks sum to 308 ms of the 350 ms forward. The other 39 ms is the prologue — embeddings, three shared modulations, RoPE tables, final layer — whose modelled cost is 2.0 ms. That prologue is m=1 GEMMs: 283 MB of modulation weights streamed to produce 92 KB of output, at a rate far below the 226 GB/s a large streaming read achieves. [D]
  2. Small-shape inefficiency (8.1×). 2.909 TFLOP in 318.9 ms is 9.1 TFLOP/s, 11 % of what the same GPU does on a 4096³ GEMM: twenty-seven layers of m=512, k=2560 GEMMs never reach the rated rate, so a model with one uniform rate cannot be right there. Next worst by the same mechanism: the VAE (5.4×, 14.9 TFLOP/s) and SA3's T5Gemma (3.6×).
  3. Nothing, in one place. The scanned 4-step denoise is 1325.2 ms against 4 × 350.2 = 1400.8 ms of separately dispatched forwards, so lax.scan is 1.057× faster than four jitted calls and dispatch overhead is not the gap. [M]

The same rows are now measured on a TPU — the chip this report is priced for. Against the measured-rate floors: the 4-step DiT is 1.47× (the L4 row above: 2.6×), the VAE 2.74× (L4: 5.4×), the text encoder 13.7× (L4: 8.1×). [M] The GEMM-bound stages get tighter on a second vendor and a decade of newer hardware — the model transfers — while the text encoder gets worse by exactly mechanism 2: small-shape GEMMs on an MXU, fed by a host instead of a launch queue. Against the datasheet floors the same three are 3.6×, 6.5× and 19×, and docs/tpu-article.md §5.1 prices that floor module by module.

Two operational numbers fall out of the same run that no FLOP model predicts. Weight loading is 37× the compute: 67.1 s to stage 13.3 GB sequentially against 1.778 s of maths for the whole pipeline ([M], 171-255 MB/s). Read that rate for what it is — a host read, not a device transfer: the same 13.3 GB over PCIe at 10-25 GB/s is 0.5-1.3 s, and into HBM at 1638 GB/s it is 8 ms. The 37× is a property of reading checkpoints off network storage on this host, so it is the number most likely to change on a TPU VM with the checkpoint in page cache, and the one to re-measure before quoting. And compile time is not proportional to FLOPs: the VAE's 2.0 TFLOP costs 12.2 s of XLA compile, the DiT's 40.8 TFLOP costs 4.0 s. [M] Peak memory, sequentially loaded: DiT 9.31 GiB, text encoder 8.25 GiB — 17.56 GiB of the two together against the L4's 16.5 GiB limit, which is why the pipeline stages them. [M]

Read sections 4-6's times as floors and multiply by 2.5 for a GEMM-dominated module and 3-6 for a convolutional or small-shape one. The rankings survive, because the correction is roughly constant within a module class.

8. Where the time goes, and what to do

Predicted floor at 8× v6e, one 512² image, four steps, largest first. [D]

hotspot floor ms % peak binding fix
×20 single: linear1 2.84 100 math int8: 988 AI fused, a pure weight-streaming GEMM
×20 single: linear2 1.26 100 math nothing: already fused with the attention output, 946 AI
×5 double: mlp 0.76 93 memory int8: 1969 AI at 1360x768 is a pure GEMM problem
×20 single: qk-norm+rope 0.46 0 memory fusion: fold the rope and qk-norm into the qkv GEMM epilogue
×27 text: mlp (SwiGLU) 0.44 64 memory int8, or FSDP-8 (runs once per prompt)
×20 single: attention 0.32 100 math splash/Pallas for VMEM only; FLOPs are unchanged
×5 double: qkv 0.24 100 math nothing: already one GEMM

Predicted floor at 8× v6e, one 120 s SA3 clip, eight steps, largest first. [D]

hotspot floor ms % peak binding fix
×24 DiT: mlp (SwiGLU) 2.24 94 memory int8; 527 unfused is memory-bound on v6e
×12 codec: ff_norm + GLU + proj_out 1.60 100 math int8: 2471 AI fused, 706 unfused, at 120 s
×12 codec: pre_norm + to_qkv 0.89 100 math int8
×24 DiT: self to_qkv 0.88 100 math int8: 674 AI fused
×24 DiT: cross to_q/to_kv 0.83 65 memory hoist K/V once per clip, not per block per step
×12 codec: qk DyT + rope ×4 0.78 0 memory fusion: 10.2 GB of traffic for 8.5 GFLOP
×24 DiT: self attention (differential ×2) 0.66 100 math Pallas/flash: the port materialises fp32 logits
  1. Do not use FSDP for either model's DiT. For klein it is 16.96 ms of collective against a 6.22 ms math floor at 8 chips, and it drags the 512² pipeline from 6.48 ms (Ulysses, 96.0 % of peak) to 23.17 ms (26.8 %). For SA3's 120 s clip the same all-gather is 6.36 ms, lifting the 8-chip floor from 8.95 ms (Ulysses, 78.4 %) to 14.96 ms (46.9 %). Ulysses costs 0.26 and 0.34 ms at the same chip count. Shard weights between requests, not within one.
  2. Do not shard the text encoder either. Qwen3's 27 executed layers have a 5.4 ms floor on one v6e chip — memory-bound, 3.2 ms of math against 8.9 GB of traffic — and fit in 5.45 GB of 32 GiB; FSDP-8 would add 11.9 ms of all-gather to save at most 4.8 ms. [D] T5Gemma is 0.046 TFLOP for the whole clip — a warm-up, not a workload.
  3. int8 on the four big GEMM stacks is the only change that moves the needle. It halves both the FLOPs and the weight traffic of single: linear1/linear2, double: mlp, SA3's mlp and SAME-L's ff. Those modules sit at 674-2471 AI, so quantization error is the only real question. v6e's fp8 is not an alternative (section 2).
  4. Fuse the norm, rope and modulate rows. They are 16 % of klein's 8-chip floor at 0.2 % of peak — a bigger share on v6e than a lower-ridge part would pay, because their traffic is fixed while the compute peak is not. Still cheap to do rather than valuable in itself.
  5. Hoist SA3's cross-attention K/V. The port already flags this (sa3jax/dit.py:245, deferred to M8); it is the one change in either model that is free.
  6. Leave the VAE and the step count alone. The VAE is 4.4 % of 512² FLOPs, and four steps is already distilled — the caching literature's 1.3-2× numbers came from 28-50-step baselines (flicker/docs/benchmarks.md:103-111).
  7. Re-measure before believing any of this. klein now has one — §7's v6e-1 rows, where the DiT and VAE gaps close and the text encoder's widens — but SA3 has no TPU measurement at all, and the correction in section 7 was priced on a GPU whose small-shape overhead does not transfer either way.