lab

One TPU chip, one image: klein-4B on a Colab v6e-1

FLUX.2 klein-4B in JAX, bf16, batch 1, four Euler steps, text encoder and VAE included, on one Colab v6e-1 chip: 32 GiB, 918 TFLOP/s, 1638 GB/s, ridge 560 FLOP/byte [S]/[D]. Resolutions 512² (L=1024), 1360×768 (L=4080) and, for the 32 GiB headroom question, 1024² (L=4096), 512 text tokens; 45.666, 150.309 and 150.920 TFLOP per image [D]. Every latency is a warmed median with jax.block_until_ready() around the region and the sample count in the JSON (10 per stage, 2 per serving shape after the first request); XLA compile, the first post-compile call and the 13.3 GB of weight staging are reported separately, and the chip's spec block travels with the numbers. The harness is tools/bench_tpu.py; the FLOP counts are tools/roofline.py's, imported, not re-derived.

1. What a single chip changes

A one-chip run deletes the communication half of the cost model. jax.device_count() is 1 [M], so every collective in docs/roofline.md §5 costs exactly zero — and is no longer a lever: FSDP-8's all-gather (16.96 ms/image at 8 chips), Ulysses (0.26 ms) and tensor parallelism (1.45 ms) [D] all vanish, and with them the layout frontier. What is left per module is max(math, memory) against the chip's own peak and bandwidth — the roofline's half of the result, §5.1.

chip bf16 TFLOP/s int8 TOPS HBM GiB HBM GB/s ridge FLOP/byte floor ms 512² / 1360×768 / 1024² floor img/s
v6e-1 918 1836 32 1638 560 49.7 / 163.7 / 164.4 20.1 / 6.1 / 6.1

notes: floor = FLOPs / peak, the roofline's max(math, memory) with comms = 0; peaks, capacity, bandwidth and ridge are the datasheet [S], the floors derived [D].

32 GiB holds all three components at once. DiT 7.75 GB + Qwen3-4B trunk 5.45 GB + VAE 99.2 MB = 13.30 GB of weights [D] — 38.7 % of the chip's HBM. Measured: v6e-1 holds them with room — load_all took 35.9 s for 14.08 GB of weights and left 22.09 GB in use against a 33.55 GB device limit — 66 % of capacity, 34 % still free [M].

component weights [D] peak residency [M, L4] v6e-1 (32 GiB)
DiT 7.75 GB 9.31 GiB fits
Qwen3-4B trunk, 27 of 36 layers 5.45 GB 8.25 GiB fits
all three resident 13.30 GB 17.56 GiB 22.09 GB peak, 66 % of the limit

These numbers are a lower bound, and that is a reporting constraint, not a footnote. Colab's managed runtime reports transparent hugepages as always [madvise] never [M] — always is not set, and it is not ours to set; JAX warns about exactly this on every run. Rule 7 follows: Colab for correctness and iteration, a gcloud TPU VM for publishable performance.

2. Correctness on a third backend

The port passes its ladder against PyTorch on a Modal L4: 16 rungs, 81 pass / 29 warn / 0 FAIL, end-to-end 33.3 dB PSNR, mae 2.6 levels [M]. On TPU the ladder is a band comparison, never a byte comparison: bf16 is not portable across hosts, and any byte diff of a bf16 tensor across machines is meaningless (docs/port-plan.md §6.4).

rung group L4 vs PyTorch [M] requirement [S] v6e-1 [M]
1, 12, 16 (token ids, unpatchify, injected noise) bit-exact exact — a convention check has no dtype freedom bit-exact
2-11, 13-15 (per-layer to per-step tensors) 27-77 dB; 159-181 dB on rung 7 cosine > 0.99 and SNR ≥ 20 dB; 20-45 dB is the bf16 band 26.18 dB minimum, 37 warns
end-to-end image 33.3 dB PSNR, mae 2.6 in-band PSNR, never byte-equal 31.0 dB, mae 2.91
totals 81 / 29 / 0 — 73 / 37 / 0

notes: same captured noise, same prompt, same fixture; the only moving parts are the codegen and the chip. Median per-key movement against the L4 is +0.31 dB, with 64 of 98 finite-SNR keys inside ±2 dB [M]. The eight extra warns are one localised shift and not scatter: txt/hidden/7..15 sat just above the 45 dB line on the L4 (45.0-46.7) and just below it on the TPU (41.0-41.9) [M]. The per-layer injected checks agree between hosts to ±1.2 dB, so it is accumulation through 27 Qwen3 layers, not a bad layer — which is the whole reason the ladder injects the reference's inputs.

The finding worth taking away: an fp32 matmul on TPU is a bf16 matmul unless you raise the precision. Measured against float64, shape- and chip-independent [M]:

precision relative error vs fp64 speed
DEFAULT 2.32-2.35e-3 (bf16-grade, 2⁻⁹) 96 % of bf16
HIGHEST 5.4-8.8e-8 (true fp32) 13-15 % of bf16

The tell was the throughput — fp32 at 96 % of bf16 is impossible for 3-pass emulation — and it only surfaced by scoring the product against float64 rather than trusting the declared dtype. It landed in the block tests as a clean split: every fp32 row containing a matmul collapsed (x_in 143.7 → 52.4 dB) while elementwise fp32 rows were untouched, which is why 4 of 10 tests failed at DEFAULT and jax_default_matmul_precision="highest" turns them into 10/10 with no bf16 row getting worse [M]. Two reproduction notes: constant operands get host-folded and hide the downgrade, so both must be jit arguments; and a degenerate K=2 dot stays fp32, so one small probe misses it.

This reaches into the port itself. nn.sdpa upcasts q/k/v to fp32 and computes both einsums in fp32 specifically to get true fp32 accumulation. On TPU at DEFAULT, those einsums run in bf16. It did not matter — the ladder is 0 FAIL either way and dit/single/0/attn scores better on the TPU (47.2 dB) than on the L4 (42.7), because PyTorch's own SDPA is mixed-precision — but "we wrote fp32 logits" was not a true statement about the hardware.

3. Resident, staged, cold: three serving shapes

Residency is a product question, and it is the largest single effect measured here.

shape what one request pays included v6e-1 [M]
cold first generate in the process staging + compile + compute 61.2 s
staged generate again, reloading 13.3 GB staging + compute 35 735 ms
resident load_all once, generate_resident after compute 262 ms

notes: 512²; the other resolutions are in §4. Cold is never scored as an MFU: a number containing XLA compile is not a latency number.

The staged-minus-resident delta is 35.47 s per request at 512², 35.97 s at 1360×768 and 35.75 s at 1024² [M] — one image of maths against a 13.3 GB reload, a factor of 136×. So the same kernels serve at 3.8 images/s or at 0.028 images/s depending only on whether the chip can hold the weights.

Two honest qualifications. First, 35.9 s for 14.08 GB is 392 MB/s, a host-I/O rate rather than a device rate — the same bytes over PCIe at 10-25 GB/s is 0.5-1.3 s [D] and into HBM at 1638 GB/s it is 8 ms. The staged penalty is real on this VM and is a property of the storage path as much as of the chip; on a VM with the checkpoint in page cache the two shapes would sit far closer. Second, this is a batch-1 measurement: at 512² a resident v6e-1 spends 262 ms per image and could pipeline, so the honest framing is that residency converts a per-request I/O cost into a one-time one.

4. The measured latency budget

stage TFLOP [D] floor ms [D] v6e-1 ms [M] %MFU [M] compile s [M] peak HBM [M]
GEMM calibration, 4096³ — — 0.36 41.4 1.40 —
text encode, S=512 2.909 3.2 104.8 3.02 2.72 7.58 GB
DiT 1 step, 512² 10.191 11.1 58.3 19.05 4.25 14.95 GB
DiT 4-step scan, 512² 40.763 44.4 158.1 28.08 4.45 14.95 GB
VAE decode, 512² 1.994 2.2 14.4 15.06 11.24 14.96 GB
DiT 4-step scan, 1360×768 139.045 151.5 779.8 19.42 4.69 14.95 GB
VAE decode, 1360×768 8.355 9.1 42.9 21.23 15.54 14.96 GB
DiT 4-step scan, 1024² 139.621 164.4 808.9 18.80 5.83 14.96 GB
resident request, 512² 45.666 49.7 262 19.0 — 22.09 GB

notes: measured ms excludes compile, the first post-compile call and weight staging; every spread is under 8 % over 10 samples and the JSON carries each one. The honest denominator is the calibration row, not the datasheet: on the DiT scan v6e-1 reaches 28.1 % of datasheet but 67.8 % of its own 4096³ rate [M] — the datasheet peak is simply not reachable, even by a plain GEMM (41.4 %). Compile time is not proportional to FLOPs, the same inversion the L4 showed: the 2.0 TFLOP VAE costs 10-16 s of XLA and the 40.8 TFLOP DiT costs 4.1-4.7 s [M]. The floor column is the math term, FLOPs ÷ datasheet peak; §5.1 prices the memory term too, and for one row it is the larger one.

The text encoder is the surprise. It is 6.4 % of a 512² image's FLOPs and it costs 104.8 ms of a 262 ms image on v6e-1 — 40 % of the wall clock — at 3.02 % of datasheet MFU, the worst module in the pipeline by a factor of five [D/M].

5. Predicted vs measured: the text-encoder prediction was wrong

The roofline's analytical model underestimates the L4 by 2.4-2.7× on the DiT and 8.1× on the text encoder, decomposed there into an m=1 prologue and small-shape inefficiency [M/D]. The prediction was that the second mechanism shrinks on TPU, because a TPU has no per-kernel host launch and XLA feeds the MXU itself. Scored against each chip's own calibrated GEMM rate:

stage calibrated floor ms [D] predicted factor [D] measured factor [M] verdict
DiT 4-step scan, 512² 107.2 1.2-1.6× 1.47× holds
DiT 4-step scan, 1360×768 365.7 1.2-1.6× 2.13× above the band
VAE decode, 512² 5.2 1.5-3.0× 2.74× holds
text encode, S=512 7.7 1.5-2.5× 13.7× falsified
end-to-end, 512² 120 1.3-1.8× 2.18× dragged up by the text row

notes: the calibrated floor is FLOPs ÷ the chip's measured 4096³ GEMM rate (380.1 TFLOP/s) [M], which is the only denominator that separates the model's error from the hardware's; the end-to-end factor is the measured resident request. The DiT and VAE predictions transferred to a different vendor and a different decade of hardware; the text encoder's did not, and it is worse on the TPU than the 8.1× the same model scored on the L4 [D].

So the mechanism the prediction rested on is not the mechanism that governs the text encoder. It is not established what does: the candidate explanations — a per-layer mix dominated by norms, rope and softmax rather than GEMMs; the fp32↔︎bf16 conversions §2 describes; the causal mask built inside the jit on every call — are hypotheses, and the falsifying measurement says the next instrument is per-layer differencing, not more modelling.

For reference against the published bars (docs/benchmarks.md §A.2, which are 8-chip numbers): one v6e-1 chip at 0.262 s for 512² is 2.4× the 8-chip target of 0.107 s at 31 % MFU — i.e. eight of these chips at the same per-chip efficiency would beat the target by 3.3×. That is arithmetic, not a measurement: §9 explains why it cannot be assumed.

5.1 The floor, module by module — the roofline's half of the same result

The tables above say where the gap is. docs/roofline.md says what the floor is, row by row: same chip, same graph, datasheet peak (918 TFLOP/s, 1638 GB/s), floor = max(math, memory) per row [D], printed by python3 tools/roofline.py all --chip v6e — the one-chip version of its §8 hotspots.

module, 512², 4 steps floor ms [D] % of peak [D] binding [D] the lever
×20 single: linear1 22.74 100 math int8 — 988 AI, pure weight streaming
×20 single: linear2 10.11 100 math already fused with the attention output, 946 AI
×5 double: mlp 6.11 93 memory int8 (1969 AI once the bytes halve)
×20 single: qk-norm+rope 3.69 0.2 memory fold rope + qk-norm into the qkv epilogue
×27 text: mlp (SwiGLU) 3.50 64 memory int8, or run once per prompt
×20 single: attention 2.55 100 math splash/Pallas for VMEM only, FLOPs unchanged
×5 double: qkv 1.89 100 math already one GEMM
×20 single: norm+modulate 1.84 0.2 memory fold into the preceding residual add
×5 double: qk-norm+rope 0.92 0.2 memory fold into the qkv epilogue
×27 text: norm+qkv 0.84 56 memory nothing above the floor
vae: up2 256×256 + upsample 0.69 100 math nothing
modulation ×3 (shared) 0.69 0.2 memory 283 MB of weights moved to produce 92 KB

notes: rows are per-row maxima, so they sum above the pipeline's 49.7 ms — that figure is a single max over the whole graph (49.74 math against 49.33 memory, within 1 % of each other) [D]. Three readings fall out of it:

6. Inside the path: where the time actually goes

jax.profiler traces of each stage, captured on v6e-1 with everything resident, one region per trace (the whole capture is 7.4 MiB, .xplane.pb included). The figures below are Perfetto renders, cropped to their window and rebased to zero because Perfetto's deep-link zoom does not take effect in headless Chromium.

These are XProf traces, not the JSON JAX writes for Perfetto. jax.profiler.trace emits both a perfetto_trace.json.gz and the .xplane.pb files, and the xplane carries materially more — the DIT scan has 47 819 events against the native trace's 20 191, 709 distinct op names against 649, and, decisively, two tracks the native file does not have: a Source code track naming the Python line that issued each op, and a Framework Name Scope track separating jit(denoise) / while / body / closed_call. It also labels /host:CPU and /device:TPU:0 as separate processes, which is what turns the host-versus-device question below from an inference into a measurement. Converting is xprof.convert.raw_to_tool_data.xspace_to_tool_data([xplane], "trace_viewer@", {}) — tools/xplane_to_perfetto.py wraps it.

Two numbers per stage matter, and they are not the same number. Latency is the warmed median from tools/bench_tpu.py. Device time is the span of top-level ops on the trace's XLA Ops track — what the TPU was actually executing. The gap between them is time the MXU was idle.

stage latency ms [M] device-op ms [M] floor ms [D] device / latency latency / floor %MFU datasheet [M]
text encode, S=512 104.8 16.1 7.7 15 % 13.7× 3.0
DiT 4-step scan, 512² 158.1 130.1 107.2 82 % 1.47× 28.1
VAE decode, 512² 14.4 9.7 5.2 67 % 2.74× 15.1
whole request 262 — 120.1 — 2.18× 19.0

notes: the floor is FLOPs ÷ the chip's own measured 4096³ GEMM rate (380.1 TFLOP/s), not the datasheet [D]; device-op spans are from the XLA Ops track of each region's trace, divided by the traced iteration count [M]; %MFU datasheet is FLOPs ÷ 918 TFLOP/s ÷ the warmed latency [M]. Latency excludes compile and weight staging. The whole-request device figure is omitted because the three stages were traced separately and their host overheads do not sum cleanly.

The DiT is textbook proportioned. Per step, the 20-block single stack takes 25.57 ms and the 5-block double stack 6.71 ms — 79 % / 21 % of the scan, against a FLOP split of 80 % / 20 % [8125 vs 2031 GFLOP, D]. Time tracks FLOPs almost exactly, which is what a compute-bound region with no dominant sub-module looks like: 1.28 ms per single block, 1.34 ms per double block, and the double block carries twice the parameters. There is no module-level anomaly to chase here. The largest single op in the whole DiT is fusion.319 at 62.5 ms over 160 instances (2 iterations × 4 steps × 20 blocks) — the single-block's fused GEMM chain, 0.39 ms each [M].

How much of the wall clock is the TPU actually working? Taking the union of top-level spans per process in the XProf trace gives the device's busy time directly. Set against the un-profiled latency from the harness — the profiler inflates wall time, so device busy over traced time would understate occupancy:

region device busy ms [M] latency ms [M] device occupancy idle
text encode 16.0 104.8 15 % 85 %
DiT 4-step scan 131.1 158.1 83 % 17 %
VAE decode 8.0 14.4 56 % 44 %
whole request 155.3 262 59 % 41 %

The device figures agree with the XLA-Ops method above to within a millisecond (16.1 vs 16.0, 130.1 vs 131.1), from an independent part of the trace. So on a fully resident v6e-1 the MXU works for 155 ms of a 262 ms request and is idle for the other 107 ms — and 85 % of the text-encoder stage is idle, which is where nearly all of that 107 ms lives. The host-side busy time from the same trace is not quoted here: it is inflated by the profiler (the traced request takes 906 ms against 262 ms un-profiled, 3.5×), so it is a relative signal, not a budget.

And the source-line track names the lines. Self time on the Source code track, DiT scan:

line what it is self ms [M] calls
sampling.py:52 the Euler lax.scan (contains the block scans) 260.25 26
dit.py:65 the 20-block single_blocks scan 204.71 328
nn.py:20 return x @ w — every linear 112.16 690
dit.py:56 the 5-block double_blocks scan 61.56 2 736
nn.py:65 / nn.py:61 attention PV / QKᵀ einsums 60.33 200
rope.py:47 the fp32→bf16 rope rotation 16.02 840

notes: sampling.py:52 contains the two block scans, so these nest; the 200 attention calls are 2 iterations × 4 steps × 25 blocks, and the 690 linears are the projections and MLPs inside them. For the text encoder the same track gives nn.py:20 at 10.66 ms over 270 calls (27 layers × 10 linears) and the three trunk scans at 32.7 ms — so of the whole 27-layer trunk, the projections and MLPs are 10.7 ms and everything else is normalisation, rope and attention.

The text encoder is the anomaly, and it is not the arithmetic. It is 6.4 % of the image's FLOPs and 40 % of its wall clock, and only 15 % of its latency is device-op time [M]. Its 27 layers occupy 14.5 ms of device time; ops that run once per call — the embedding gather, the RoPE tables, the causal mask, the tap concatenation — occupy 14.2 ms, as much as all 27 layers combined [M]. The remaining ~89 ms of its 104.8 ms is host-side: the trace window is filled with H2D Dispatch / XlaLinearize / Linearize bands on the host threads, not with device work.

That is a different failure than the one predicted in §4. On a GPU the text encoder's problem was small-shape GEMM inefficiency; on TPU it is that the MXU sits idle while the host feeds it. Both produce the same symptom — a 13.7× gap to the calibrated floor — and the fix is the same shape: fewer, larger device ops. Concretely, the ~14 ms of once-per-call prologue is nearly free to remove (build the RoPE tables and the mask once, outside the traced region, and skip re-deriving them per prompt), and the 27-layer trunk is already a scan.

The VAE is the opposite: 67 % device-op time, a flat sequence of ~870 small fusions with the largest at 1.1 ms, and a 2.74× gap [M]. It is small (5.5 % of the image), near its own floor, and not worth optimizing.

MFU, end to end and per component. The whole 512² request runs at 19.0 % of the v6e-1 datasheet (174.3 TFLOP/s) and 45.9 % of what this chip actually achieves on a 4096³ GEMM [D/M]. Per component, against the same two denominators: DiT 28.1 % / 67.8 %, VAE 15.1 % / 36.4 %, text encoder 3.0 % / 7.3 %.

The traces

The whole 512² request with everything resident. Process 3 carries the XLA op launches; the wide teal band is the 4-step DiT scan, and the green blocks are the text encoder and VAE either side of it. Duration 318 ms for two traced iterations.

Perfetto timeline of a whole resident 512² request on v6e-1, one process, showing the DiT scan dominating

The 4-step scan in detail: while.27 is the Euler scan, while.29 the 20-block single stack inside each step, while.28 the 5-block double stack, and the individual fusion slices between them are the per-block GEMM chains.

The 4-step DiT scan: the Euler while loop, the single and double block scans, and per-block fusion slices

A wider window over the same trace, showing how regularly the blocks repeat — 20 single and 5 double per step, at constant width.

Five double-block and twenty single-block iterations of one DiT step, uniform in width

The same 17 ms window again, from the XProf capture. Three things appear that the native trace does not have: a Source code track carrying /content/src/flicker/sampling.py:52 and dit.py:65, a Framework Name Scope track showing jit(denoise) → while → body → closed_call, and /device:TPU:0 and /host:CPU as separate processes — which is what makes the busy-time table above a measurement rather than an inference.

The same DiT window from the XProf capture, with Python source-line attribution and separate host and device processes

The text encoder. This is the finding in one picture: the window is entirely host dispatch machinery — H2D Dispatch, XlaLinearize, Linearize — wrapped around block_until_ready. The device work it is feeding is the thin band at the top.

The text-encoder trace window, dominated by host dispatch bands rather than device work

The VAE decoder, by contrast: approximately 870 small ops, individually visible, tightly packed.

The VAE decoder trace: a dense sequence of small fused ops

Why the traces were re-derived

The first version of this section used the perfetto_trace.json.gz files directly, because they open in Perfetto with no conversion. They are a lossy view: they carry the op names but not the Python frame that issued them, and they flatten host and device into one process. Converting the .xplane.pb alongside them cost one local command and produced the source-line table and the busy times above, so the figures and the attribution here are from the XProf traces. The two agree where they overlap, which is the useful part: the device-op spans and the device busy time were computed from different tracks and land within a millisecond.

Two caveats on every number in this section. The profiler inflates wall time — the traced full_resident iteration measured 421 ms against the harness's 262 ms — so the traces are used for structure and relative attribution, never for latency; every millisecond quoted above comes from the bench harness. And XLA records a while op and the ops inside its body, so the tracks nest: summing every row of a trace double-counts, which is why the table uses top-level spans and says so.

7. The same GPU, both backends

§6 says where the time goes on the TPU; it cannot say whether the port itself is the problem. So both backends were run on one L4 (22.03 GiB, one container per framework — JAX does not hand GPU memory back after clear_caches(), and torch's weight load OOM'd against the 22.01 of 22.03 GiB the first half was still holding), each fed the captured fixture's own inputs (sample/x0, txt/ctx, the step-3 latent). Same weights on both sides (dit.safetensors — the reference's own file — and the diffusers VAE), two warmups, median of ten device-synchronised calls, compile excluded.

stage, 512² JAX L4 [M] BFL torch L4 [M] torch / JAX JAX v6e-1 [M] v6e-1 vs torch-L4
devices on this row Modal L4, 22.03 GiB Modal L4, 22.03 GiB — Colab v6e-1, 32 GiB —
one DiT step 366.7 253.8 0.69× 58.3 4.4×
DiT, 4 steps 1385.1 1021.0 0.74× 158.1 6.5×
VAE decode 145.8 98.1 0.67× 14.4 6.8×
parity vs golden, u8 mae 0.94 0.00 — 0.09 —

On a GPU, XLA is behind eager torch at these shapes, and by a consistent margin. The reference's denoise is 1.36× faster end to end and its decoder 1.49× — the same factor in two structurally unrelated graphs, a 25-block attention stack and a conv decoder. That points at codegen rather than at a bug, and the parity column agrees: this port's image is 0.94 u8 levels from the reference image on the same device, while the reference reproduces its own capture bit-exactly by construction. On the TPU the same comparison inverts.

Against the L4 the TPU port wins by 6.5× on the DiT and 6.8× on the VAE — the L4 is a 22.03 GiB card, the chip a 32 GiB one, so this is also a larger-faster-against-smaller-slower trade. The same JAX code on the L4 is 8.8× and 10.1× slower, while the two chips' measured GEMM rates differ by only 4.7× (80.6 vs 380.1 TFLOP/s) [M] — the port extracts ~1.9× more of the accelerator ratio than the FLOP ratio predicts. The datasheet tells the opposite story, and both are true [S]:

L4 (121 TFLOP/s) datasheet L4 measured (80.6) v6e-1 (918) datasheet v6e-1 measured (380.1)
JAX DiT 4-step 24.3 % 36.5 % 28.1 % 67.8 %
BFL torch DiT 4-step 33.0 % 49.6 % — —
JAX VAE 11.2 % 16.8 % 15.1 % 36.4 %
BFL torch VAE 16.6 % 25.0 % — —

Against their published peaks the L4 numbers are the better ones, and the reason is §1's: this v6e delivers only 41.4 % of its datasheet on a 4096³ GEMM, so a fraction-of-datasheet comparison flatters any accelerator whose peak is reachable. Against what each chip demonstrably achieves, the TPU port is 1.9-2.2× closer to its own roofline than the reference is to the L4's — which is the honest statement of where the port stands.

The text encoder is deliberately absent from this table (§2): the reference loads Qwen/Qwen3-4B-FP8 and this port runs 27 of Qwen3's 36 layers from the klein bf16 shards, so there is no same-input text comparison to run. The cross-check that the protocol is sound is internal: this harness measures four dispatched DiT steps at 1466.8 ms against the scan's 1385.1 ms — the same 1.06× that the earlier L4 run reported [M].

8. What to optimize next

  1. The text encoder, before anything else — and §6 now says which part. 40 % of the wall clock for 6.4 % of the FLOPs at 3.0 % MFU, and only 15 % of its latency is device-op time. Half of that device time is ops that run once per call — RoPE tables, the causal mask, the embedding gather, the tap concat — which is 14.2 ms of the 104.8 ms and is nearly free to hoist. The other half is the 27-layer trunk in 14.5 ms. Whatever the remaining ~89 ms of host-side time is, it is the single largest recoverable block in the pipeline.
  2. int8 on the four big GEMM stacks. single: linear1 and linear2 are 74 % of DiT FLOPs [D]; they sit at unfused AI 522-864 against the 560 ridge [D], so doubling the rate is legitimate where a row is math-bound. The roofline prices the saving at 6.6 ms of the 49.7 ms 512² floor (13 %) and 33.8 ms of the 163.7 ms floor at 1360×768 (21 %) [D] — the fraction is well under a pure FLOP halving because double: mlp (AI 522) is memory-bound too, which is exactly the correction the ridge point predicts (§5.1 prices it in milliseconds). v6e's fp8 is not the alternative: it publishes the same rate as bf16, the signature of emulation [S].
  3. Fuse the norm, rope and modulate rows. The four such rows in the DiT cost 6.9 ms of the 49.7 ms 512² floor (14 %) [D] — a larger share than a lower-ridge part would pay, because their traffic is fixed while the compute peak is not. All twelve module floors are in §5.1, which also shows why: seven of the twelve are memory-bound, and linear1 alone is 46 % of the floor.
  4. splash/Pallas only at long sequences. Attention is 7.2 % of DiT FLOPs at 512² but 18.8 % at 1360×768 [D]; the reason to care is VMEM, not FLOPs. At 512² it buys nothing.
  5. Leave the VAE and the sampler alone. The decoder is 4.4 % of 512² FLOPs, and four steps is already distilled — the caching literature's 1.3-2× came from 28-50-step baselines [S].
  6. Do not chase dispatch overhead. lax.scan already beat four dispatched forwards by 1.057× on the L4 [M], and here it is 1.48× faster per step than a single dispatched step (158.1 ms for four against 58.3 ms for one) [M].

9. What this does not establish

Nothing about the 8-chip layout question. One single chip cannot see a collective: FSDP-8's all-gather is 6.78 GB per chip and 16.96 ms per image at 8 chips, Ulysses is 0.26 ms, TP is 1.45 ms [D] — enough to move the same 512² floor from 6.48 ms (96 % of peak) to 23.17 (27 %). The ranking that follows — shard between requests, not within one — is a claim about ICI that this run does not test.

Nothing publishable as a performance figure. Colab is a managed, shared, ephemeral runtime where transparent hugepages cannot be set (rule 7), so every latency above is a lower bound; the published configuration is the same harness on a gcloud v6e-1 spot VM ($0.243/chip-hour in europe-west4-a [S]), which needs no code change.

Nothing about the 9B. These are 4B numbers; the 9B has 8 double and 24 single blocks at hidden 4096 and does not fit the same envelope.

What it does establish: the first TPU latency, MFU, per-stage profiler attribution and residency measurements for FLUX.2 klein-4B anywhere — no published TPU number for this model exists, from any source — plus four transferable results: an fp32 matmul on TPU is bf16 unless asked otherwise; a single 32 GiB chip serves this model resident at 262 ms using 66 % of its HBM; the analytical model that predicted the DiT and the VAE to within 1.4-2.7× was wrong by 5× about the text encoder; and the profiler attribution that localises it, with the DiT shown to be proportioned exactly to its FLOPs (79 % of scan time against 80 % of scan FLOPs) and the text encoder's MXU idle for 85 % of its latency. §5.1 gives that finding its other half: the model's per-module floor for the same graph.

Where the numbers come from

tools/bench_tpu.py --chip v6e1 --stages cold,staged,resident and then --stages calib,text,dit,vae, one JSON object per run on stdout, ten samples per stage; correctness from tools/tpu_correctness.py; the L4 head-to-head from tools/bench_backends.py (one Modal app, two L4 functions, out/bench/backends.json). The runs are out/bench/*.json and out/tpu/*.json, the transport is the private dataset tensorkelechi/flicker-tpu, and the full bring-up log with every command is docs/tpu-bringup.md. Colab shipped jax 0.7.2 against the 0.11.2 this port is developed on, and the port needed no change [M].

One operational note that is worth more than it looks: keep one session, tear it down at the end. The measurements here took ~14 separate colab run invocations, each of which booted a fresh VM, re-installed jax and re-downloaded the 15.97 GB checkpoint before doing any work — ~91 minutes of VM lifetime for about four minutes of actual compute. colab run self-cleans, which is why nothing was ever orphaned, and that safety is exactly what the provisioning costs. For an iterative phase, colab new -s <name> once and colab exec -s <name> -f script.py per step pays it once.