From 92be0be23ee89db29bd2c35db77bd5d4e5ec6df2 Mon Sep 17 00:00:00 2001
From: Mike Lothian <mike@fireburn.co.uk>
Date: Sun, 27 Sep 2026 13:29:32 +0100
Subject: [PATCH 13/18] vulkan : keep Q fragments in registers in coopmat1
 flash attention

Assisted-by: Claude Code (Claude Opus 5.5)
---
 .../vulkan-shaders/flash_attn_cm1.comp           | 16 +++++++++++++++-
 1 file changed, 15 insertions(+), 1 deletion(-)

diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp
index 4e1259522..11feb6530 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp
@@ -150,6 +150,16 @@ void main() {
     uint32_t mask_opt_bits = 0;
     f16vec4 mask_cache[Bc * Br / 4 / WorkGroupSize];
 
+    // Q does not change over the KV loop: keep its B fragments in registers (head sizes up to 256)
+    const bool Q_IN_REGS = HSK_pad <= 256;
+    coopmat<FLOAT_TYPE, gl_ScopeSubgroup, 16, MatBr, gl_MatrixUseB> QRegs[Q_IN_REGS ? HSK_pad / 16 : 1];
+    if (Q_IN_REGS) {
+        barrier();
+        [[unroll]] for (uint32_t d = 0; d < HSK_pad / 16; ++d) {
+            coopMatLoad(QRegs[d], Qf, d * 16 / 4, qstride, gl_CooperativeMatrixLayoutColumnMajor);
+        }
+    }
+
     [[dont_unroll]]
     for (uint32_t j = start_j; j < end_j; ++j) {
 
@@ -319,7 +329,11 @@ void main() {
                 coopMatLoad(KMat, kvsh, coord, kvsh_stride, gl_CooperativeMatrixLayoutRowMajor);
             }
 
-            coopMatLoad(QMat, Qf, d * 16 / 4, qstride, gl_CooperativeMatrixLayoutColumnMajor);
+            if (Q_IN_REGS) {
+                QMat = QRegs[d];
+            } else {
+                coopMatLoad(QMat, Qf, d * 16 / 4, qstride, gl_CooperativeMatrixLayoutColumnMajor);
+            }
 
             SfMat = coopMatMulAdd(KMat, QMat, SfMat);
         }
-- 
2.55.0

