~/satyajit

KV caching in flow models: a cache you have to train for

mdjsonmcp

2026-10-07 · 25 min · kv-cache · diffusion · diffusion-transformers · image-generation · flow-matching · attention · qwen

Why read this

Notabletop 60%

Why a flow transformer can cache K/V at all, read from the diffusers code, with a FLOP model that predicts all five measured speedups.

  • Interactive explanations
  • Original analysis
  • Widely used

Image & video generationNeeds a workstation GPUNon-commercial licencePractitioner guide

How this was scored
Is it new?
1 of 3: An incremental tweak
Can I trust it?
2 of 3: Measures key facts from files, code or configs
Can I run it?
1 of 3: API-only, gated or restrictive licence
Will I understand it?
3 of 3: Mechanism carried by interactives built from real code or data
Can I act on it?
2 of 3: A concrete recipe, numbers or comparison
Will it last?
2 of 3: A reference for a year or more
Does it affect many?
2 of 3: A widely used model, tool or lab release
Only here?
2 of 3: A teardown or measurement few others did

Score 66 of 100, ranked 158 of 454 rated articles. Each question is answered 0–3 by hand, and a 3 is rare. How articles are scored

When Sayak Paul's post on KV caching in flow models came through, I did not believe the premise. A KV cache works in a language model because the past does not change: token 41 has the same keys and values whether you are generating token 42 or token 4,200. A flow model has no past in that sense. Every denoising step feeds the transformer a new noisy latent, and every token in the sequence can attend to every other one. If everything can see everything, and one thing changes every step, then everything changes every step. There should be nothing to cache.

It turns out the premise is right and my objection is right too, and the gap between them is the interesting part. You can cache K and V in a diffusion transformer only for tokens that have been forbidden from looking at the noisy latents and from knowing what step it is. Nothing in the inference code grants that. The model has to be trained under those rules. FLUX.2 Klein 9B KV is a separate checkpoint for exactly this reason, and Qwen-Image-2.1 ships with the rule switched on in its config.

Sayak's post builds up the idea with pseudocode in the style he credits to Karpathy, then benchmarks the speed and memory trade-off on an A100. He was prompted by the Qwen-Image-2.1 implementation (he thanks @Kun11664638 for it on X) and went back to Klein KV, which he calls "probably the first one to have done it (but only for reference image tokens)". I read the code both of them run in diffusers, rebuilt the cost from the model shapes, and checked his numbers against it. Most of them hold up very well. Two of his side experiments say less than they appear to.

What an LLM's cache gets for free

It helps to be exact about why the language-model cache is valid, because a flow model has to buy both of those properties back.

The first is the causal mask. A token's hidden state at layer ℓ\ell depends only on tokens at or before it, so appending a new token cannot change anything already computed. The second is easy to miss because it is an absence: a language model has no clock. Nothing in the forward pass depends on when you run it. Put those together and the K and V of every earlier token, at every layer, are fixed forever. The cache grows by one row per generated token and each decode step computes only the new row. (The site's walk through LLM inference covers the prefill and decode economics of that.)

Now look at the block a flow model runs. This is the MMDiT block from the Stable Diffusion 3 paper, which Sayak uses as his starting point:

Block diagram of one MMDiT block. Two parallel columns, one for the caption tokens c and one for the noisy image tokens x, each with its own LayerNorm, modulation, linear QKV projections, and MLP. The timestep conditioning y feeds the modulation (alpha, beta, gamma, delta) of both columns. The Q, K and V of the two columns are concatenated into a single joint attention operation, then split back into the two columns.
An MMDiT block: separate weights for text and image, one joint attention over both, and the timestep modulating both streams (Esser et al. 2024, reproduced in Sayak Paul's post, Figure 1).

Both properties are gone. The joint attention is bidirectional, so a text token's output at layer 1 depends on the noisy image tokens, which changes its input to layer 2, which changes its K and V there. And the timestep y drives the shift, scale and gate of the text stream as well as the image stream, so even at layer 1, before any attention has happened, the text tokens' normalised input is different at every step. In a standard MMDiT, like FLUX.1 or SD3, the prompt's keys and values are different at every step and every layer. The text encoder's output is constant, and that is not the same thing.

So the condition for a cache is a set of tokens whose input is constant, whose modulation does not depend on the step, and which cannot read anything that does change. Text and reference images satisfy the first condition by default. The other two have to be imposed on the model.

Klein KV: references that look only at themselves

FLUX.2 Klein 9B KV imposes them on the reference images and nothing else. The whole trick is in one attention function in transformer_flux2.py. At the first step the sequence is [text, reference, target], and the function splits the queries:

# diffusers/src/diffusers/models/transformers/transformer_flux2.py:158-170
# txt+img attend to all tokens
q_txt_img = torch.cat([q_txt, q_img], dim=1)
k_all = torch.cat([k_txt, k_ref, k_img], dim=1)
v_all = torch.cat([v_txt, v_ref, v_img], dim=1)
attn_txt_img = dispatch_attention_fn(q_txt_img, k_all, v_all, backend=backend)
attn_txt = attn_txt_img[:, :ref_start]
attn_img = attn_txt_img[:, ref_start:]
 
# ref tokens self-attend only
attn_ref = dispatch_attention_fn(q_ref, k_ref, v_ref, backend=backend)
 
return torch.cat([attn_txt, attn_ref, attn_img], dim=1)

Text and target queries see every key, references included. Reference queries see only reference keys. A reference token therefore never reads the noisy latent, so the latent's changes cannot reach it through attention. The second leak, the clock, is closed in the model's forward, which builds a separate modulation for the reference tokens from a fixed timestep and splices it into the image stream's modulation at the reference positions:

# transformer_flux2.py:1305-1307
# Ref tokens use a fixed timestep for modulation
ref_timestep = torch.full_like(timestep, ref_fixed_timestep * 1000)
ref_temb = self.time_guidance_embed(ref_timestep, guidance)

ref_fixed_timestep defaults to 0.0 and the pipeline never passes it, so the references are modulated as if at the clean end of the trajectory. That is the natural choice for an input that is a clean image. Sayak's post describes the fixed timestep as "typically the timestep associated with the first iteration", which in this flow convention would be the noisy end. For Klein the code says otherwise; it is a small point, and Qwen-Image-2.1, which he does describe as using zero, makes the same choice.

With both leaks closed, the attention processor writes the reference K and V into the cache on the first step, after RoPE, one entry per block:

# transformer_flux2.py:455-458
if kv_cache_mode == "extract" and kv_cache is not None and num_ref_tokens > 0:
    ref_start = num_txt_tokens
    ref_end = num_txt_tokens + num_ref_tokens
    kv_cache.store(key[:, ref_start:ref_end].clone(), value[:, ref_start:ref_end].clone())

On every later step the pipeline stops feeding reference latents to the model at all. The input is [text, target], and each attention layer splices the stored keys back between the two:

# transformer_flux2.py:134-141
if kv_cache is not None:
    # Cached mode: inject ref K/V between txt and img
    k_ref, v_ref = kv_cache.get()
 
    k_all = torch.cat([key[:, :num_txt_tokens], k_ref, key[:, num_txt_tokens:]], dim=1)
    v_all = torch.cat([value[:, :num_txt_tokens], v_ref, value[:, num_txt_tokens:]], dim=1)
 
    return dispatch_attention_fn(query, k_all, v_all, backend=backend)

The loop in pipeline_flux2_klein_kv.py:794-825 is the prefill and decode split you would expect: step 0 runs with kv_cache_mode="extract" and num_ref_tokens=image_latents.shape[1], and the remaining steps run with kv_cache_mode="cached". The cache covers all 32 blocks, the 8 double-stream and the 24 single-stream ones, because a single-stream block concatenates text and image into one sequence and still has to keep the references from reading forward.

Sayak checks the claim empirically, by recomputing the reference K and V at every step instead of caching them and measuring how far they drift:

Two line charts of the change from step 1, in percent, over denoising steps 1 to 8. Left, reference image tokens: keys and values sit flat at 0% at every step, annotated '0% change, K and V stay unchanged'. Right, generated image tokens: keys and values climb steadily from 0% at step 1 to about 63% at step 8.
Recomputed reference K/V do not move across steps, while the generated tokens' K/V drift by over 60%; a tiny Flux2KleinKVPipeline with random weights, float32 on CPU, 3 runs (Sayak Paul's post, Figure 3).

Note what this was run on: a tiny model with random weights. That is fine, and it is actually the stronger demonstration. The zero is structural. It comes from the mask and the fixed modulation, so it holds for any weights, trained or not. A trained model can only make it hold more convincingly by coincidence, never less.

The cache never reads the prompt

Something falls out of that attention function that neither the post nor the PR mentions. Reference queries attend to k_ref only, so reference tokens never see the text either. Their hidden states at every layer depend on the reference image, its position ids, the fixed timestep and the weights. They do not depend on the prompt.

So the cache Klein builds is prompt-independent. If you edit the same photo ten times with ten different prompts, the reference K and V are the same all ten times, and in principle you could compute them once per image rather than once per call. The pipeline does not expose this: it builds the cache inside __call__ and clears it at the end (pipeline_flux2_klein_kv.py:864-865). I have not tried reusing a cache across calls, so treat this as what the code implies and not as something I ran. It is also the cost of the design. A reference that cannot read the prompt cannot specialise its features to the edit you asked for; all of that has to happen in the text and target tokens that read it.

That trade is why Klein KV is a checkpoint and not a flag. A model trained with full joint attention, where references and text read each other and everything reads the latent, has learned to rely on those paths. Cut them at inference time and you are running a different network from the one you trained. Sayak's benchmark repository makes the same point from the other side: its uncached baseline is a patch that recomputes the references every step while keeping the fixed timestep and the reference-only mask, because, in its README's words, setting kv_cache_mode=None "would change those rules". The fair comparison is the same network with and without the cache, and that is the comparison he ran.

What the speedup should be

Once you know which tokens run at which step, the cost is arithmetic, and you can predict the speedup before anything is timed. That is what the widget below does, with the real token counts from Sayak's benchmarks.

tokens through the blocks at step 1 (7,424 in the full sequence)
reference
noisy target
text: 512
recomputed
reference: 2,816
recomputed
noisy target: 4,096
recomputed
attention at step 1: query rows, key columns
textreferencetarget
textcomputedcomputedcomputed
referencemaskedcomputedmasked
targetcomputedcomputedcomputed
this step: 125.7 TFLOP
full sequence
steps 1-1, cached: 125.7 TFLOP
uncached: 125.7 TFLOP
whole run of 4: 25.9% fewer transformer FLOPs
measured end-to-end latency saving on an A100: 23.0%
cache held: 1.375 GiB in bf16
2 x 32 layers x 2,816 tokens x 4096 x 2 bytes
Derived from the configs and the diffusers masks, not timed. Transformer FLOPs only: the text encoder and VAE are outside the count, which is why the measured saving sits a few points under the model. Solid blue is recomputed this step; hatched teal is K/V read from the cache.

The shapes are these. Klein 9B and Qwen-Image-2.1 are both 32 transformer layers at width 4,096 (32 heads of 128) with a SwiGLU MLP of ratio 3. Per token, per layer, the linear layers cost 13d213d^2 multiply-accumulates: 4d24d^2 for Q, K, V and the output projection and 9d29d^2 for the MLP. That checks out against the names: eight double blocks at 26d226d^2 and 24 single blocks at 13d213d^2 come to 8.72B parameters for Klein, and 32 blocks at 13d213d^2 come to 6.98B for Qwen, whose card says 7B. Attention adds 4d4d FLOPs per query-key pair per layer. The Klein 9B config is gated, so I took the layer split from the benchmark repo's TaylorSeer configuration ("all 8 double and 24 single blocks") and the width from its cache-size formula; the MLP ratio is the open Klein 4B config's.

The token counts come from the benchmark records. The output is 1024 × 1024, which is 4,096 latent tokens in both models. Klein pads the prompt to 512 tokens (max_sequence_length in the pipeline), and a 1024 × 704 reference is 64 × 44 = 2,816 tokens. Three references make 8,448. Qwen's own audit in the benchmark gives a prefix of 4,246 tokens, of which the 54 × 78 reference latent is 4,212, which leaves 34 for the text.

For Klein with one reference, an uncached step pushes 7,424 tokens through the linear layers and a cached step pushes 4,608, because the 2,816 reference tokens are gone. The attention shrinks less, since the text and target queries still read every reference key. That makes the full step 125.7 TFLOP and the cached step 82.3. Only the first step of a run pays full price, so over a run of nn steps the saving is

1−Ffull+(n−1) Fcachedn Ffull1 - \frac{F_{\text{full}} + (n-1)\,F_{\text{cached}}}{n\,F_{\text{full}}}

and it rises with nn toward the per-step saving. Set that next to what Sayak measured, all on an A100 80GB in bf16:

RunStepsPredicted transformer FLOP savingMeasured latency saving
Klein KV, 1 reference425.9%23.0%
Klein KV, 1 reference830.3%28.5%
Klein KV, 3 references446.4%42.1%
Klein KV, 3 references854.1%51.3%
Qwen-Image-2.1, 1 reference4046.6%44.7%

The prediction is above the measurement in every row, by 1.8 to 4.3 points, and that is the direction it has to miss in. His latency is the full pipeline call: text encoding, VAE encoding of the references, the denoising loop, VAE decoding and PIL conversion. My count is the transformer only. The fixed costs the cache cannot touch dilute the saving a little, and they dilute it most in the short four-step Klein runs, which is where the gap is largest. Nothing in his numbers needs explaining beyond the arithmetic, which is what you want from a benchmark.

Two bar charts for FLUX.2 Klein 9B with one reference image on an A100-SXM4-80GB in bfloat16 at 1024 by 1024. Left, full-pipeline latency: 3.09 s without cache and 2.37 s with cache at 4 steps; 5.88 s and 4.21 s at 8 steps. Right, peak GPU memory: 34.73 GiB without and 35.48 GiB with cache at both step counts.
Klein 9B KV with one reference: median of 5 measured runs after 3 warmups, seed 42 (Sayak Paul's benchmark repository, flux2-kv/pretrained results chart).

The table also says when to bother. Klein is distilled for four steps, and at four steps the first step, which pays full price, is a quarter of the run. One reference saves under a quarter of the time. Three references, where the reference tokens outnumber the target two to one, save 42%. The cache pays in proportion to how much of the sequence is conditioning, and in how many steps you amortise the prefill over. A text-to-image call on Klein has no references, so the pipeline falls back to the plain forward and saves nothing.

The outputs did not change. The benchmark repo hashed all 32 saved PNGs from the three-reference run and found one pixel hash per step count across both modes: cached and uncached images are identical, pixel for pixel.

Memory

The cache size is easy to compute and Sayak's post gives the useful rule: for these models each 1,024 cached tokens cost half a GiB. Two tensors, K and V, times 32 layers times 1,024 tokens times 4,096 channels times 2 bytes is exactly 0.5 GiB. One Klein reference is 1.375 GiB of cache, three are 4.125 GiB, and Qwen's 4,246-token prefix is 2.073 GiB.

Peak memory is a different number. With one reference the measured peak rose from 34.73 GiB to 35.48 GiB, which is 0.75 GiB and not 1.375; with three references it rose by 4.124 GiB, almost exactly the payload. The peak is whatever is live at the worst moment, weights included, and the cached run also carries fewer activations in its later steps. I cannot attribute the one-reference gap from the published numbers alone, and the benchmark README is careful to say the same: peak memory "reflects all tensors live at the peak, rather than just the cache payload". For planning, budget the payload. At a 2048 × 2048 target, which is 16,384 tokens, a same-size reference costs 8 GiB of cache.

Qwen-Image-2.1 caches the whole prefix

Klein closes the leaks for references and leaves the prompt alone. Qwen-Image-2.1 closes them for the entire conditioning prefix, prompt included, and it does it with a mask that looks much more like a language model's. The rule is one line in transformer_qwenimage21.py:

# transformer_qwenimage21.py:294-295
same_image_block = (q_image_id == kv_image_id) & (q_image_id >= 0)
allowed = ((q_idx >= kv_idx) | same_image_block) & key_valid[batch_idx, kv_idx]

The sequence is causal, so nothing can read anything after it, except that the tokens of one image may read each other in both directions. Text is strictly causal. A reference can read the text and the references before it. The noisy target comes last and is one image block, so it reads everything. Sayak's figure draws it:

An attention mask grid titled 'Who can attend to whom?', QwenImage 2.1 block-causal attention. Query rows and key columns are grouped as Text A, Reference image 1, Text B, Reference image 2 and Noisy latent tokens, two tokens per group. Text groups form causal staircases, each reference image's block is fully allowed within itself and toward everything earlier, all cells to the right of a prefix group are blocked, and the noisy latent rows are fully allowed across every column.
Qwen-Image-2.1's block-causal mask: no prefix token can attend to the noisy latents, which is what makes the prefix cacheable (Sayak Paul's post, Figure 8).

The mask removes the attention leak for every prefix token. The clock is handled by the causal_condition flag in the transformer config, which is true in the released checkpoint. The model computes one extra timestep row at zero (transformer_qwenimage21.py:932) and every non-target token takes its modulation from that row. The model refuses a cache without it, and the error message says the whole argument in one sentence:

# transformer_qwenimage21.py:895-898
raise ValueError(
    "kv_cache requires `causal_condition=True`. The cache is only valid because text and condition-image "
    "tokens modulate from t=0, which makes their activations independent of the denoising step."
)

On a cached step the model cuts the sequence down to the target before the first block (joint_hidden_states[:, prefix_len:], line 957), and each attention layer prepends the stored prefix keys (key = torch.cat([cached_k, key], dim=1), line 362). The target rows of the mask are all allowed, so the decode step needs no mask at all. All the block-causal structure is paid for once, in the prefill.

Two bar charts for Qwen-Image-2.1, 40 steps on an A100 80 GB in bfloat16 at 1024 by 1024 with one reference. End-to-end latency: 32.20 s without cache, 17.80 s with cache. Peak GPU memory: 36.81 GiB without cache, 38.89 GiB with cache.
Qwen-Image-2.1 with and without the prefix cache: 3 warmups and 5 timed runs, median latency, no classifier-free guidance (Sayak Paul's benchmark repository, qwenimage21-benchmark run chart).

The pipeline turns the cache on by default (use_kv_cache=True) and makes two of them when you use classifier-free guidance, one per branch (pipeline_qwenimage21.py:763-764). Guidance doubles the cache as well as the compute. Sayak benchmarked without guidance, so his 38.89 GiB is the single-cache figure.

One difference from Klein is worth knowing before you compare images. Klein's cached and uncached outputs matched pixel for pixel. Qwen's do not, and the pipeline's docstring says why: caching changes the sequence layout the decode step attends over, so the kernels tile differently and round differently in bf16, and "a one-ULP difference at the first block is then amplified by 32 blocks and every sampler step, so the two settings give equally valid but visibly distinct samples". Both match an fp32 reference to the same tolerance. If you are A/B testing an image, fix the flag.

The site has a longer read of Qwen-Image-2.1 that times this same cache at 2.55x with two references and works through the .clone() detail that keeps it from pinning the whole prefill; and a note on its few-step LoRAs, which shrink exactly the step count this cache amortises over.

Caching Klein's prompt is a guess

This is the experiment Sayak flagged on X as needing more validation, and he is right to. Having seen Qwen cache its text, he tried caching Klein's text K and V from the first step as well.

On Qwen that is exact. On Klein it cannot be, and the reason is in two lines. The text stream's modulation is computed from the current step's embedding, with no fixed-timestep blend like the one the references get:

# transformer_flux2.py:1291-1292
double_stream_mod_img = self.double_stream_modulation_img(temb)
double_stream_mod_txt = self.double_stream_modulation_txt(temb)

And in the attention function above, text queries attend to k_all, which includes the noisy target. Both leaks are open. Klein's text K and V at step 3 genuinely differ from those at step 1, and reusing step 1's is an approximation. His post says as much: it is "an approximation, rather than the equivalent computation provided by reference KV caching".

A two-by-two grid of generated images of a black-and-white cat in a blue wizard hat and cloak, leaning on a wooden ledge against a green wall. Top left: fresh text K/V, the baseline. Top right: text K/V cached in double-stream blocks; the hat and cloak details differ. Bottom left: cached in single-stream blocks. Bottom right: cached in all blocks. All four are coherent wizard-cat edits that differ in costume detail and framing.
FLUX.2 Klein 9B KV at 4 steps, seed 42, with text K/V reused from step 1 in different block types; reference caching is on in every panel and text queries stay fresh (Sayak Paul's post, Figure 6).

The images are coherent. They are also visibly different from the baseline, which is what an approximation looks like. With four steps, the step-1 text states being reused were computed at the noisiest point of the trajectory and then reused at the cleanest. The gain on top of the reference cache was small: 2.385 s to 2.231 s at four steps (6.5%) and 4.227 s to 3.861 s at eight (8.7%), for 0.7% more peak memory. Klein's prompt is padded to 512 tokens against a 4,096-token target, so there was never much to save. I would not ship it on one seed and one prompt, and he does not suggest anyone should.

Feature caching works on a different axis

The diffusion literature already had "caching" before any of this, which makes the name confusing. TeaCache, FORA and TaylorSeer all cache across steps for the same tokens. FORA reuses attention and MLP outputs for a few steps before recomputing them; TeaCache decides when an entire model call can be skipped by watching how much the timestep-modulated input has changed; TaylorSeer stops reusing and starts forecasting, fitting a Taylor expansion to a feature's recent trajectory and extrapolating it. All three bet that a noisy token's features change smoothly from step to step. All three are approximations, and their error grows with the distance between computed steps.

KV caching as Klein and Qwen do it runs on the other axis. It caches across steps for different tokens, the ones that were made constant by construction. It is exact, it needs a model trained for it, and it does nothing for the noisy target, which is still recomputed in full at every step. The two compose in principle, because one shrinks the sequence and the other skips steps for what is left.

What is reusedExact?Needs a model trained for it?
LLM KV cacheK/V of earlier tokensyesno, the causal mask is the model
Condition KV cache (Klein KV, Qwen-Image-2.1)K/V of text and reference tokensyes, when they never read the latent or the clockyes
Feature caching (FORA, TeaCache, TaylorSeer)activations of the noisy tokens from earlier stepsnono

Sayak measured the combination on Klein, with TaylorSeer on top of the reference cache: 2.407 s to 2.033 s at four steps (15.5% faster), 4.264 s to 3.105 s at eight (27.2%), and peak memory from 35.48 to 37.24 GiB, because TaylorSeer keeps its own per-module factors. He reports noticeable quality loss, and the images show it:

A two-by-two grid of the same wizard-cat edit. Left column: reference KV cache only, at 4 and 8 denoising steps; clean fur and smooth fabric. Right column: with TaylorSeer added. At 4 steps the fur and the cloak look rough and oversharpened, with a mottled texture on the hat and stronger background texture; at 8 steps the differences are smaller but texture and contrast still shift.
Reference KV cache alone (left) against reference KV cache plus TaylorSeer (right), FLUX.2 Klein 9B KV, seed 42 (Sayak Paul's post, Figure 10).

I agree with the observation and not quite with how it is framed. The post says the combination "can lead to noticeable degradation", which reads as an interaction between the two. The benchmark has no TaylorSeer-only control, so the data cannot separate an interaction from TaylorSeer simply degrading a four-step model on its own. And the configuration points at the second. The repo's README says the first three steps compute fully and, at four steps, "step 4 is predicted": the last step of a model distilled to four steps, the one that lays down the fine texture, is extrapolated from the three before it. Rough fur and oversharpened fabric are what you would expect from that, cache or no cache. Since the reference cache is exact, I would expect TaylorSeer's error to be the same with or without it. I have not run the control. It is the obvious next run.

Who did it first

Sayak hedges on X: Klein KV was "probably the first". The Klein KV checkpoint went up on Hugging Face in March 2026 and its diffusers support landed as PR #13262, titled "klein 9b kv" and merged on 12 March. Qwen-Image-2.1 arrived with its prefix cache in PR #14804 in September. The idea is older than either. OminiControl2, from March 2025, describes "a conditional feature reuse mechanism that computes condition token features only once and reuses them across denoising steps" for image-conditioned DiTs, which is the same move for the same kind of token. I have read its abstract and not its code, so I cannot say how close the masks are. The hedge in "probably" was the right call.

The post itself is careful about its own provenance in a way I wish more were. A note at the end says the text is primarily his and that Codex was used for "language polishing, code snippets, and coordination of the experiments", and the benchmark repo says he reviewed the code and patches by hand. The benchmarks deserve a word too: pinned checkpoint revisions, 3 warmups and 5 timed runs per setting, alternating mode order, every output PNG hashed, and a tiny float32 CPU check of cache correctness before each GPU job. It is the reason the arithmetic above could land so close. Sloppy benchmarks do not agree with a FLOP count to within four points.

What I would take away

The question I started with has a clean answer. A flow transformer can cache K and V for exactly the tokens that cannot see the noisy latent and cannot see the timestep, and nothing else. A language model gets both properties from its causal mask and the absence of a clock. A diffusion transformer has to be trained with a mask that blocks conditioning tokens from reading forward and a modulation that pins them to one timestep. Klein KV does this for references; Qwen-Image-2.1 does it for the whole prefix. You cannot add it to a model trained with full joint attention, and anything you cache from such a model, like Klein's text, is an approximation however good the pictures look.

The rest is arithmetic. The saving is the share of the sequence that is conditioning, amortised over every step after the first; the cache costs half a GiB per 1,024 tokens on these 4,096-wide, 32-layer models, twice that with classifier-free guidance. It pays most in multi-reference editing at many steps, and least in text-to-image, where Klein saves nothing and Qwen caches a few dozen prompt tokens. If you want feature caching on top, measure it against feature caching alone before you blame the combination.

How I checked

I read Sayak's post in full, along with his X thread and its replies, and took the figures from the Hugging Face bucket his post serves them from. I shallow-cloned huggingface/diffusers at commit 122b1e1 (7 October 2026) and read transformer_flux2.py, pipeline_flux2_klein_kv.py, transformer_qwenimage21.py and pipeline_qwenimage21.py; every quote above gives its file and line at that commit. I shallow-cloned his benchmark repository and read each run's README, the uncached-baseline patch and the TaylorSeer configuration, and pulled summary.json, metadata.json and validation.json from the bucket for the exact medians, token counts and hardware. The PR titles, authors and merge dates come from the GitHub API, and the OminiControl2, FORA, TeaCache and TaylorSeer descriptions from their arXiv abstracts.

The Qwen-Image-2.1 transformer config is public and gave its shape directly. The FLUX.2 Klein 9B configs are gated, so its 8 + 24 layer split and 4,096 width come from the benchmark repo; the MLP ratio of 3 comes from the open Klein 4B config. The resulting 8.72B parameter count matching the "9B" name is my consistency check. The FLOP model counts linear layers and attention in the transformer only, ignoring norms, RoPE, modulation and the embedders, which are small at these widths. I ran no model: the latencies, memory figures and images are all Sayak's, from an A100 80GB. The claims that Klein's reference cache could be reused across prompts and that TaylorSeer alone would show the same degradation follow from the code and the configuration, and neither has been run.

Cite this article

For attribution, please use the following reference or BibTeX:

Satyajit Ghana, "KV caching in flow models: a cache you have to train for", ai.thesatyajit.com, October 2026.

bibtex
@misc{ghana2026kvcachingflowmodels,
  author = {Satyajit Ghana},
  title  = {KV caching in flow models: a cache you have to train for},
  url    = {https://ai.thesatyajit.com/articles/kv-caching-flow-models},
  year   = {2026}
}
share