KV Caching in Flow Models

KV caching in flow models for image generation.
flow
diffusion-transformers
diffusion
kv-caching
Published

October 7, 2026

KV caching [1] stores the key and value tensors from tokens a transformer has already processed. When generating the next token, the model computes keys and values only for that new token and reuses the cached ones for earlier tokens. This speeds up generation, especially for long sequences, but the cache uses memory that grows with sequence length.

Historically, KV caching has always been linked to autoregressive models (like language models) where the token being generated can only attend to the previous tokens. So, the information flow is unidirectional here (left-to-right). The type of attention used here is causal. We enforce causality through a causal mask.

# Pseudo code. Don't chase me if it doesn't run.
def attention(x, cache=None):
    q = project_q(x)
    k = project_k(x)
    v = project_v(x)

    if cache is not None:
        k = concat(cache.keys, k, dim="sequence")
        v = concat(cache.values, v, dim="sequence")

    scores = (q @ k.transpose(-2, -1)) / sqrt(head_dim)
    scores = apply_causal_mask(scores)  # No token can see future tokens.
    weights = softmax(scores, dim=-1)
    output = project_out(weights @ v)

    return output, KVCache(keys=k, values=v)

There is another family of models, MMDiT (Multimodal Diffusion Transformers [2]), typically used in media generation (such as images and videos). These models often operate across multiple modalities. A popular example is text-to-image generation. There, we need to attend between text and image tokens in their respective representation spaces. In these models, the information flow is made bidirectional. This is made clear in code:

def mmdit_attention(image_tokens, text_tokens):
    qi, ki, vi = image_qkv(image_tokens)
    qt, kt, vt = text_qkv(text_tokens)

    q = concat(qi, qt, dim="sequence")
    k = concat(ki, kt, dim="sequence")
    v = concat(vi, vt, dim="sequence")

    weights = softmax((q @ k.transpose(-2, -1)) / sqrt(head_dim), dim=-1)
    attended = weights @ v  # Bidirectional: image and text can see each other.

    image_out, text_out = split(
        attended, lengths=[len(image_tokens), len(text_tokens)]
    )
    return image_out_proj(image_out), text_out_proj(text_out)

MMDiTs are generally used in flow-based [3] media generation pipelines. This makes MMDiTs iterative — we start from pure random Gaussian noise, which is iteratively denoised over a few steps to produce the final output. In a flow-based pipeline, MMDiT serves as the “denoiser”. The final output depends on the space we’re operating on; it can be pixel-space (for images and videos) or latent-space.

Note

For the rest of this post, we’ll assume latent-space models, but the commentary and findings should carry over to pixel-space models, too. Flow and diffusion are often intertwined but they are different. Even though the post focuses on flow-based models but it is applicable to diffusion models, too.

For each step in this denoising process, we’re always computing the QKV projections for the input tokens across all the MMDiT blocks. Similar to autoregressive models, can we reuse any KV projections from earlier steps in the current step? Let’s assess our options:

If we were to draw parallels to KV caching in autoregressive models, the condition tokens seem like the best fit for KV caching in this case. Several questions should arise at this point:

The rest of this post discusses these questions.

As you might have guessed, we’ll use code snippets to cover some of these sections. So, expect pseudo Python and PyTorch code. Benchmarking code is available in this repository.

Note

The readers are expected to know (at least at a high level) MMDiTs and how they’re implemented and used in text-to-image generation. If you’re looking for a quick intro, the Flux.1 official codebase is a good one.

An MMDiT block

This post focuses on a particular class of models — MMDiTs. Let’s see how a single MMDiT block looks like:

Figure 1: MMDiT block diagram. Taken from the original paper [2].

Translating this to code is straightforward:

  • Each of the two modalities c (text tokens) and x (noisy latent tokens) share the same set of operations. Those operations could be contained in a block, with each modality having its own block so that their parameters are different.
  • y is timestep. It is needed because, in the grand scheme of a flow-based generation pipeline, the MMDiT needs to be aware of its position in the iterative denoising process. Timestep is modulated across both modalities through a set of parameters.
class MMDiTBlock:
    # Add parameters in init.
    ...
    
    def forward(self, image_tokens, text_tokens, cond):
        # `cond` is timesteps.
        # Each stream gets its own conditioning parameters.
        i_shift_a, i_scale_a, i_gate_a, i_shift_m, i_scale_m, i_gate_m = (
            self.image_modulation(cond)
        )
        t_shift_a, t_scale_a, t_gate_a, t_shift_m, t_scale_m, t_gate_m = (
            self.text_modulation(cond)
        )

        image_for_attn = modulate(
            self.image_norm1(image_tokens), i_shift_a, i_scale_a
        )
        text_for_attn = modulate(
            self.text_norm1(text_tokens), t_shift_a, t_scale_a
        )
                
        # `mmdit_attention()` is from above.
        image_attn, text_attn = mmdit_attention(
            image_for_attn, text_for_attn
        )
        image_tokens += i_gate_a * image_attn
        text_tokens += t_gate_a * text_attn

        image_tokens += i_gate_m * self.image_mlp(
            modulate(self.image_norm2(image_tokens), i_shift_m, i_scale_m)
        )
        text_tokens += t_gate_m * self.text_mlp(
            modulate(self.text_norm2(text_tokens), t_shift_m, t_scale_m)
        )

        return image_tokens, text_tokens

def modulate(x, shift, scale):
    return x * (1 + scale) + shift

This resembles the original MMDiT block. More recent models use variants of it. It should also be noted that for simplicity, the snippet doesn’t show the additional techniques that are normally included, such as [4], QK-norm, and parallel attention blocks [5].

The mmdit_attention code snippet shows how attention is generally computed in MMDiT architectures. This form of attention is also referred to as “joint attention”.

Put simply, in joint attention, the QKV projection is performed separately (with separate sets of parameters) on each of the two modalities shown above (image_tokens and text_tokens). Before computing attention, these projections are concatenated. Quoting the paper that introduced this style of attention [2]:

Since text and image embeddings are conceptually quite different, we use two separate sets of weights for the two modalities. […], this is equivalent to having two independent transformers for each modality, but joining the sequences of the two modalities for the attention operation, such that both representations can work in their own space yet take the other one into account.

A more modern variant of MMDiT utilizes part joint attention and part self-attention:

# Regular MMDiT blocks.
for block in self.double_blocks:
    context, x = block(
        image_tokens=image_tokens,
        text_tokens=text_tokens, 
        ...
    )

# Concatenate.
context_x = torch.cat([context, x], dim=1)  

# Continue with the rest.
for block in self.single_blocks:
    x = block(image_tokens=context_x, ...) 

Modern media generation pipelines go beyond just text-to-image. They’re capable of accommodating image-level conditions, which are useful for tasks like image editing and pose-guided image generation. How do we incorporate these additional conditions into MMDiTs?

Note

It is better to think of the process of going from an input (like a prompt) to the final output (like an image) as a “pipeline” rather than a model. This is because it involves multiple models, non-parametric components, and coordination between them. Figure 2 presents a simple visual representation of how information usually flows in a latent-space flow pipeline.

Figure 2: Information flow in standard flow-based generation pipelines (for illustration purposes only).

Conditions in MMDiTs

First, we need a way to encode the incoming conditions and reduce them to a representation that the MMDiT was trained on. For text, this is usually done with a dedicated language model. For images, a dedicated vision encoder (such as DINO [6], SigLIP [7], etc.) can be used. One can also leverage the encoder component of the Autoencoder involved in a latent-space flow pipeline. Popular pipelines like Flux2 [8] and QwenImage Edit [9] follow this approach.

Once the image representations are computed, they’re concatenated with the noisy image latent being denoised:

num_generated = latents.size(1)
num_reference = image_latents.size(1)

for _ in range(denoising_iters):
    latent_image_inputs = concat(latents, image_latents, dim="sequence")
         
    noise_pred = denoiser(
        image_tokens=latent_image_inputs, 
        text_tokens=text_tokens, 
        ...
    )
    noise_pred = noise_pred[
        :, num_reference : num_reference + num_generated
    ]
         
    latents = # compute latents from the prediction with Euler
         

Within latent_image_inputs, only latents change with each denoising step. This is why we slice the prediction.

As mentioned earlier, text-level and image-level representations differ. So, we keep them separate for the input conditions and don’t mix them. The image tokens now are a concatenation of the condition image tokens and the noisy latent tokens. Once their KV projections are computed, do the projections belonging to the condition image tokens change across steps? Let’s find out.

Room for reuse?

Figure 3 confirms our intuition that the KV projections for the condition, aka reference image tokens, don’t change over the denoising steps.

Figure 3: Recomputed KV projections of the condition image tokens across denoising steps. Generated image tokens correspond to the noisy ones being denoised.

This means that after the KV projections are computed in the first step, they can be cached and reused in the later steps. Let’s call this first step “prefill” and later steps “decode”.

There are some important details to keep in mind:

  • Reference tokens use a fixed timestep throughout the denoising process (typically the timestep associated with the first iteration).
  • The concatenated image tokens are only passed to the denoiser during the first step. Once the KV projections are computed and cached, reference image tokens are discarded.
  • Reference image tokens only self-attend; however, the noisy latent and text tokens attend to all tokens. How is this related to KV caching, though? This merits a section of its own.

KV cache design in MMDiT

Note

There are multiple ways in which a KV cache can be designed and integrated into an MMDiT. Here we present just one approach. Readers are encouraged to brainstorm other feasible approaches.

To integrate KV caching for the reference image tokens into our earlier MMDiT blocks, what do we need?

A standard model would have multiple MMDiT blocks. Each block would have its own set of parameters to compute QKV projections. So, we need to cache KV projections per block. Each layer-level cache should be able to store the KV projections, and users should be able to fetch them as needed. Mapping these requirements to a class would look something like this:

class KVLayerCache:
    def __init__(self):
        self.k_ref: torch.Tensor | None = None
        self.v_ref: torch.Tensor | None = None

    def store(self, k_ref: torch.Tensor, v_ref: torch.Tensor):
        self.k_ref = k_ref
        self.v_ref = v_ref

    def get(self) -> tuple[torch.Tensor, torch.Tensor]:
        return self.k_ref, self.v_ref

    def clear(self):
        self.k_ref = None
        self.v_ref = None

The model-level cache would then consist of KVLayerCache:

class KVCache:
    def __init__(self, num_double_layers: int, num_single_layers: int):
        self.double_block_caches = [KVLayerCache() for _ in range(num_double_layers)]
        self.single_block_caches = [KVLayerCache() for _ in range(num_single_layers)]
        self.num_ref_tokens: int = 0

    def get_double(self, layer_idx: int) -> KVLayerCache:
        return self.double_block_caches[layer_idx]

    def get_single(self, layer_idx: int) -> KVLayerCache:
        return self.single_block_caches[layer_idx]

    def clear(self):
        for cache in self.double_block_caches:
            cache.clear()
        for cache in self.single_block_caches:
            cache.clear()
        self.num_ref_tokens = 0

Recall that modern variants of MMDiT do the classic MMDiT-style attention for the first couple of blocks, and then they turn to pure self-attention, operating on a single concatenated representation space. KVCache accommodates that design.

Now, for the more important part: how is KVCache integrated into the model? When the model is invoked for the first time (prefill), we’ll extract the cache. For the subsequent steps, we’ll reuse the cache. Let’s trace this from the forward() of the model:

def forward(
    ...,
    kv_cache: "KVCache | None" = None,
    kv_cache_mode: str | None = None,
    num_ref_tokens: int = 0,
    ref_fixed_timestep: float = 0.0,
):
    ...
    if kv_cache_mode == "extract" and num_ref_tokens > 0:
        kv_cache = KVCache(
            num_double_layers=len(self.transformer_blocks),
            num_single_layers=len(self.single_transformer_blocks),
        )
        kv_cache.num_ref_tokens = num_ref_tokens
        
        # modulate reference timestep embeddings, blend modulation
        # params, etc.
        ...

How is a single KVLayerCache populated inside KVCache?

# Within the same forward
def forward(
    ...,
    kv_cache: "KVCache | None" = None,
    kv_cache_mode: str | None = None,
    num_ref_tokens: int = 0,
    ref_fixed_timestep: float = 0.0,
):
    ...
    for block in enumerate(self.transformer_blocks):
        if kv_cache_mode is not None and kv_cache is not None:
            kv_cache = kv_cache.get_double(index_block)
            encoder_hidden_states, hidden_states = block(
            ..., 
            kv_cache=kv_cache
        )

The next important detail is about how the kv_cache is used during attention computation, depending on the kv_cache_mode.

def mmdit_attention(
    image_tokens,
    text_tokens,
    kv_cache=None,
    kv_cache_mode=None,
    num_ref_tokens=0,
):
    # Extract mode: image_tokens = [ref, current image].
    # Cached mode:  image_tokens = [current image].
    qi, ki, vi = image_qkv(image_tokens)
    qt, kt, vt = text_qkv(text_tokens)

    q = concat(qt, qi, dim="sequence")
    k = concat(kt, ki, dim="sequence")
    v = concat(vt, vi, dim="sequence")

    n_txt = num_tokens(text_tokens)
    n_image_tokens = num_tokens(image_tokens)

    if kv_cache_mode == "extract":
        # First denoising step: separate reference and current-image tokens.
        n_img = n_image_tokens - num_ref_tokens
        lengths = [n_txt, num_ref_tokens, n_img]
                
        # Split q, k, v into 
        # q_txt, q_ref, q_img
        # k_txt, k_ref, k_img
        # v_txt, v_ref, v_img
        ...

        kv_cache.store(k_ref.clone(), v_ref.clone())

        # Text + current image attend to all tokens; refs attend to refs.
        txt_img = attend(concat(q_txt, q_img, dim="sequence"), k, v)
        ref = attend(q_ref, k_ref, v_ref)

        txt_out, img_out = split(txt_img, [n_txt, n_img], dim="sequence")
        attended = concat(txt_out, ref, img_out, dim="sequence")

    else:
        if kv_cache_mode == "cached":
            # Subsequent denoising steps: insert saved reference K/V.
            k_ref, v_ref = kv_cache.get()
            k = concat(k[:, :n_txt], k_ref, k[:, n_txt:], dim="sequence")
            v = concat(v[:, :n_txt], v_ref, v[:, n_txt:], dim="sequence")

        # Cached mode or ordinary attention without caching.
        attended = attend(q, k, v)

    text_out, image_out = split(
        attended, [n_txt, n_image_tokens], dim="sequence"
    )
    return image_out_proj(image_out), text_out_proj(text_out)

Now, we know the different pieces needed to integrate KVCache into an MMDiT-based model. Otherwise, the model won’t have access to the populated cache. The final piece remaining is how we pass the KVCache to the forward() the model, because it accepts a KVCache:

def forward(
    ...,
    kv_cache: "KVCache | None" = None,
    kv_cache_mode: str | None = None,
    num_ref_tokens: int = 0,
    ref_fixed_timestep: float = 0.0,
):

This is important because it’d also reveal how the KVCache is used in the iterative denoising process:

kv_cache = None

for i, t in enumerate(denoising_iters):
    if i == 0 and image_latents is not None:
        # First step: process reference tokens and cache their K/V.
        noise_pred, kv_cache = transformer(
            ...,
            hidden_states=concat(image_latents, latents, dim="sequence"),
            kv_cache_mode="extract",
            num_ref_tokens=num_tokens(image_latents),
        )

    elif kv_cache is not None:
        # Later steps: reuse cached reference K/V.
        noise_pred = transformer(
            ...,
            hidden_states=latents,
            kv_cache=kv_cache,
            kv_cache_mode="cached",
        )[0]

    else:
        noise_pred = transformer(..., hidden_states=latents)[0]

    latents = # compute latents from the prediction with Euler

With KV caching enabled, the first denoising step would go as usual, but we should expect to see a speedup in the subsequent steps. For multi-reference image editing, this can yield substantial speedups. This speedup, however, would come at the expense of increased memory consumption because of the on-memory storage. Depending on the image resolution being used, this increased memory consumption can pose challenges. Let’s discuss these numbers next.

This KV cache design is inspired by the PR that added support for Flux.2-Klein KV.

Speed-memory trade-off in KV caching

We took the Flux.2-klein-9b-kv pipeline and benchmarked it with and without KV caching. Unless otherwise specified, we always use a single reference image (Figure 4) with this prompt:

Dress this cat as a wizard wearing a blue hat and cloak, keeping its face and pose.

Figure 4: Reference cat image.

The table below demonstrates the impact of KV caching on speed and memory consumption of the pipeline:

Measurement ⬇️ Without cache With cache
Latency (4 steps) 3.085s 2.374s (23.1%)
Latency (8 steps) 5.883s 4.209s (28.4%)
Peak GPU memory 34.74GB 35.48GB

KV caching, as one would expect, does not affect quality:

Figure 5: Output images from Flux.2-klein-9b-kv with and without KV caching. Reproduce with the code from here.

Readers should take the qualitative comparisons with a grain of salt because these comparisons are presented with a single example, which is far from being conclusive. Instead, readers should run experiments on a set of prompts and determine metrics like PSNR and SSIM of the outputs to have a more robust comparison.

The speed and memory numbers scale with multiple reference images. With 3 reference images, we get between 42.1% and 51.3% speedups.

Measurement ⬇️ Without cache With cache
Latency (4 steps) 5.950s 3.446s (42.1%)
Latency (8 steps) 11.459s 5.580s (51.3%)
Peak GPU memory 35.45GB 39.58GB

The input prompt is a condition, just like the reference images. Since the prompt stays the same throughout denoising, could we cache its KV projections too?

The distinction is between the fixed text-encoder output and the text representations inside the denoiser. The pipeline computes the prompt embeddings once and reuses them at every denoising step. Inside the denoiser, however, those embeddings undergo further processing.

Flux.2 Klein has two kinds of transformer blocks:

  • The classic MMDiT blocks, which process text and image representations with separate parameters while allowing them to interact through joint attention.
  • More regular DiT blocks, which process a concatenated sequence of text and image tokens.

In both kinds of blocks, text representations depend on the current denoising timestep. Text tokens also attend to the changing noisy latent tokens, so their updated states and the KV projections in subsequent layers depend on the current noisy latent tokens.

Reference image tokens in the KV variant behave differently: they use a fixed timestep and attend only to reference tokens. Their representations still evolve across layers, but each layer’s reference KV remains unchanged across denoising steps.

Reusing text KV across steps would therefore be an approximation, rather than the equivalent computation provided by reference KV caching. It might still produce plausible images; let’s find out.

We ran a small experiment to also cache text KV similar to the reference image KV and to assess its impact on the final generated image.

Figure 6: Output images from Flux.2 Klein 9B KV with text KV caching enabled.

As seen in Figure 6, reusing text KV values produced coherent wizard-cat edits. The outputs changed in costume details and framing, but showed no obvious visual failure in this small experiment. Text KV caching is therefore an approximation worth evaluating. In terms of numbers:

Measurement ⬇️ Reference KV cache + Text KV cache
Latency (4 steps) 2.385s 2.231s (6.31%)
Latency (8 steps) 4.227s 3.861s (8.7%)
Peak GPU memory 35.48GB 35.73GB

A more recent (at the time of writing) model, QwenImage 2.1 [10], directly caches KV for both the text tokens and reference image tokens, discussed next.

QwenImage 2.1 results

Similar speed gains transfer to QwenImage 2.1:

Measurement ⬇️ Without cache With KV cache
Latency (40 steps) 32.202s 17.798s (44.7%)
Peak GPU memory 36.81GB 38.89GB
Figure 7: Output images for QwenImage 2.1 with and without KV caching.

QwenImage 2.1 uses a single stream of text and image tokens through its transformer blocks, unlike the design discussed so far. The broad design principles to accommodate this don’t change much; however, the internal details warrant a separate discussion.

KV caching in QwenImage 2.1

In QwenImage 2.1, reference image tokens occupy positions within the conditioning sequence, and the noisy latent tokens come last. A simplified sequence looks like this:

+-----------+-------------------+-----------+-------------------+---------------------+
|   Text A  | Reference image 1 |   Text B  | Reference image 2 | Noisy latent tokens |
+-----------+-------------------+-----------+-------------------+---------------------+
<--------------------- conditioning prefix ---------------------> <----- denoised ---->

Let’s call everything before the noisy latent tokens the conditioning prefix. Unlike the earlier reference-only caching example, the cache now contains KV projections for this entire prefix, separately for each transformer block. What makes it interesting is how these tokens attend to each other.

Groups of tokens and token influence

From our earlier discussions, we know that each reference image can be encoded into multiple latent tokens. QwenImage 2.1 assigns one image ID to all latent tokens from that reference, and a different ID to the noisy latent tokens being denoised. For the sequence above, the groups are:

Tokens Image ID
All latent tokens from reference image 1 0
All latent tokens from reference image 2 1
All noisy latent tokens being denoised 2
Text tokens, in either text segment -1 (no image ID)

An image block means one of these groups with a non-negative image ID. A token can attend to earlier tokens and itself. It can also attend to later tokens within its own image block. For example, two latent tokens from reference image 1 can attend to each other in both directions. Reference image 1 cannot attend to reference image 2, which comes later and has a different image ID. Reference image 2 can attend to reference image 1, which comes earlier. This causal feature is implemented using the following rule.

query_position is the index of the token doing the attending, and key_position is the index of the token being attended to, both in the assembled sequence.

same_image_block = (query_image_id >= 0) and (query_image_id == key_image_id)
allowed = (query_position >= key_position) or same_image_block

The noisy latent tokens all have ID 2, so they also attend to one another in both directions. This gives us three cases:

  • Text tokens attend to preceding tokens and themselves. Preceding tokens can include reference images.
  • Reference image tokens attend to all preceding text and reference image tokens, as well as every latent token encoded from that particular reference image.
  • Noisy latent tokens attend to the entire conditioning prefix and all noisy latent tokens, including themselves.

Figure 8 provides a visual way to think about this setup.

Figure 8: QwenImage 2.1 attention mask. Text attention is causal; each reference’s latent tokens can attend to one another in both directions. Noisy latent queries can attend to the entire sequence, while the conditioning prefix cannot attend to noisy latent tokens.

Reading the figure

Read the matrix by choosing a query row and a key column: teal means attention is allowed; light gray means it is blocked. Thin white lines separate individual tokens, and wider white gaps separate the labeled groups. Each group contains two tokens for illustration.

Let’s look at the two rows labeled Reference image 2. Their cells are teal under Text A, Reference image 1, Text B, and Reference image 2. Both tokens can therefore read all those tokens, including each other. Their cells under “Noisy latent tokens” are gray, so those connections are blocked. Reading in the opposite direction gives a different result: Reference image 1 cannot attend to the later “Reference image 2”.

Segmented attention

So, reference images can influence later text and later reference images. But no token in the conditioning prefix can attend to the noisy latent tokens. As these tokens change during denoising, they cannot change the prefix’s representations through attention.

So far, we have established a vague notion of different tokens belonging to different segments (denoted by “Image ID”s). Thus far, we have also discussed how the different tokens are represented within a single stream of tokens. How is attention computed in these streams?

We can compute this attention one segment at a time. A segment could be a text segment or an image segment. Here, a text segment is a consecutive run of text tokens: Text A and Text B in our example are separate segments because Reference image 1 sits between them. A reference image segment contains all the latent tokens encoded from one reference image. Each reference forms a separate segment, even when two references appear next to each other without text between them.

Each boxed group below is one segment; its label stands for all the tokens in that group:

+-----------+-------------------+-----------+-------------------+---------------------+
|   Text A  | Reference image 1 |   Text B  | Reference image 2 | Noisy latent tokens |
| segment 1 |     segment 2     | segment 3 |     segment 4     |      segment 5      |
+-----------+-------------------+-----------+-------------------+---------------------+
<--------------------- conditioning prefix ---------------------> <----- denoised ---->

Segments 1–4 form the conditioning prefix and are processed by the prefix_segments loop below. Segment 5 contains the noisy latent tokens, whose attention is computed after that loop.

For each segment, we take its queries and attend to the keys and values allowed by the causal rule discussed above. We also compute attention for the noisy latent queries, then concatenate all the outputs in sequence order.

# No batching, masking, output projections, and all that yet.
def segmented_attention(q, k, v, prefix_segments, prefix_len):
    outputs = []

    for start, end, is_text in prefix_segments:
        mask = None

        if is_text:
            query_positions = arange(start, end)
            key_positions = arange(end)
            mask = key_positions[None, :] <= query_positions[:, None]

        outputs.append(
            attend(q[start:end], k[:end], v[:end], mask=mask)
        )

    outputs.append(attend(q[prefix_len:], k, v))
    return concat(outputs, dim="sequence")

QKV projection preserves token order, so q[start:end] selects the queries belonging to that segment. With two tokens per group and indices starting at zero, Reference image 2 occupies positions 6 and 7. Its attention call is therefore attend(q[6:8], k[:8], v[:8]):

Figure 9: Attention for Reference image 2. Query rows 6 and 7 attend to key columns 0 through 7. Noisy latent columns 8 and 9 are excluded by the K/V slice.

Slicing KV at end excludes all later segments. For a text segment, an additional causal mask prevents each query from seeing later text positions within that segment. A reference image segment needs no such internal mask: it contains the latent tokens encoded from one reference, which can all attend to one another. Finally, queries from the noisy latent tokens attend to all keys and values.

These separate calls implement the block-causal attention rule exactly. QwenImage 2.1 also supports expressing the rule as a block mask for FlexAttention [11], which computes the attention in a single call.

Fixed timestep modulation

Like Flux.2, QwenImage 2.1 also uses fixed timesteps (zero in this case) for the text and reference image tokens, aka the conditioning prefix tokens.

Putting everything together

The attention mechanism and the fixed timestep modulation on conditioning prefix tokens make their computation independent of the denoising step. Its inputs stay fixed, its modulation doesn’t change, and it cannot receive information from the latents being denoised in the current step.

Despite the technical differences, the cache lifecycle should now look familiar. During the first denoising step, each transformer block processes the full sequence and saves the prefix KV. This is where the block-causal restriction applies: text and reference image queries must be prevented from reading later segments (Figure 8). The segmented computation above enforces it through KV slicing and text masks.

q, k, v = prepare_qkv(all_tokens)
layer_cache.store(k[:prefix_len].clone(), v[:prefix_len].clone())
output = segmented_attention(q, k, v, prefix_segments, prefix_len)

We cache only the prefix’s per-layer KV. To compute an attention output for a noisy latent token, we need its own query and the keys and values it can read. That computation does not use the queries of the text or reference image tokens.

During the first step, prefix queries help compute the prefix’s hidden states as they pass through the blocks, allowing us to populate each layer’s KV cache. At a given layer, those hidden states stay constant across denoising steps because of the fixed timestep modulation and block-causal attention. Once every layer has its prefix KV cached, later denoising steps skip computing prefix queries, attention outputs, and feed-forward outputs. The noisy latent tokens continue to read the cached prefix KV, so the conditioning still influences denoising.

On subsequent steps, each block therefore computes fresh QKV only for the noisy latent tokens and retrieves its saved prefix KV:

q_noisy, k_noisy, v_noisy = prepare_qkv(noisy_latent_tokens)
k_prefix, v_prefix = layer_cache.get()

noisy_latent_output = attend(
    q_noisy,
    concat(k_prefix, k_noisy, dim="sequence"),
    concat(v_prefix, v_noisy, dim="sequence"),
)

Looking back at Figure 8, we keep only the bottom two rows: the noisy latent queries. Every cell in those rows is teal because each query can read the entire prefix and all noisy latent tokens. There are therefore no block-causal restrictions to enforce on the remaining queries. With caching disabled, every denoising step recomputes the full sequence, so the block-causal restriction must be enforced at every step.

Room for optimization?

One problem with KV caching is that it increases memory consumption because we need to keep the cache in accelerator memory. Our cache is indexed with respect to the current layer being computed. Therefore, caches for other layers are not necessary. We could offload that cache to CPU (or even secondary storage) and onload when needed.

However, during these onloading and offloading phases, the accelerator must not sit idle, those phases should be as non-blocking as possible. To reduce this I/O overhead, we could overlap compute with communication (offloading and onloading the cache in this case).

It’d be a fun exercise to implement this in practice and see how far we can push the speed-memory trade-offs.

A specific class of caching techniques for flow-based pipelines exists, such as TaylorSeer [12], SeaCache [13], and TeaCache [14], which cache intermediate outputs within the denoiser. For example, TaylorSeer uses Taylor series expansions to approximate and cache intermediate activations across denoising steps. This family of caching techniques can be combined with KV caching, too.

Measurement ⬇️ Reference KV cache + TaylorSeer
Latency (4 steps) 2.407s 2.033s (15.5%)
Latency (8 steps) 4.264s 3.105s (27.2%)
Peak GPU memory 35.48GB 37.24GB

However, this combination can lead to noticeable degradation in output quality, as shown in Figure 10.

Figure 10: Combining KV caching and TaylorSeer caching may lead to quality degradations for Flux.2-klein-9b-kv.

Conclusion

KV caching has mostly been a technique tightly related to autoregressive models. It is promising to see it getting extended to other families of models and providing similar benefits.

Even in the compressed latent space, image and video generation can involve tens of thousands of visual tokens. For example, for 1024x1024 pixel images, Flux.2 Klein operates on a sequence of 4096 tokens. Doubling that resolution increases this to 16384 tokens. Caching these representations also has a substantial memory cost. In Flux.2 Klein 9B, each additional 1024 reference tokens requires approximately 0.5GB of BF16 KV storage across all layers. This adds up quickly as we scale up the number of reference images and the resolution. For videos, as one would imagine, these numbers would be even higher.

Therefore, it’d remain interesting to see how KV caching is leveraged without blowing memory consumption too much.

NoteAI assistance

The text content of this post comes primarily from the author. AI assistance (Codex) was used for the following: language polishing, code snippets, and coordination of the experiments.

References

1.
2.
Esser P, Kulal S, Blattmann A, et al (2024) Scaling rectified flow transformers for high-resolution image synthesis. In: Salakhutdinov R, Kolter Z, Heller K, et al (eds) Proceedings of the 41st international conference on machine learning. PMLR, pp 12606–12633
3.
Lipman Y, Chen RTQ, Ben-Hamu H, Nickel M, Le M (2023) Flow matching for generative modeling. In: The eleventh international conference on learning representations
4.
Su J, Ahmed M, Lu Y, Pan S, Bo W, Liu Y (2024) RoFormer: Enhanced transformer with rotary position embedding. Neurocomput 568(C). https://doi.org/10.1016/j.neucom.2023.127063
5.
Dehghani M, Djolonga J, Mustafa B, et al (2023) Scaling vision transformers to 22 billion parameters
6.
Oquab M, Darcet T, Moutakanni T, et al (2024) DINOv2: Learning robust visual features without supervision. Transactions on Machine Learning Research
7.
Zhai X, Mustafa B, Kolesnikov A, Beyer L (2023) Sigmoid loss for language image pre-training. In: 2023 IEEE/CVF international conference on computer vision (ICCV). pp 11941–11952
8.
Black Forest Labs (2025) FLUX.2: Frontier visual intelligence
9.
Wu C, Li J, Zhou J, et al (2025) Qwen-Image technical report
10.
11.
12.
Liu J, Zou C, Lyu Y, Chen J, Zhang L (2025) From reusing to forecasting: Accelerating diffusion models with TaylorSeers. In: Proceedings of the IEEE/CVF international conference on computer vision (ICCV). pp 15853–15863
13.
Chung J, Hyun S, Lee M, et al (2026) SeaCache: Spectral-evolution-aware cache for accelerating diffusion models. In: Proceedings of the IEEE/CVF conference on computer vision and pattern recognition (CVPR). pp 14283–14294
14.
Liu F, Zhang S, Wang X, et al (2025) Timestep embedding tells: It’s time to cache for video diffusion model. In: Proceedings of the IEEE/CVF conference on computer vision and pattern recognition (CVPR). pp 7353–7363