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));