From c0ed46d0b2f3c6c45989727435b1d1d221287e7d 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 14/18] vulkan : store the flash attention quantized K/V
 scratch as 16x16 tiles

Assisted-by: Claude Code (Claude Opus 5.5)
---
 ggml/src/ggml-vulkan/ggml-vulkan-types.h      |  2 ++
 ggml/src/ggml-vulkan/ggml-vulkan.cpp          | 19 +++++++++++----
 .../vulkan-shaders/dequant_q8_0.comp          | 22 +++++++++++++++++-
 .../vulkan-shaders/flash_attn_base.glsl       |  2 ++
 .../vulkan-shaders/flash_attn_cm1.comp        | 12 ++++++++--
 .../vulkan-shaders/vulkan-shaders-gen.cpp     |  2 ++
 tests/test-backend-ops.cpp                    | 23 +++++++++++++++++--
 7 files changed, 72 insertions(+), 10 deletions(-)

diff --git a/ggml/src/ggml-vulkan/ggml-vulkan-types.h b/ggml/src/ggml-vulkan/ggml-vulkan-types.h
index 21cdd044e..de6b06390 100644
--- a/ggml/src/ggml-vulkan/ggml-vulkan-types.h
+++ b/ggml/src/ggml-vulkan/ggml-vulkan-types.h
@@ -800,6 +800,8 @@ struct vk_device_struct {
 
     vk_pipeline pipeline_dequant[GGML_TYPE_COUNT];
     vk_pipeline pipeline_dequant_transpose[GGML_TYPE_COUNT]; // fused dequant+transpose for FA quant-KV
+    vk_pipeline pipeline_dequant_ktile[GGML_TYPE_COUNT];     // dequant to per-head 16x16 K tiles for FA cm1
+    vk_pipeline pipeline_dequant_vtile[GGML_TYPE_COUNT];     // dequant to per-head 16x16 transposed V tiles for FA cm1
     vk_pipeline pipeline_dequant_mul_mat_vec_f32_f32[DMMV_WG_SIZE_COUNT][GGML_TYPE_COUNT][mul_mat_vec_max_cols];
     vk_pipeline pipeline_dequant_mul_mat_vec_f16_f32[DMMV_WG_SIZE_COUNT][GGML_TYPE_COUNT][mul_mat_vec_max_cols];
     vk_pipeline pipeline_dequant_mul_mat_vec_id_f32[DMMV_WG_SIZE_COUNT][GGML_TYPE_COUNT];
diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
index 13b7f3268..f1cd6a636 100644
--- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp
+++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
@@ -1372,7 +1372,7 @@ 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_mq = false) {
+                                                  bool use_mq = false, bool use_tiled_kv = 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);
 
@@ -1381,7 +1381,8 @@ vk_fa_pipeline_state get_fa_pipeline_state(const vk_device& device, const vk_fa_
                      (use_logit_softcap ? 4 : 0) |
                      (old_amd_windows   ? 8 : 0) |
                      (use_sparse        ? 16 : 0) |
-                     (use_mq            ? 32 : 0);
+                     (use_mq            ? 32 : 0) |
+                     (use_tiled_kv      ? 64 : 0);
 
     const uint32_t subgroup_size = params.disable_subgroups ? 0 : params.subgroup_size;
 
@@ -3092,6 +3093,8 @@ void ggml_vk_load_shaders(vk_device& device, vk_pipeline requested) {
     ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q5_1], "dequant_q5_1", dequant_q5_1_len, dequant_q5_1_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1);
     ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q8_0], "dequant_q8_0", dequant_q8_0_len, dequant_q8_0_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1);
     ggml_vk_create_pipeline(device, device->pipeline_dequant_transpose[GGML_TYPE_Q8_0], "dequant_q8_0_transpose", dequant_q8_0_transpose_len, dequant_q8_0_transpose_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1);
+    ggml_vk_create_pipeline(device, device->pipeline_dequant_ktile[GGML_TYPE_Q8_0], "dequant_q8_0_ktile", dequant_q8_0_ktile_len, dequant_q8_0_ktile_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1);
+    ggml_vk_create_pipeline(device, device->pipeline_dequant_vtile[GGML_TYPE_Q8_0], "dequant_q8_0_vtile", dequant_q8_0_vtile_len, dequant_q8_0_vtile_data, "main", 2, 5 * sizeof(uint32_t), {256 * 16, 1, 1}, {}, 1);
     ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q2_K], "dequant_q2_k", dequant_q2_k_len, dequant_q2_k_data, "main", 2, 5 * sizeof(uint32_t), {256 * 64, 1, 1}, {}, 1);
     ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q3_K], "dequant_q3_k", dequant_q3_k_len, dequant_q3_k_data, "main", 2, 5 * sizeof(uint32_t), {256 * 64, 1, 1}, {}, 1);
     ggml_vk_create_pipeline(device, device->pipeline_dequant[GGML_TYPE_Q4_K], "dequant_q4_k", dequant_q4_k_len, dequant_q4_k_data, "main", 2, 5 * sizeof(uint32_t), {256 * 32, 1, 1}, {}, 1);
@@ -8169,9 +8172,15 @@ 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 && 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);
+    // coopmat1 reads the f16 scratch as 16x16 tiles, one contiguous fragment per load
+    const bool use_tiled_kv = use_dequant_kv && !use_sparse && aligned && gqa_nq == 1 &&
+                              tuning_params.path == FA_COOPMAT1 && !tuning_params.shmem_staging &&
+                              HSK % 16 == 0 && HSV % 16 == 0 && KV % 16 == 0 &&
+                              ctx->device->pipeline_dequant_ktile[k->type] != nullptr &&
+                              ctx->device->pipeline_dequant_vtile[v->type] != nullptr;
     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,
-                                                                   gqa_nq > 1);
+                                                                   gqa_nq > 1, use_tiled_kv);
 
     vk_pipeline pipeline = nullptr;
 
@@ -8367,8 +8376,8 @@ void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx, const
             ctx->prealloc_size_x = k_f16_sz + v_f16_sz;
             ggml_vk_preallocate_buffers(ctx, subctx);
         }
-        vk_pipeline tr_k = ctx->device->pipeline_dequant_transpose[k->type];
-        vk_pipeline tr_v = ctx->device->pipeline_dequant_transpose[v->type];
+        vk_pipeline tr_k = use_tiled_kv ? ctx->device->pipeline_dequant_ktile[k->type] : ctx->device->pipeline_dequant_transpose[k->type];
+        vk_pipeline tr_v = use_tiled_kv ? ctx->device->pipeline_dequant_vtile[v->type] : ctx->device->pipeline_dequant_transpose[v->type];
         ggml_pipeline_request_descriptor_sets(ctx, tr_k, 1);
         ggml_pipeline_request_descriptor_sets(ctx, tr_v, 1);
         if (ctx->prealloc_x_need_sync) {
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q8_0.comp b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q8_0.comp
index 3b3fbbe89..764d0c63f 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q8_0.comp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/dequant_q8_0.comp
@@ -18,7 +18,21 @@ void main() {
         return;
     }
 
-#ifdef DEQUANT_TRANSPOSE
+#if defined(DEQUANT_KTILE) || defined(DEQUANT_VTILE)
+    // read [HS, NH, KV, NS], write per head 16x16 tiles (kv/16, h/16): K row-major, V transposed
+    const uint HS = p.M, NH = p.K, KVn = p.stride_a;
+    const uint e0 = ib * 32 + 16 * il;
+    const uint kv = (e0 / (HS * NH)) % KVn;
+    const uint head_base = ((e0 / (HS * NH * KVn)) * NH + (e0 / HS) % NH) * (HS * KVn);
+    const uint tile = (kv / 16) * (HS / 16) + (e0 % HS) / 16;
+#if defined(DEQUANT_KTILE)
+    const uint b_idx = head_base + tile * 256 + (kv % 16) * 16;
+    const uint b_step = 1;
+#else
+    const uint b_idx = head_base + tile * 256 + (kv % 16);
+    const uint b_step = 16;
+#endif
+#elif defined(DEQUANT_TRANSPOSE)
     // read [HS, NH, KV, NS], write [HS, KV, NH, NS]
     const uint HS = p.M, NH = p.K, KVn = p.stride_a;
     const uint e0 = ib * 32;
@@ -35,8 +49,14 @@ void main() {
 
     const uint q_idx = 16*il;
 
+#if defined(DEQUANT_KTILE) || defined(DEQUANT_VTILE)
+    [[unroll]] for (uint l = 0; l < 16; ++l) {
+        data_b[b_idx + l * b_step] = D_TYPE(d * data_a[ib].qs[q_idx + l]);
+    }
+#else
     [[unroll]] for (uint l = 0; l < 16; l += 2) {
         data_b[b_idx + l    ] = D_TYPE(d * data_a[ib].qs[q_idx + l    ]);
         data_b[b_idx + l + 1] = D_TYPE(d * data_a[ib].qs[q_idx + l + 1]);
     }
+#endif
 }
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 732ddc4eb..5275a3d1c 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_base.glsl
@@ -29,6 +29,8 @@ 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;
+// K/V are f16 scratch in 16x16 tiles (K row-major, V transposed): each fragment is one contiguous 512 byte load
+const bool TILED_KV        = (Flags & 64) != 0;
 
 // Round up head sizes to a multiple of 16, for coopmat1/coopmat2 paths
 const uint32_t HSK_pad = (HSK + 15) & ~15;
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 11feb6530..6df261711 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_cm1.comp
@@ -320,6 +320,9 @@ void main() {
             if (stage_k) {
                 uint coord = (gl_SubgroupID * MatBc) * kvsh_stride;
                 coopMatLoad(KMat, kvsh, coord, kvsh_stride, gl_CooperativeMatrixLayoutRowMajor);
+            } else if (TILED_KV) {
+                const uint kb = (j * Bc + gl_SubgroupID * MatBc) / 16;
+                coopMatLoad(KMat, data_kv4, k_offset / 4 + (kb * (HSK / 16) + d) * 64, 4, gl_CooperativeMatrixLayoutRowMajor);
             } else {
                 const uint coord = k_offset / 4 + (j * Bc + gl_SubgroupID * MatBc) * k_stride / 4 + d * 16 / 4;
                 coopMatLoad(KMat, data_kv4, coord, k_stride / 4, gl_CooperativeMatrixLayoutRowMajor);
@@ -512,8 +515,13 @@ void main() {
                     if (!USE_DECODE_V && !KV_bounds_check && !USE_SPARSE) {
                         // F16/BF16 values can be loaded directly from global memory
                         const uint v_tile_row = j * Bc + bc_chunk * MatBc;
-                        const uint v_tile_offset = v_offset / 4 + v_tile_row * v_stride / 4 + hsv_offset / 4;
-                        coopMatLoad(QMat, data_vv4, v_tile_offset, v_stride / 4, gl_CooperativeMatrixLayoutRowMajor);
+                        if (TILED_KV) {
+                            const uint vt = v_offset / 4 + ((v_tile_row / 16) * (HSV / 16) + hsv_offset / 16) * 64;
+                            coopMatLoad(QMat, data_vv4, vt, 4, gl_CooperativeMatrixLayoutColumnMajor);
+                        } else {
+                            const uint v_tile_offset = v_offset / 4 + v_tile_row * v_stride / 4 + hsv_offset / 4;
+                            coopMatLoad(QMat, data_vv4, v_tile_offset, v_stride / 4, gl_CooperativeMatrixLayoutRowMajor);
+                        }
                     } else {
                         const uint v_tile_offset = bc_chunk * MatBr * v_cols + gl_SubgroupID * (MatBc / 4);
                         coopMatLoad(QMat, kvsh, v_tile_offset, vsh_stride, gl_CooperativeMatrixLayoutRowMajor);
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp
index 115534893..eb14e065c 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/vulkan-shaders-gen.cpp
@@ -831,6 +831,8 @@ void process_shaders() {
         // Fused dequant+transpose variant for FA quant-KV (per-head-contiguous f16 scratch).
         if (tname == "q8_0") {
             string_to_spv("dequant_" + tname + "_transpose", "dequant_" + tname + ".comp", merge_maps(base_dict, {{data_a_key, "1"}, {"D_TYPE", "float16_t"}, {"DEQUANT_TRANSPOSE", "1"}}));
+            string_to_spv("dequant_" + tname + "_ktile", "dequant_" + tname + ".comp", merge_maps(base_dict, {{data_a_key, "1"}, {"D_TYPE", "float16_t"}, {"DEQUANT_TRANSPOSE", "1"}, {"DEQUANT_KTILE", "1"}}));
+            string_to_spv("dequant_" + tname + "_vtile", "dequant_" + tname + ".comp", merge_maps(base_dict, {{data_a_key, "1"}, {"D_TYPE", "float16_t"}, {"DEQUANT_TRANSPOSE", "1"}, {"DEQUANT_VTILE", "1"}}));
         }
 
         shader = (tname == "f32" || tname == "f16" || tname == "bf16") ? "get_rows.comp" : "get_rows_quant.comp";
diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp
index 7bb3fca89..7ab411442 100644
--- a/tests/test-backend-ops.cpp
+++ b/tests/test-backend-ops.cpp
@@ -10981,6 +10981,17 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
         }
     }
 
+    // quantized K/V prefill (large batch) with GQA, e.g. 24 Q / 4 KV heads
+    for (int hs : { 128, 256, }) {
+        for (int nb : { 64, 100, 512, }) {
+            for (int kv : { 512, 4096, }) {
+                for (bool mask : { true, false }) {
+                    test_cases.emplace_back(new test_flash_attn_ext(hs, hs, 4, {6, 1}, kv, nb, mask, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, {0, 2, 1, 3}, false));
+                }
+            }
+        }
+    }
+
     // asymmetric head_dim (hsk != hsv) with one or both sides not 64-aligned
     test_cases.emplace_back(new test_flash_attn_ext(72, 64, 4, {1, 1}, 256, 2, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16));
     test_cases.emplace_back(new test_flash_attn_ext(64, 72, 4, {1, 1}, 256, 2, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16));
@@ -11003,9 +11014,9 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
     // q8_0 KV cases: decode and prompt batches, KV pad, permuted KV, feature flags, and long context
     test_cases.emplace_back(new test_flash_attn_ext(256, 256, 2, {16, 1},   113,   1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0));
     test_cases.emplace_back(new test_flash_attn_ext(256, 256, 2, {16, 1},  1024,   1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0));
-    test_cases.emplace_back(new test_flash_attn_ext(256, 256, 2, {16, 1},  1024,   1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, {0, 2, 1, 3}));
+    test_cases.emplace_back(new test_flash_attn_ext(256, 256, 2, {16, 1},  1024,   1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, {0, 2, 1, 3}, false));
     test_cases.emplace_back(new test_flash_attn_ext(256, 256, 2, {16, 2},  1025,   1, true, true,  8, 30, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0));
-    test_cases.emplace_back(new test_flash_attn_ext(256, 256, 2, {16, 1},  1025,  64, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, {0, 2, 1, 3}));
+    test_cases.emplace_back(new test_flash_attn_ext(256, 256, 2, {16, 1},  1025,  64, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, {0, 2, 1, 3}, false));
     test_cases.emplace_back(new test_flash_attn_ext(256, 256, 2, {16, 1}, 16384,   1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0));
 
     // MLA shape: the V cache is a sub-view of the K cache, with quantized KV
@@ -11270,8 +11281,16 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {
 
 // Test cases for performance evaluation: should be representative of real-world use cases
 static std::vector<std::unique_ptr<test_case>> make_test_cases_perf() {
+
     std::vector<std::unique_ptr<test_case>> test_cases;
 
+    // Qwen3.x 27B style attention: 24 Q / 4 KV heads, head size 256, q8_0 KV cache (prefill + decode)
+    for (int kv : { 8192, 32768, 33792, 34816, 65536, 66560, }) {
+        for (int nb : { 4, 512, 1024, }) {
+            test_cases.emplace_back(new test_flash_attn_ext(256, 256, 4, {6, 1}, kv, nb, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, {0, 2, 1, 3}, false));
+        }
+    }
+
     // Qwen3.x 27B style decode: 24 Q heads / 4 KV heads, head size 256, q8_0 KV cache
     for (int kv : { 1024, 4096, 8192, 16384, 65536, }) {
         test_cases.emplace_back(new test_flash_attn_ext(256, 256, 4, {6, 1}, kv, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0));
-- 
2.55.0

