lab

944 → 346 ms: optimizing FLUX.2 klein-4B on one TPU v6e chip

FLUX.2 klein-4B, raw JAX, bf16, 1024², batch 1, 4 Euler steps, text encoder and VAE included, all weights resident on one Colab TPU v6e-1 (32 GiB, 918 TFLOP/s bf16 datasheet, 1.64 TB/s). From a frozen baseline of 944.67 ms to 345.9 ms per image — 2.73×, 1.06 → 2.89 img/s, 17.5 % → 47.5 % datasheet MFU — with every kept step gated against the PyTorch oracle. Every number is a warmed p50 around jax.block_until_ready(), and every A/B ran old and new code in the same session. Tags: [M] measured, [D] derived, [S] spec.

1. Result

stage request p50 text DiT (4 steps) VAE img/s vs baseline
frozen baseline 944.67 ms 104.81 808.92 42.98 1.059 —
T1 residency + one scan 832.21 ms 7.41 784.04 37.54 1.202 1.135×
D1 flash attention 435.7 ms 7.4 387.7 37.5 2.295 2.17×
D3c fused qk-prep kernel 345.9 ms 7.3 298.0 37.2 2.891 2.73×

notes: [M]. Latest production: 436.3 TFLOP/s achieved, 47.5 % of the 918 datasheet peak and 63.8 % of the 683 TFLOP/s a chain of 8192³ bf16 GEMMs actually reaches on this chip; the device is busy 98.8 % of the request.

Request p50 per stage on one v6e-1; dashed: official BFL and sglang-diffusion on an RTX PRO 6000 (different hardware); dotted: upstream maxdiffusion on a v6e-1.

Per-stage p50 against the official implementation (RTX PRO 6000), upstream maxdiffusion (v6e-1) and the frozen baseline.

For scale, on an RTX PRO 6000 Blackwell (different hardware — an implementation reference, not a chip comparison): sglang-diffusion 675.60 ms, BFL's official repo 739.72 ms [M].

On the same chip type, upstream maxdiffusion (Google's JAX port, commit 1bc5481, unmodified, its shipped splash-attention config) runs the same request in 463.5 ms — text 8.7 / DiT 421.3 / VAE 33.5 ms, p50 of 20 reps [M]. flicker is 1.34× faster, all of it in the DiT (298.0 vs 421.3 ms); their VAE decode is 3.7 ms faster than ours. Not a same-session A/B: maxdiffusion ran on jax 0.11.2 / libtpu 0.0.49, flicker on jax 0.7.2 / libtpu 0.0.23, and its quality was not gated (Appendix B, 2026-10-01).

1b. Same chip, same session: flicker vs maxdiffusion at 512² and 1024²

Both engines back to back in one v6e-1 session, same prompt, same timing boundary (text + denoise + VAE, p50 of 20 reps), each on its own toolchain [M].

batch 1 flicker maxdiffusion (best setting) maxdiffusion (shipped flash)
512² 95.3 ms · 10.5 img/s 112.2 ms · 8.9 img/s 149.7 ms · 6.7 img/s
1024² 343.0 ms · 2.9 img/s 462.9 ms · 2.2 img/s 462.9 ms

Time for one image at 512² and 1024², flicker and maxdiffusion, batch 1.

flicker is 1.35× faster at 1024² and 1.18× at 512² (1.57× against maxdiffusion's shipped config). The gap is the DiT; text and VAE are level at 512².

Per-stage time at 512², batch 1: flicker, maxdiffusion, and maxdiffusion with its flash attention.

At 512² maxdiffusion's flash attention is slower than its own dot-product (132.7 vs 95.2 ms): its block stays 4608 rows for a 1536-token sequence. It also fails to compile above batch 1, so the batch curves below use its dot-product path.

Time per batch against images per second at 512², batch 1 to 16, flicker and maxdiffusion.

Images per second by batch size at 512², flicker and maxdiffusion.

Batching buys nothing at 512² either: flicker's batch 8 equals batch 1 (10.5 img/s), maxdiffusion peaks at batch 1.

2. The two inference paths

stage official (BFL, PyTorch) this port (JAX, TPU)
tokenizer HF AutoTokenizer, chat template Rust tokenizers, same template byte-for-byte, host numpy
text encoder HF Qwen3ForCausalLM, 36 layers, output_hidden_states 27 layers only (taps 9/18/27), one lax.scan, taps from ys
conditioning (1, 512, 7680) = hidden states 9 ‖ 18 ‖ 27 same, padding kept unmasked
DiT attention torch SDPA (flash kernel on GPU) Pallas flash attention, tiles q 2304 / k 1536
q/k/v prep separate ops one Pallas pass: head transpose + QK-RMSNorm + rope
DiT loop 4 Python-level Euler steps 4 steps in one lax.scan, 5 double + 20 single blocks scanned
VAE diffusers AutoencoderKLFlux2 decoder same decoder, NHWC, hand-written ops
weights loaded per process all 13.3 GB resident as device arrays

Numerics are matched cast point for cast point (fp32 RMSNorm → bf16 → × bf16 scale; fp32 rope → bf16; bf16 1000·t; bf16 Euler step; truncating uint8 conversion). Parity: 16-rung ladder 0 FAIL on L4 and on v6e-1 [M].

3. How the latency hid

wall = device-busy + device-idle
       ├─ compute-bound  (limited by FLOP/s)
       └─ memory-bound   (limited by bytes/s)
                           └─ host work, H2D copies, dispatch, waits

Three questions decided every step: (1) is the device idle (occupancy = busy ÷ wall)? (2) for busy time, is each op compute- or memory-bound — achieved FLOP/s and bytes/s against the ridge, 683 / 1.64 ≈ 420 FLOP per byte? (3) does the fix change numerics — then gate it; if not, prove bit-identity.

Frozen baseline, one full 974 ms request in Perfetto: the device is nearly idle for the first ~100 ms (text encoder waiting on a host copy), then the 4-step DiT scan, then the VAE.

Before and after, one second each. The same 1000 ms window from the first text encode, same scale: the frozen baseline fits one request; production bf16 fits two with room to spare; the int8 tier two shorter ones.

Frozen baseline: one 944 ms request fills the 1 s window.

Production bf16 (flash attention + fused qk-prep): two ~347 ms requests in the same 1 s window.

Opt-in int8 tier: two ~300 ms requests.

3b. The TPU facts that picked the fixes

Each property below is a v6e fact that ruled a candidate in or out — the same idea (fuse, quantize, batch) lands differently on a GPU.

v6e property consequence here decision it drove
One TensorCore, one MXU (256×256 systolic array, bf16 in / fp32 accumulate) fed by a vector unit; no thousands of threads to hide latency [S] latency is hidden by DMA prefetch into VMEM, not by switching warps; each tile must carry enough MXU work to cover the next tile's copy large attention tiles (q 2304, k 1536); 128×128 tiles were 3× slower than naive
Large on-chip VMEM (tens of MiB per core, vs ~100-200 KB shared memory per GPU SM) [S] tiles 10× the size GPU flash kernels use fit on chip flash/splash with 1536-2304-row tiles; qk-prep holds a whole (rows, 128) head tile
Vector registers are 8 sublanes × 128 lanes [S] a trailing axis of 2 (rope pairs) or 1 wastes most of each register and forces relayout copies rope as two (L, 128) tables and a lane rotate (pltpu.roll), never a (…, 64, 2) view
Mosaic compiles Pallas kernels to 2-D tiles only [M] strided lane slices lower to an unsupported gather; 3-D in-register transposes are rejected kernel layout work goes into BlockSpec index maps — the DMA does the head transpose for free
Grid steps run sequentially on the core, a BlockSpec addresses whole blocks [S] grid order is a loop, not parallelism; index maps return block indices batch became one more grid axis at no cost; an offset-vs-block bug corrupted only grids ≥ 3 steps
Ridge ≈ 420 FLOP/byte (683 TFLOP/s measured / 1.64 TB/s) [D] anything under ~420 FLOP/B waits on HBM: softmax at ~1 FLOP/B, rope at 0.25 TFLOP/s fuse memory-bound chains (flash, qk-prep) rather than speed up math
DEFAULT precision runs fp32 matmuls as a bf16 MXU pass [M] nn.sdpa's fp32 upcast of q/k/v doubled bytes and bought no precision dropping the upcast (in the flash kernel) cost 0 dB
int8 MXU mode, up to 2× the bf16 rate [S] only realised on large GEMMs (1.6-1.7× at 4096-4608 rows × 18-27 K columns; ~1.0× elsewhere) [M] int8 restricted to the two large GEMMs; the rest stay bf16
32 GiB HBM, one chip [S] all 13.3 GB of weights fit resident; no collectives at all resident serving; batch 1; no sharding study needed
Host → device over PCIe at ~8.4 GB/s measured, dispatch asynchronous [M] a per-call NumPy weight costs its size ÷ 8.4 GB/s, invisible to host timers every weight a device array at load (T1)
XLA is the fuser, and it fuses elementwise chains but not reductions + lane permutes + transposes into one pass [M] layout rewrites at the XLA level move bytes between ops, they do not remove them a Pallas kernel, not a fourth XLA rewrite, for the q/k/v prep
jax and libtpu are a matched pair [M] Colab ships jax 0.7.2 on libtpu 0.0.21.1; XLA tolerates it, Pallas/Mosaic modules do not deserialize pin libtpu 0.0.23 for Pallas stages; compare only within a session
Colab VMs: no hugepage control, occasional mid-session loss [M] absolute MFU is a lower bound; late shipping can lose results ship reports before long lanes; a gcloud VM for publishable absolutes

4. What worked

4.1 Device-resident weights — text 104.8 → 7.4 ms

Frozen baseline text encode: a 70 ms H2D Dispatch of the 742 MiB embedding table, then a thin sliver of actual encoder work.

Production text encode: two back-to-back ~7 ms encodes, no host copy in front.

4.2 One 27-layer scan — text 16.9 → 7.2 ms

Three 9-layer scans over in-jit slices of the stacked (27, …) weights made XLA copy ~1.8 GB per call. One scan whose body emits each layer's output as ys gives the taps as ys[8], ys[17], ys[26] for 70 MB and no copies [M]. 7.2 ms is at the compute bound for the measured GEMM rate. Bit-identical on CPU and L4, not on TPU (max |Δ| 12-102): the two loop shapes compile to different bf16 accumulation orders. Quality: equal distance to the oracle for both (§6).

4.3 Pallas flash attention — DiT 780.6 → 387.7 ms

Frozen baseline DiT, zoomed into single blocks: attention (QKᵀ, fp32 softmax, P·V) dominates each block.

Production DiT, zoomed into single blocks: dense GEMMs now dominate each block; attention is one fused Pallas call.

XProf op-profile buckets of the 4-step DiT across the three kernel generations.

4.4 Fused qk-prep kernel — DiT 387.6 → 298.6 ms

5. What did not work (measured, closed)

attempt result [M] why
per-call jax.jit hoisting 0.096 ms JAX already caches traces
scan unroll 2/4/5 7-10 % slower, +14-20 s compile —
batching B = 2/4/8/16 2.28 / 2.18 / 2.70 / 2.65 img/s vs 2.89 at B=1 4608 tokens per GEMM already saturate the MXU; DiT per image rises with B (298 → 336 ms); VAE is 3× worse per image at B=2/4
step-invariant prologue hoist ~0 txt_in(ctx) + rope tables ≈ 0.04 ms once weights are resident
VAE sub-pixel upsample (fold nearest-2× + 3×3 conv into one low-res conv) 37.2 → 48.9 ms exact (131 dB fp32) and gate-passing, but XLA runs the 4·C_out conv worse
one fused jit for the whole request −2.7 ms not byte-identical (max 5 uint8): one graph changes fusion at stage boundaries
int8 W8A8 (per-token act, per-channel weight) DiT 297.7 → 245.5 ms gate FAIL: LPIPS 0.0327, step SNR 25.3 dB
SmoothQuant int8, two largest GEMMs, α 0.5 / 0.35 DiT ~251 ms, request ~299 ms gate FAIL at a plateau: LPIPS 0.012, DINO p5 0.974-0.979, step SNR 32.3-32.5 dB

Throughput vs batch size for the full request: batch 1 is the optimum at 1024².

Per-GEMM int8 speedup on the real DiT shapes, quantize and dequantize included.

int8 detail. v6e runs int8 at up to 2× bf16, but only on large shapes: the single-block linear1 1.70× and the double image mlp.0 1.60× with quant + dequant included; the rest 0.89-1.13× [M]. Those two GEMMs carry ~98 % of the gain and nearly all the error (quantizing linear1 alone: 30.4 dB latent SNR vs 30.7 for both). Per-token W8A8 plus SmoothQuant saturates at LPIPS ≈ 0.012 — 1.9× the implementation floor, yet ~3.4× closer to the oracle than the native text encoder. Passing needs a different scheme (per-group scales, a low-rank bf16 branch as in SVDQuant, or quantization-aware tuning). Shipped as an opt-in lossy tier on main: Klein4B(gemm="smooth_large") plus a calibration file from flicker.calibrate.absmax (8 held-out prompts, α 0.35) — request ~299 ms, 3.34 img/s, LPIPS 0.012. bf16 stays the default.

6. Quality gating

Floors were calibrated on the isolated DiT/VAE lane (oracle ctx, JAX on L4 vs PyTorch on L4, 32 prompts × 3 seeds = 96 pairs): LPIPS median ≤ 0.00638, DINO p5 ≥ 0.98345, step-latent SNR ≥ 33.19 dB [M]. Never retuned.

lane (96 pairs vs oracle) LPIPS median DINO p5 step SNR verdict
frozen v6e-1 lane 0.00548 0.98571 33.92 pass
D1 flash 0.00528 0.98568 33.75 pass
D3c qk-prep (production) 0.00497 0.98435 35.65 pass
SmoothQuant α 0.35 0.01197 0.97930 32.34 fail

LPIPS median vs the PyTorch oracle over 96 pairs; dashed: the calibrated floor.

Three gating lessons: (1) compare to the oracle, never to the previous implementation — a gate that measured the one-scan encoder against the three-scan one failed a pure loop restructure; against the oracle the two are equidistant (ΔLPIPS +0.00003, 16 closer / 16 further). (2) Pre-register the rule; prefer paired statistics to per-lane tails (a p5 over 32 prompts is one or two images). (3) Use the sample the floors were calibrated on: the unchanged production lane scores LPIPS 0.00653 on seed 0 alone — above the floor — and 0.00548 on all 96 pairs.

The text encoder's own gap to PyTorch (LPIPS 0.041 end to end) is bf16 sensitivity, not a porting bug: a loop-shape change alone moves padded-token conditioning to 22.5 dB and images to LPIPS 0.023.

7. Where the time goes now

bucket (DiT, 297.6 ms/scan) ms share rate [M]
dense GEMMs 161.3 54 % 701 TFLOP/s — at the measured bf16 peak
flash attention 97.6 33 % 270 TFLOP/s (exp/rescale on the vector unit)
copies (single-block linear1 output split) 22.7 8 % —
qk-prep kernel 12.3 4 % memory-bound
norms/modulate, rope 2.6 1 % —

Per stage, frozen baseline → production, with the opt-in int8 tier.

Datasheet MFU of the full request per stage; dashed: the measured bf16 GEMM ceiling.

Full request 344.9 ms of device time per image, 4.3 ms idle. What is left is compute: bf16 GEMMs at peak (only lower precision moves them) and attention's vector-unit work.

8. Method, transferable

  1. Occupancy first, then per-op roofline (XProf op_profile, not name heuristics — a heuristic rollup gave GEMMs an impossible 66 ms).
  2. Calibrate the denominator. One 4096² GEMM reads 400 TFLOP/s (launch-ramp limited); 20 chained 8192² GEMMs read 683. Every "%MFU" before that was 1.7× pessimistic.
  3. Same-session A/B, versions recorded. A vendored copy of the committed code (flicker_base) loaded beside the new code gives an exact old-vs-new in one session, one weight load.
  4. Prove the B arm ran. A 0.9996× "speedup" with identical HLO was jit reusing a trace cached before a monkeypatch. Kernel choices are now static config fields — part of the cache key.
  5. Microbenchmarks mislead on fusible ops; judge at the DiT level.
  6. Pallas on TPU: put layout in BlockSpec index maps (block indices, not offsets); keep the body 2-D on (rows, 128) tiles; interpret mode checks numerics, the chip checks compilability.
  7. Dry-run every non-accelerator path on CPU with tiny weights — it caught five plumbing bugs (stacked-dict indexing, AOT static arguments, the patch guard, index maps, a scorer overwrite).
  8. Ship results before the long tail. The Colab VM died late in two sessions; results are now shipped before the quality lanes.

Appendix A: optimization mental model (teaching notes)

Written at the flash-attention stage (435.7 ms, 2026-09-26); the later stages are in Appendix B. Every number is measured.

FLUX.2 [klein] 4B, 1024², batch 1, 4 steps, bf16, one TPU v6e-1. Every number is measured; sources in results.md, metrics-progress.md, out/tpu/*.

0. Where we are

stage request p50 text DiT VAE vs baseline
frozen baseline 944.67 ms 104.81 808.92 42.98 —
T1 residency + scan27 832.21 ms 7.41 784.04 37.54 1.135×
D1 flash attention 435.7 ms 7.4 387.7 37.5 2.17×

Quality held at every kept step: the D1 image gate (96 pairs vs the PyTorch oracle) lands at or better than the frozen TPU lane. sglang-diffusion on an RTX PRO 6000 Blackwell is 675.6 ms.

1. The mental model: where latency hides

wall time = device-busy time  +  device-idle time
            ├─ compute-bound ops   (limited by FLOP/s)
            └─ memory-bound ops    (limited by bytes/s)
                                    └─ host work, H2D/D2H transfers, dispatch, waits

Three questions, in order, decide every optimization:

  1. Is the device idle? occupancy = device-busy ÷ wall. Text encoder was 15 % → the problem was not the math. DiT was 97 % → the problem was the math (or its memory traffic).
  2. For busy time: compute- or memory-bound? Compare each op's achieved FLOP/s and bytes/s to the chip's two ceilings. v6e-1: ~683 TFLOP/s (measured, chained bf16 GEMM) and ~1.64 TB/s HBM. The ridge is 683/1.64 ≈ 420 FLOP per byte. Below it an op waits on memory; above it, on the MXU.
  3. Does the fix change numerics? If yes, it is gated against the oracle; if not, prove bit-identity.

The two cheap sanity tools that caught most mistakes:

2. What worked, and why

2.1 Device-resident weights (text 104.8 → 7.4 ms; DiT −25 ms; VAE −5 ms)

2.2 One 27-layer scan instead of three sliced ones (text 16.9 → 7.2 ms)

2.3 Fused (flash) attention in the DiT (DiT 780.6 → 387.7 ms; request 1.90×)

3. What did not matter (measure before fixing)

4. Method lessons (the transferable part)

  1. Occupancy first, then per-op roofline. Idle device → host/transfer problem. Busy device → compare achieved FLOP/s and B/s to the ceilings; a far-below-peak op with high B/s is memory-bound.
  2. Calibrate the denominator. A single 4096² GEMM read 400 TFLOP/s (launch-ramp limited); 20 chained 8192² GEMMs in one jit read 683. Every "%MFU" before that was understated by 1.7×.
  3. Same-session A/B only. Different sessions/toolchains move numbers (libtpu 0.0.21 → 0.0.23 changes the compiler). Always time the old path next to the new one.
  4. Prove the B arm ran. D0's first full-DiT A/B was 0.9996× with identical HLO: jit reused a trace cached before a monkeypatch. Guard: the two outputs must differ; better, make the choice a static config field so it is part of the cache key.
  5. Toolchains are pairs. jax 0.7.2 declares libtpu 0.0.23; Colab shipped 0.0.21. XLA tolerated it; Pallas (serialised Mosaic modules) did not. Custom kernels break first.
  6. Gates must measure distance to the truth. Comparing a change to the previous implementation (G1) fails any reordering. Compare to the oracle, pre-register the rule, prefer paired statistics over per-lane tails (a p5 over 32 samples is one or two images), and use the sample size the floors were calibrated on (96 pairs, not seed 0 alone).
  7. Dry-run everything that is not the accelerator. Tiny-weight CPU runs of each bench caught three plumbing bugs (stacked-dict indexing, AOT static args, the patch guard) before paying for a TPU.

5. Transfer checklist (e.g. hummingbird / Stable Audio 3)

6. What is left in the DiT (from the op profile, pre-flash)

Dense GEMMs ~161 ms (already ~700 TFLOP/s — lever is int8, D4), rope ~87 ms (fp32 on an f32[1,24,4608,64,1,2] layout — memory-bound, D3), copies/concats ~34 ms, norms/modulate ~15 ms, attention now ~116 ms. Step-invariant work (D2) is small but free.

Appendix B: optimization log

The rolling log kept during the work, one entry per plan or session, unedited. Paths refer to the flicker repository.

2026-09-26 — text encoder, phase T1 started

2026-09-26 — T1 landed (uncommitted), reviewed

2026-09-26 — T2 landed: one v6e-1 session, residency + scan27 priced, G1 fails

Code d1f5e17, bundle 4f89a325…, jax 0.7.2, TPU v6 lite. Full write-up: results.md; evidence out/tpu/t2/; archive t2-v6e1-4f89a325a2a754f8.tar.gz on tensorkelechi/flicker-tpu.

2026-09-26 — T2 review (planner)

2026-09-26 — G1′ launched; DiT plan drafted

2026-09-26 — G1′ result

2026-09-26 — D0 part A: exact op profile of the T2 DiT trace (local, no TPU)

tools/d0_op_profile.py reads the shipped dit_scan xplane with XProf's op_profile tool data: per-HLO-op self time, FLOPs and bytes accessed (rawTime is picoseconds and inclusive of children, so self = node - sum(children); that partition sums back to the root exactly for all three metrics). Normaliser: the trace holds iterations = 3 DiT scans, read from out/tpu/t2/t2_out/report.json, never guessed. Artifact out/tpu/d0/op-profile.json.

bucket ms/scan % dev TFLOP/s GB/s
attention.pv_dot (bhls,bhsd->bhld) 185.705 23.66 71 1125
dense_gemm (projections/MLPs) 161.176 20.53 702 431
attention.softmax (exp/reduce-sum/div over 24x4608x4608) 157.752 20.09 1 1292
attention.qk_dot (bhld,bhsd->bhls, fused reduce-max) 134.216 17.10 98 1540
rope 87.118 11.10 0.25 1279
copies/transposes/concats 33.854 4.31 0.55 1647
norms/modulate 15.103 1.92 1.41 1025
other 8.221 1.05 0.01 5276

Device 785.0 ms/scan, idle 1.9 → device-busy 783.1 ms/scan, against 783.2 ms measured independently by the trace-union occupancy in T2 (-0.00 %). 139.7 TFLOP/scan at 177.9 TFLOP/s (reproduces the roofline 1.3962e14 and T2's 178.08 TFLOP/s), 915 GB/scan at 1166 GB/s.

2026-09-26 — D1 implemented (attention switch in src/, stage d1, scorer)

2026-09-26 — D2+D3 implemented (rope relayout, hoisted prologue, stage d2)

2026-09-27 — D2+D3 session: no win, not committed

2026-09-27 — D3b implemented (reshape+flip swap, load-time linear1 split, stage d3b)

2026-09-27 — D3b: wash, branch not merged

2026-09-27 — D3c implemented (fused Pallas qk-prep kernel, stage d3c)

Branch opt/qk-prep from main (d734102); opt/dit-prep (b29c399) is not merged and not built on. The three XLA-level attempts moved the q/k/v prep between buckets, so this is the kernel.

2026-09-27 — D4 implemented (batch sweep, W8A8 int8 block GEMMs, load-time linear1 split; stage d4)

Branch opt/d4 from main (5f1b731) — D3c shipped — plus the uncommitted batch-axis change to kernels.qk_prep (grid (batch, row_tile, head), index maps take (b, i, h)); opt/dit-prep is still unmerged and not built on. The D4 A/B base is main itself (5f1b731), not D3c's d734102.

2026-09-27 — D4: batching, int8, full-request profile

2026-09-27 — D5 implemented (SmoothQuant int8, sub-pixel VAE upsample, one fused request; stage d5)

Branch opt/d5 from main (9c8b6d8); the D4 int8 code was taken over from opt/d4 as four files (git checkout opt/d4 -- src/flicker/{nn,loader,config,modules}.py) and extended. opt/d4 is not merged and not built on. The A/B base is main itself.

2026-09-27 — D5 session 1 (partial)

2026-09-27 — D5 session 2 (VM died after the trace; quality archives shipped)

2026-09-27 — D6 implemented (four one-variable SmoothQuant folds, fused byte-identity, early shipping; stage d6)

Branch opt/d5 (D6 is built on the D5 tree, uncommitted and not committed by this work). The A/B base is still main 9c8b6d8: D6 spends no new mechanism, it spends the neighbourhood of the D5 near-miss (sm-a050-large: both large GEMMs, α=0.5, DiT 250.8 ms, gate FAIL LPIPS 0.0122 / DINO p5 0.9743 / step SNR 32.50 vs 33.19) — and it re-measures that arm in the same session as the bar.

2026-09-27 — D6 (session completed, early ship worked)

2026-09-27 — gallery session + article assets

2026-10-01 — reference: upstream maxdiffusion klein-4B on v6e-1 (outside flicker)

AI-Hypercomputer/maxdiffusion @ 1bc5481, unmodified, own venv on a Colab v6e-1 (session mdx, stopped): python 3.12, jax 0.11.2, libtpu 0.0.49, flax 0.12.10, tokamax 0.0.14 (their generated_requirements, resolved today). 1024², 4 steps, bf16, default prompt, 20 timed reps after AOT compile + warmup. Tool tools/bench_maxdiffusion.py; logs, images, report.json in out/tpu/maxdiffusion/.

variant stage sum p50 e2e p50 Qwen3 / denoise / VAE (ms) img/s AOT compile
flash (splash, shipped config), B=1 463.5 ms 553.4 8.7 / 421.3 / 33.5 2.157 68.7 s
dot_product, B=1 512.6 ms 602.8 8.7 / 470.5 / 33.4 1.951 70.0 s
dot_product, B=4 2581.5 ms 2929.3 27.9 / 2152.8 / 400.9 1.549 22.7 s
flash, B=4 compile fails — — — —

2026-10-01 — 512² and a same-session A/B: flicker vs maxdiffusion on one v6e-1

One Colab v6e-1 session (res, stopped), both engines back to back, same prompt, text / denoise / VAE each closed by block_until_ready, p50 of 20 reps after compile + warmup. flicker 8d11fc5 on jax 0.7.2 / libtpu 0.0.23; maxdiffusion 1bc5481 on jax 0.11.2 / libtpu 0.0.49. Tool tools/bench_res.py; logs and report.json in out/tpu/res512/.

case flicker maxdiffusion, dot-product maxdiffusion, flash (shipped)
512² B=1, stage sum 95.3 ms (7.3 / 79.6 / 8.4) 112.2 ms (8.6 / 95.2 / 8.4) 149.7 ms (8.6 / 132.7 / 8.4)
512² B=1, img/s 10.49 8.91 6.68
1024² B=1, stage sum 343.0 ms (7.3 / 298.5 / 37.2) — 462.9 ms (8.6 / 420.9 / 33.4)
512² img/s by batch B=1 B=2 B=4 B=8 B=16
flicker 10.49 9.06 8.73 10.51 10.16
maxdiffusion, dot-product 8.91 7.04 6.25 7.09 6.58