Skip to content

Qwen Image 2.1 prefix cache: store K/V as F16 when flash attention is on (half the VRAM, same output) #2040

Description

@brt16

Thanks for implementing the prefix cache so quickly in #2035 (from #2031). It works great.

One suggestion: the cache always stores K/V as FP32 (qwen_image_2_1.hpp#L194). But when flash attention is enabled, ggml_ext_attention_ext converts K and V to F16 right before the kernel (ggml_extend.cpp#L675 and #L685). So with flash attention the stored FP32 values are rounded to F16 on every step anyway. Storing them as F16 in the first place would give identical results with half the memory.

The same holds only when:

  • flash attention is on
  • sage attention is off (its path uses F32 K)
  • kv_scale == 1, since the scale is applied before the cast (#L673)

In every other case FP32 is still needed.

Why it matters: the docs estimate about 4 GiB per condition for a 4096-token prefix. Multi-reference edits reach that quickly. Two 1024x1024 references are over 8192 tokens, which is about 8+ GiB at FP32 versus 4+ GiB at F16. On a 16 GB card that difference decides whether the cache fits next to the weights.

The patch attached to #2031 did this (prefix_cache_type()). On Vulkan (RX 6800 XT, --diffusion-fa), its output was pixel-identical to the uncached path. I haven't tested other backends.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions