Commit b3daa077a for llama.cpp

commit b3daa077a56cfda820b22fc28941a6d82080e8e0
Author: François-Xavier Gsell <fxgsell@gmail.com>
Date:   Mon Oct 5 16:37:54 2026 +0800

    vulkan: sparse flash attention for quantized K/V (#29639)

    * vulkan: sparse flash attention for quantized K/V

    Assisted-by: Claude

    * vulkan: single-scan sparse FA index compaction

    The compaction ran one workgroup per mask row and walked the row in
    BLOCK_SIZE chunks, with a workgroup scan per chunk. For decode that is
    one workgroup doing KV/1024 barrier-bound iterations, so at 128k cells
    it cost more than the sparse attention it feeds.

    Split the row into contiguous segments instead: one per subgroup with
    ballot counting over coalesced loads, or one per thread without
    subgroups. A single scan over the segment counts then gives each
    segment its output offset. The index list stays ascending.

diff --git a/ggml/src/ggml-vulkan/ggml-vulkan.cpp b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
index a4c7bcc8c..0588a0040 100644
--- a/ggml/src/ggml-vulkan/ggml-vulkan.cpp
+++ b/ggml/src/ggml-vulkan/ggml-vulkan.cpp
@@ -8145,11 +8145,14 @@ void ggml_vk_flash_attn(ggml_backend_vk_context * ctx, vk_context& subctx, const
     // Sparse mask hint (op_params[4]): compact the <= n_kv_max finite positions and gather only those.
     const int32_t n_kv_max = mask ? ggml_get_op_params_i32(dst, 4) : 0;
     static const bool disable_sparse = getenv("GGML_VK_FA_SPARSE_DISABLE") != nullptr;
+    const bool kv_f16 = k_type_eff == GGML_TYPE_F16 && v_type_eff == GGML_TYPE_F16;
     // cm2 dense is fast, so it needs a larger reduction to win.
-    const int64_t min_ratio = tuning_params.path == FA_COOPMAT2 ? 4 : 2;
+    // With quantized K/V, sparse only breaks even around 16x (measured on RDNA3/RDNA4).
+    const int64_t min_ratio = tuning_params.path == FA_COOPMAT2 ? 4 : (kv_f16 ? 2 : 16);
     const bool use_sparse = !disable_sparse && 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 &&
+                            // the cm2 sparse gather only reads f16
+                            (kv_f16 || tuning_params.path != FA_COOPMAT2) &&
                             nem0 == KV &&
                             (int64_t)KV >= std::max<int64_t>(4096, min_ratio * (int64_t)n_kv_max) &&
                             (gqa_ratio > 1 || (tuning_params.path == FA_SCALAR && N == 1));
diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp
index 107d44aaa..5bba47834 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn.comp
@@ -285,8 +285,9 @@ void main() {
                 const uint32_t block = ib % (HSK / 32);
                 if (idx + gl_WorkGroupSize.x <= quant_iters || c < Bc) {
                     const uint buf_ib = c * qf_stride + block;
-                    if (!KV_bounds_check || j * Bc + c < KV) {
-                        const uint global_ib = (j * Bc + c) * k_stride + block;
+                    uint32_t kcol;
+                    if (fa_kv_index(j * Bc + c, kcol)) {
+                        const uint global_ib = kcol * k_stride + block;
                         k_block_to_shmem(buf_ib, global_ib, iqs, k_offset);
                     } else {
                         k_block_to_shmem_zero(buf_ib, iqs);
@@ -363,7 +364,8 @@ void main() {
                                 (hsk4 % 2 == 0) ? 2 : 1;

         [[unroll]] for (uint32_t c = 0; c < cols_per_thread; ++c) {
-            if (KV_bounds_check && j * Bc + c * cols_per_iter + col_tid >= KV) {
+            uint32_t kcol;
+            if (!fa_kv_index(j * Bc + c * cols_per_iter + col_tid, kcol)) {
                 continue;
             }

@@ -400,7 +402,7 @@ void main() {
                         }
                     }
                 } else {
-                    const uint coord = (j * Bc + c * cols_per_iter + col_tid) * k_stride * BLOCK_SIZE_K + 4 * (d_tid * (HSK_per_thread / 4) + d_block);
+                    const uint coord = kcol * k_stride * BLOCK_SIZE_K + 4 * (d_tid * (HSK_per_thread / 4) + d_block);
                     const uint ib = coord / BLOCK_SIZE_K;
                     const uint iqs = (coord % BLOCK_SIZE_K);

diff --git a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_sparse_compact.comp b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_sparse_compact.comp
index 3d3136266..50eeebece 100644
--- a/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_sparse_compact.comp
+++ b/ggml/src/ggml-vulkan/vulkan-shaders/flash_attn_sparse_compact.comp
@@ -26,14 +26,22 @@ layout (push_constant) uniform parameter {
 } p;

 #ifdef USE_SUBGROUPS
-shared uvec4 ballots_sh[NUM_SUBGROUPS];
+shared uint counts_sh[NUM_SUBGROUPS];
 #else
 shared uint scan[BLOCK_SIZE];
 #endif

+bool is_selected(const uint m_idx) {
+    const float v = float(data_m[m_idx]);
+    return !isinf(v) && !isnan(v);
+}
+
 // One workgroup per mask row: compact the finite-mask KV positions into a
-// per-row index list of length n_kv_max, -1 padded. Emitted in ascending KV
-// order so the downstream attention accumulation is deterministic.
+// per-row index list of length n_kv_max, -1 padded, in ascending KV order so
+// the downstream attention accumulation is deterministic.
+// The row is split into contiguous segments, one per subgroup (or per thread
+// without subgroups), so it needs a single workgroup scan instead of one per
+// BLOCK_SIZE chunk.
 void main() {
     const uint i1 = gl_WorkGroupID.x;
     const uint i2 = gl_WorkGroupID.y;
@@ -43,60 +51,75 @@ void main() {
     const uint m_base   = i3 * p.nbm3 + i2 * p.nbm2 + i1 * p.nbm1;
     const uint out_base = ((i3 * p.nem2 + i2) * p.nem1 + i1) * p.n_kv_max;

-    uint base = 0;
-    for (uint chunk = 0; chunk < p.KV; chunk += BLOCK_SIZE) {
-        const uint k = chunk + tid;
-        bool selected = false;
-        if (k < p.KV) {
-            const float v = float(data_m[m_base + k]);
-            selected = !isinf(v) && !isnan(v);
-        }
-
 #ifdef USE_SUBGROUPS
+    const uint sg   = gl_SubgroupID;
+    const uint lane = gl_SubgroupInvocationID;
+    const uint seg  = (p.KV + gl_NumSubgroups - 1) / gl_NumSubgroups;
+    const uint seg_begin = min(sg * seg, p.KV);
+    const uint seg_end   = min(seg_begin + seg, p.KV);
+
+    // Lanes read consecutive positions, so each step is one coalesced load.
+    uint count = 0;
+    for (uint k0 = seg_begin; k0 < seg_end; k0 += gl_SubgroupSize) {
+        const uint k = k0 + lane;
+        count += subgroupBallotBitCount(subgroupBallot(k < seg_end && is_selected(m_base + k)));
+    }
+    if (subgroupElect()) {
+        counts_sh[sg] = count;
+    }
+    barrier();
+
+    uint slot = 0;
+    uint total = 0;
+    for (uint s = 0; s < gl_NumSubgroups; ++s) {
+        slot  += s < sg ? counts_sh[s] : 0u;
+        total += counts_sh[s];
+    }
+
+    for (uint k0 = seg_begin; k0 < seg_end && slot < p.n_kv_max; k0 += gl_SubgroupSize) {
+        const uint k = k0 + lane;
+        const bool selected = k < seg_end && is_selected(m_base + k);
         const uvec4 ballot = subgroupBallot(selected);
-        if (subgroupElect()) {
-            ballots_sh[gl_SubgroupID] = ballot;
+        const uint pos = slot + subgroupBallotExclusiveBitCount(ballot);
+        if (selected && pos < p.n_kv_max) {
+            data_i[out_base + pos] = int32_t(k);
         }
-        barrier();
+        slot += subgroupBallotBitCount(ballot);
+    }
+#else
+    const uint run   = (p.KV + BLOCK_SIZE - 1) / BLOCK_SIZE;
+    const uint begin = min(tid * run, p.KV);
+    const uint end   = min(begin + run, p.KV);
+
+    uint count = 0;
+    for (uint k = begin; k < end; ++k) {
+        count += is_selected(m_base + k) ? 1u : 0u;
+    }

-        uint subgroup_base = 0;
-        uint total = 0;
-        [[unroll]] for (uint s = 0; s < gl_NumSubgroups; ++s) {
-            if (s == gl_SubgroupID) {
-                subgroup_base = total;
-            }
-            total += subgroupBallotBitCount(ballots_sh[s]);
+    // Hillis-Steele inclusive prefix sum of the per-thread counts.
+    scan[tid] = count;
+    barrier();
+    for (uint off = 1; off < BLOCK_SIZE; off <<= 1) {
+        uint add = 0;
+        if (tid >= off) {
+            add = scan[tid - off];
         }
         barrier();
-
-        const uint slot = base + subgroup_base + subgroupBallotExclusiveBitCount(ballot);
-#else
-        // Hillis-Steele inclusive prefix sum over the workgroup.
-        scan[tid] = selected ? 1u : 0u;
+        scan[tid] += add;
         barrier();
-        for (uint off = 1; off < BLOCK_SIZE; off <<= 1) {
-            uint add = 0;
-            if (tid >= off) {
-                add = scan[tid - off];
-            }
-            barrier();
-            scan[tid] += add;
-            barrier();
-        }
-
-        const uint inclusive = scan[tid];
-        const uint total = scan[BLOCK_SIZE - 1];
-        const uint slot = base + inclusive - 1u;
-#endif
+    }

-        if (selected && slot < p.n_kv_max) {
+    const uint total = scan[BLOCK_SIZE - 1];
+    uint slot = scan[tid] - count;
+    for (uint k = begin; k < end && slot < p.n_kv_max; ++k) {
+        if (is_selected(m_base + k)) {
             data_i[out_base + slot] = int32_t(k);
+            ++slot;
         }
-        base += total;
-        barrier();
     }
+#endif

-    for (uint s = min(base, p.n_kv_max) + tid; s < p.n_kv_max; s += BLOCK_SIZE) {
+    for (uint s = min(total, p.n_kv_max) + tid; s < p.n_kv_max; s += BLOCK_SIZE) {
         data_i[out_base + s] = int32_t(-1);
     }
 }
diff --git a/tests/test-backend-ops.cpp b/tests/test-backend-ops.cpp
index bc32aae77..d8893025e 100644
--- a/tests/test-backend-ops.cpp
+++ b/tests/test-backend-ops.cpp
@@ -11359,6 +11359,14 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_eval() {

     // Qwen QSA: 256/256, gqa 12, budget 2048.
     test_cases.emplace_back(new test_flash_attn_ext(256, 256, 2, {12, 1}, 8192, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false, 2048));
+    // quantized cache, deep enough to take the sparse path
+    test_cases.emplace_back(new test_flash_attn_ext(256, 256, 2, {12, 1}, 32768, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, {0, 1, 2, 3}, true, false, 2048));
+    test_cases.emplace_back(new test_flash_attn_ext(256, 256, 2, {12, 1}, 32768, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q4_0, GGML_TYPE_Q4_0, {0, 1, 2, 3}, true, false, 2048));
+    // single head with quantized K (MMQ on Vulkan)
+    test_cases.emplace_back(new test_flash_attn_ext(128, 128, 1, { 1, 1}, 8192, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, {0, 1, 2, 3}, true, false, 512));
+    test_cases.emplace_back(new test_flash_attn_ext(128, 128, 1, { 1, 1}, 8192, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q4_1, GGML_TYPE_Q4_1, {0, 1, 2, 3}, true, false, 512));
+    // KV not a multiple of the compaction workgroup size.
+    test_cases.emplace_back(new test_flash_attn_ext(128, 128, 1, { 8, 1}, 5003, 4, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false,  512));

     // more V-is-sub-view-of-K cases: other head shapes, and full views with equal head sizes
     test_cases.emplace_back(new test_flash_attn_ext(320, 256, 1, {32, 1}, 512, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, true));
@@ -11830,6 +11838,8 @@ static std::vector<std::unique_ptr<test_case>> make_test_cases_perf() {
         test_cases.emplace_back(new test_flash_attn_ext(512, 512, 1, { 8, 1}, kv, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false,  512));
         test_cases.emplace_back(new test_flash_attn_ext(576, 512, 1, {16, 1}, kv, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, true,   512));
         test_cases.emplace_back(new test_flash_attn_ext(256, 256, 2, {12, 1}, kv, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16, {0, 1, 2, 3}, true, false, 2048));
+        test_cases.emplace_back(new test_flash_attn_ext(256, 256, 2, {12, 1}, kv, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, {0, 1, 2, 3}, true, false, 2048));
+        test_cases.emplace_back(new test_flash_attn_ext(256, 256, 2, {12, 1}, kv, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_Q8_0, GGML_TYPE_Q8_0, {0, 1, 2, 3}, true, false, 0));
     }

     test_cases.emplace_back(new test_flash_attn_ext(64, 64, 8, {8, 1}, 7680, 1, true, false, 0, 0, GGML_PREC_F32, GGML_TYPE_F16, GGML_TYPE_F16));