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.
This commit is contained in:
François-Xavier Gsell
2026-10-05 10:37:54 +02:00
committed by GitHub
parent c173a53bdf
commit b3daa077a5
4 changed files with 93 additions and 55 deletions
+5 -2
View File
@@ -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));
@@ -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);
@@ -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();
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]);
}
barrier();
const uint slot = base + subgroup_base + subgroupBallotExclusiveBitCount(ballot);
slot += subgroupBallotBitCount(ballot);
}
#else
// Hillis-Steele inclusive prefix sum over the workgroup.
scan[tid] = selected ? 1u : 0u;
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;
}
// 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();
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) {
data_i[out_base + slot] = int32_t(k);
}
base += total;
scan[tid] += add;
barrier();
}
for (uint s = min(base, p.n_kv_max) + tid; s < p.n_kv_max; s += BLOCK_SIZE) {
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;
}
}
#endif
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);
}
}
+10
View File
@@ -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));