From 1d34a600e21e3b5fad4d511386984731a0dde22a Mon Sep 17 00:00:00 2001
From: Mike Lothian <mike@fireburn.co.uk>
Date: Wed, 16 Sep 2026 15:38:10 +0100
Subject: [PATCH 02/18] vulkan : share the K/V pass between query rows in
 coopmat1 GQA flash attention

Assisted-by: Claude Code (Claude Opus 5.5)
---
 ggml/src/ggml-vulkan/ggml-vulkan.cpp          | 31 ++++++++--
 .../vulkan-shaders/flash_attn_base.glsl       | 33 ++++++++++-
 .../vulkan-shaders/flash_attn_cm1.comp        | 59 +++++++++++++++----
 3 files changed, 103 insertions(+), 20 deletions(-)

diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
index 414e14f85..f0c99fa79 100644
--- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp
+++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
@@ -1371,7 +1371,8 @@ vk_fa_tuning_params get_fa_tuning_params(const vk_device& device, uint32_t hsk,
 }
 
 vk_fa_pipeline_state get_fa_pipeline_state(const vk_device& device, const vk_fa_tuning_params& params, uint32_t hsk, uint32_t hsv, bool aligned, bool f32acc,
-                                                  bool use_mask, bool use_mask_opt, bool use_logit_softcap, bool use_sparse, ggml_type k_type, ggml_type v_type) {
+                                                  bool use_mask, bool use_mask_opt, bool use_logit_softcap, bool use_sparse, ggml_type k_type, ggml_type v_type,
+                                                  bool use_mq = false) {
     const bool old_amd_windows = device->vendor_id == VK_VENDOR_ID_AMD && device->driver_id == vk::DriverId::eAmdProprietary &&
                                  (device->architecture == AMD_GCN || device->architecture == AMD_RDNA1 || device->architecture == AMD_RDNA2);
 
@@ -1379,7 +1380,8 @@ vk_fa_pipeline_state get_fa_pipeline_state(const vk_device& device, const vk_fa_
                      (use_mask          ? 2 : 0) |
                      (use_logit_softcap ? 4 : 0) |
                      (old_amd_windows   ? 8 : 0) |
-                     (use_sparse        ? 16 : 0);
+                     (use_sparse        ? 16 : 0) |
+                     (use_mq            ? 32 : 0);
 
     const uint32_t subgroup_size = params.disable_subgroups ? 0 : params.subgroup_size;
 
@@ -8086,6 +8088,19 @@ void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx, const
 
     tuning_params = get_fa_tuning_params(ctx->device, HSK, HSV, N, KV, k_type_eff, v_type_eff, f32acc);
 
+    // Multi-query GQA: fold a small batch of query rows into the GQA rows so they share one K/V pass.
+    uint32_t gqa_nq = 1;
+    {
+        static const bool disable_mq = getenv("GGML_VK_FA_DISABLE_MQ") != nullptr;
+        const float mq_max_bias = ((const float *) dst->op_params)[1];
+        if (!disable_mq && gqa_ratio > 1 && neq1 > 1 && neq1 <= 8 && neq3 == 1 && nem3 <= 1 &&
+            tuning_params.path == FA_COOPMAT1 && sinks == nullptr && mq_max_bias == 0.0f && mask != nullptr) {
+            gqa_nq = (uint32_t)neq1;
+            N = gqa_nq * gqa_ratio;
+            tuning_params = get_fa_tuning_params(ctx->device, HSK, HSV, N, KV, k_type_eff, v_type_eff, f32acc);
+        }
+    }
+
     float scale         = 1.0f;
     float max_bias      = 0.0f;
     float logit_softcap = 0.0f;
@@ -8103,7 +8118,7 @@ void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx, const
     static const bool disable_sparse = getenv("GGML_VK_FA_SPARSE_DISABLE") != nullptr;
     // cm2 dense is fast, so it needs a larger reduction to win.
     const int64_t min_ratio = tuning_params.path == FA_COOPMAT2 ? 4 : 2;
-    const bool use_sparse = !disable_sparse && n_kv_max > 0 && mask &&
+    const bool use_sparse = !disable_sparse && gqa_nq == 1 && n_kv_max > 0 && mask &&
                             max_bias == 0.0f && logit_softcap == 0.0f &&
                             k_type_eff == GGML_TYPE_F16 && v_type_eff == GGML_TYPE_F16 &&
                             nem0 == KV &&
@@ -8143,10 +8158,11 @@ void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx, const
     }
 
     // Only use mask opt when the mask is fairly large. This hasn't been tuned extensively.
-    bool use_mask_opt = mask && !use_sparse && nem1 >= 32 && nem0 * nem1 > 32768 && nem0 >= tuning_params.block_cols * 16
+    bool use_mask_opt = mask && !use_sparse && gqa_nq == 1 && nem1 >= 32 && nem0 * nem1 > 32768 && nem0 >= tuning_params.block_cols * 16
                         && (ctx->device->architecture != vk_device_architecture::AMD_GCN || HSK > 256 || HSV > 256);
     vk_fa_pipeline_state fa_pipeline_state = get_fa_pipeline_state(ctx->device, tuning_params, HSK, HSV, aligned, f32acc,
-                                                                   mask != nullptr, use_mask_opt, logit_softcap != 0, use_sparse, k_type_eff, v_type_eff);
+                                                                   mask != nullptr, use_mask_opt, logit_softcap != 0, use_sparse, k_type_eff, v_type_eff,
+                                                                   gqa_nq > 1);
 
     vk_pipeline pipeline = nullptr;
 
@@ -8191,6 +8207,10 @@ void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx, const
 
     GGML_ASSERT(Br == pipeline->wg_denoms[0]);
     const uint32_t Tr = CEIL_DIV(N, Br);
+    if (gqa_nq > 1) {
+        // one workgroup per tile of Br folded rows
+        workgroups_x = Tr;
+    }
 
     // Try to use split_k when KV is large enough to be worth the overhead.
     // Sparse: split_kv carries n_kv_max, split_k partitions its blocks for occupancy.
@@ -16483,4 +16503,3 @@ void ggml_vk_debug_label::begin(vk_context & ctx, const std::string & name) {
     subctx->debug_labels.push_back(name);
     ggml_vk_cmd_label_begin(subctx->s->buffer->buf, subctx->debug_labels.back().c_str());
 }
-
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl
index 2e0e23bc1..732ddc4eb 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl
@@ -26,6 +26,9 @@ const bool LOGIT_SOFTCAP   = (Flags & 4) != 0;
 const bool OLD_AMD_WINDOWS = (Flags & 8) != 0;
 // Sparse: gather binding-7 indices instead of scanning [0,KV); p.split_kv = n_kv_max.
 const bool USE_SPARSE      = (Flags & 16) != 0;
+// Multi-query GQA: several query rows fold into the GQA rows of one tile and share one K/V pass.
+// Row r maps to query row r / gqa_ratio and head r % gqa_ratio.
+const bool USE_MQ          = (Flags & 32) != 0;
 
 // Round up head sizes to a multiple of 16, for coopmat1/coopmat2 paths
 const uint32_t HSK_pad = (HSK + 15) & ~15;
@@ -155,7 +158,12 @@ void init_indices()
     N = p.N;
     KV = p.KV;
 
-    if (p.k_num > 1) {
+    if (USE_MQ) {
+        // tiles of Br rows share gl_WorkGroupID.x with split_k
+        gqa_iq1 = 0;
+        split_k_index = gl_WorkGroupID.x % p.k_num;
+        i = gl_WorkGroupID.x / p.k_num;
+    } else if (p.k_num > 1) {
         if (p.gqa_ratio > 1) {
             i = 0;
             // batch and split_k share gl_WorkGroupID.x
@@ -226,6 +234,29 @@ void init_indices()
     }
 }
 
+// Element offset of a Q row relative to q_offset.
+uint32_t fa_q_elem(const in uint32_t row) {
+    if (USE_MQ) {
+        return (row / p.gqa_ratio) * p.nb01 + (row % p.gqa_ratio) * q_stride;
+    }
+    return row * q_stride;
+}
+
+// Mask offset of a Q row relative to m_offset.
+uint32_t fa_m_row(const in uint32_t row) {
+    if (USE_MQ) {
+        return (row / p.gqa_ratio) * KV;
+    }
+    return row * m_stride;
+}
+
+// Store an output row in multi-query gqa mode. o_base is the offset of query row 0 for this split/iq3.
+void mqStore(const in uint32_t row, const in uint32_t c, const in O_TYPEV4 elems, const in uint32_t o_base, const in uint32_t o_row_stride)
+{
+    uint32_t offset = o_base + (row / p.gqa_ratio) * o_row_stride + ((iq2 + row % p.gqa_ratio) * HSV) / 4 + c;
+    data_ov4[offset] = D_TYPEV4(elems);
+}
+
 // Resolve a linear KV slot to a real column; false for inactive (sparse padding/-1, or dense OOB).
 bool fa_kv_index(uint lin, out uint kv_col) {
     if (USE_SPARSE) {
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 aa9dd624b..4e1259522 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp
@@ -97,7 +97,7 @@ void main() {
         uint32_t r = (idx + tid) / (HSK / 4);
         if (r < Br && d < HSK / 4 &&
             i * Br + r < N) {
-            Qf[r * qstride + d] = FLOAT_TYPEV4(data_qv4[q_offset / 4 + (i * Br + r) * q_stride / 4 + d] * p.scale);
+            Qf[r * qstride + d] = FLOAT_TYPEV4(data_qv4[q_offset / 4 + fa_q_elem(i * Br + r) / 4 + d] * p.scale);
         }
     }
     barrier();
@@ -186,25 +186,25 @@ void main() {
                                 m = f16vec4(mv);
                                 max_mask = max(max_mask, float(mv));
                             } else if (!nem1_bounds_check || i * Br + r * 4 + 3 < p.nem1) {
-                                m = f16vec4(data_m[m_offset + (i * Br + r * 4    ) * m_stride + (j * Bc + c)],
-                                            data_m[m_offset + (i * Br + r * 4 + 1) * m_stride + (j * Bc + c)],
-                                            data_m[m_offset + (i * Br + r * 4 + 2) * m_stride + (j * Bc + c)],
-                                            data_m[m_offset + (i * Br + r * 4 + 3) * m_stride + (j * Bc + c)]);
+                                m = f16vec4(data_m[m_offset + fa_m_row(i * Br + r * 4) + (j * Bc + c)],
+                                            data_m[m_offset + fa_m_row(i * Br + r * 4 + 1) + (j * Bc + c)],
+                                            data_m[m_offset + fa_m_row(i * Br + r * 4 + 2) + (j * Bc + c)],
+                                            data_m[m_offset + fa_m_row(i * Br + r * 4 + 3) + (j * Bc + c)]);
                                 max_mask = max(max(max(max(max_mask, float(m[0])), float(m[1])), float(m[2])), float(m[3]));
                             } else if (i * Br + r * 4 + 2 < p.nem1) {
-                                m = f16vec4(data_m[m_offset + (i * Br + r * 4    ) * m_stride + (j * Bc + c)],
-                                            data_m[m_offset + (i * Br + r * 4 + 1) * m_stride + (j * Bc + c)],
-                                            data_m[m_offset + (i * Br + r * 4 + 2) * m_stride + (j * Bc + c)],
+                                m = f16vec4(data_m[m_offset + fa_m_row(i * Br + r * 4) + (j * Bc + c)],
+                                            data_m[m_offset + fa_m_row(i * Br + r * 4 + 1) + (j * Bc + c)],
+                                            data_m[m_offset + fa_m_row(i * Br + r * 4 + 2) + (j * Bc + c)],
                                             0.0);
                                 max_mask = max(max(max(max_mask, float(m[0])), float(m[1])), float(m[2]));
                             } else if (i * Br + r * 4 + 1 < p.nem1) {
-                                m = f16vec4(data_m[m_offset + (i * Br + r * 4    ) * m_stride + (j * Bc + c)],
-                                            data_m[m_offset + (i * Br + r * 4 + 1) * m_stride + (j * Bc + c)],
+                                m = f16vec4(data_m[m_offset + fa_m_row(i * Br + r * 4) + (j * Bc + c)],
+                                            data_m[m_offset + fa_m_row(i * Br + r * 4 + 1) + (j * Bc + c)],
                                             0.0,
                                             0.0);
                                 max_mask = max(max(max_mask, float(m[0])), float(m[1]));
                             } else if (i * Br + r * 4 < p.nem1) {
-                                m = f16vec4(data_m[m_offset + (i * Br + r * 4    ) * m_stride + (j * Bc + c)],
+                                m = f16vec4(data_m[m_offset + fa_m_row(i * Br + r * 4) + (j * Bc + c)],
                                             0.0,
                                             0.0,
                                             0.0);
@@ -550,7 +550,28 @@ void main() {
     // If there is split_k, then the split_k resolve shader does the final
     // division by L. Store the intermediate O value and per-row m and L values.
     if (p.k_num > 1) {
-        if (p.gqa_ratio > 1) {
+        if (USE_MQ) {
+            const uint32_t o_base = HSV * p.ne1 * (split_k_index + p.k_num * p.ne2 * iq3) / 4;
+            const uint32_t o_row_stride = HSV * p.ne1 * p.k_num / 4;
+            const uint32_t lm_base = HSV * p.ne1 * p.k_num * p.ne2 * p.ne3 + p.ne1 * 2 * (split_k_index + p.k_num * p.ne2 * iq3);
+
+            [[unroll]] for (uint32_t r = 0; r < rows_per_thread; ++r) {
+                const uint row = i * Br + tile_row(r);
+                if (row < N) {
+                    [[unroll]] for (uint32_t d0 = 0; d0 < HSV / 4; d0 += threads_per_rowgroup) {
+                        const uint d = d0 + col_tid;
+                        if (d >= HSV/4) break;
+                        const uint d_local = d0 / threads_per_rowgroup;
+                        mqStore(row, d, Of[r][d_local], o_base, o_row_stride);
+                    }
+                    if (col_tid == 0) {
+                        const uint lm_offset = lm_base + p.ne1 * 2 * p.k_num * (row / p.gqa_ratio) + iq2 + row % p.gqa_ratio;
+                        data_o[lm_offset] = D_TYPE(Lf[r]);
+                        data_o[lm_offset + p.ne1] = D_TYPE(Mf[r]);
+                    }
+                }
+            }
+        } else if (p.gqa_ratio > 1) {
             // note: O and Q have swapped coord 1,2.
             uint32_t o_offset = HSV * p.ne1 * (split_k_index + p.k_num * (gqa_iq1 + p.ne2 * iq3)) / 4;
 
@@ -637,7 +658,19 @@ void main() {
 
     uint32_t o_offset = (gqa_iq1*p.ne1*HSV + iq3*p.ne2*p.ne1*HSV) / 4;
 
-    if (p.gqa_ratio > 1) {
+    if (USE_MQ) {
+        [[unroll]] for (uint32_t r = 0; r < rows_per_thread; ++r) {
+            const uint row = i * Br + tile_row(r);
+            if (row < N) {
+                [[unroll]] for (uint32_t d0 = 0; d0 < HSV / 4; d0 += threads_per_rowgroup) {
+                    const uint d = d0 + col_tid;
+                    if (d >= HSV / 4) break;
+                    const uint d_local = d0 / threads_per_rowgroup;
+                    mqStore(row, d, Of[r][d_local], o_offset, p.ne1 * HSV / 4);
+                }
+            }
+        }
+    } else if (p.gqa_ratio > 1) {
         [[unroll]] for (uint32_t r = 0; r < rows_per_thread; ++r) {
             if (tile_row(r) < N) {
                 [[unroll]] for (uint32_t d0 = 0; d0 < HSV / 4; d0 += threads_per_rowgroup) {
-- 
2.55.0

