mirror of
https://github.com/ggml-org/llama.cpp.git
synced 2026-10-08 22:10:37 +02:00
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:
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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));
|
||||
|
||||
Reference in New Issue
Block a user