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.
mediumhas 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 ownsa3_dit/sa3_codec/t5_trunkso 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:
- 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=1GEMMs: 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] - 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=2560GEMMs 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×). - 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.scanis 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 |
- 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.
- 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. - 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'smlpand SAME-L'sff. Those modules sit at 674-2471 AI, so quantization error is the only real question. v6e's fp8 is not an alternative (section 2). - 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.
- 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. - 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). - 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.