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.
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_extconverts 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:
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.