CUDA + ggml: add sparse-fa for DSV4/GLM - #27970
Conversation
ggerganov
left a comment
There was a problem hiding this comment.
Should we also try instead of extracting indices from the mask, to directly construct a dense KV cache (i.e. get_rows-style) and run the existing FA kernels without modifications?
| const bool need_f16_V = type_V == GGML_TYPE_F16; | ||
| constexpr size_t nbytes_shared = 0; | ||
| launch_fattn<D, cols_per_block, 1>(ctx, dst, fattn_kernel, nwarps, nbytes_shared, D, need_f16_K, need_f16_V, false); | ||
| launch_fattn<D, cols_per_block, 1>(ctx, dst, fattn_kernel, nwarps, nbytes_shared, D, need_f16_K, need_f16_V, false, false); |
There was a problem hiding this comment.
Where are the "vec" FA kernels located? Just from the filenames, it looks as if the sparse indices are not used in the "vec" case.
There was a problem hiding this comment.
Yeah we use the mma kernel for the vec (bs=1) case when f16 kv cache is used. I think the vec kernel is only used for quantized kv cache - I may be wrong
| static constexpr bool ggml_cuda_flash_attn_ext_mma_f16_may_use_sparse( | ||
| const int DKQ, const int DV, const int ncols1, const int ncols2) { | ||
| return (DKQ == 512 && DV == 512 && ncols1 == 1 && ncols2 == 8) || | ||
| (DKQ == 576 && DV == 512 && ncols1 == 1 && ncols2 == 16); | ||
| } |
There was a problem hiding this comment.
Do we need to restrict to only these head sizes?
Could you remind me what ncols1 and ncols2 refer to in the CUDA backend?
There was a problem hiding this comment.
ncols1 is the number of query tokens and ncols2 is the k/v tokens per block. It's restricted here because I only ran test this for dsv4, for qwen4 it was slower for prefill and faster for decode but qwen4 has other issues w.r.t to the mask so I left for a future PR
At least for the CUDA backend this was slower than not doing anything at all for dsv4. |
|
I have done a lot of attempts at optimizing sparse attention in DSv4 in a Vulkan focused llamacpp fork, and the biggest improvements was using dense attention for the continuous block of keys, sparse attention for the selected keys and then ultimately combining the result. The commit for this exact change can be found here: Nathanw1014@4bbe53e Then I did some further memory usage improvements by changing the implementation to a tiled version in a follow-up commit here: Nathanw1014@80728f3 These changes were vibe coded with oversight from me, but the idea itself might be useful for optimizing sparse attention. My apologies if this PR already does something similar. I am on my phone so I haven't really look at the code yet, but I thought I would share my experiences with optimizing sparse attention when I saw this PR. |
|
I have a Metal implementation that works with DSv4. But I also want to test it with Qwen4. Do you have a suggestion for a patch that enables sparse attention with the Qwen4 graph? |
|
@ggerganov you can just pass n_kv_max (which should be |
|
I run some tests on my container |
|
Rebased onto current master and benchmarked against it on an RTX PRO 6000 Blackwell, DeepSeek-V4-Flash IQ2_XXS, f16 KV cache, single sequence. Clear win at long context, small cost below 32k, which matches the shape of your own table.
Two passes per binary with the run order alternated, so the long context numbers are not a thermal artifact. One correctness note: in Patch : |
|
@ServeurpersoCom thanks, fixed in be029fd. |
|
Optional nits : Small cleanup patch on top of the PDL fix, nothing behavioural except one extra test case. It names the block size and the two activation thresholds, replaces the serial prefix sum with a warp scan over the eight warp counts, adds a static_assert for the ncols1 == 1 assumption the index addressing relies on, and drops a dead VARS_TO_STR18 plus a test case that was listed twice. The one real change is in the mask generator: it now reduces one row in every thirty two to a single finite entry, which is the degenerate case of the first token of a sequence and was not covered before. Built and ran the full FLASH_ATTN_EXT suite on a Blackwell card, 2948 pass, one fewer than before because of the removed duplicate. Take or leave any of it, none of it is load bearing. Running a greedy check next and I will approve after that. |
JohannesGaessler
left a comment
There was a problem hiding this comment.
Any comments regarding performance are only suggestions.
|
|
||
| // Use finite mask entries as a sparse K/V set. Set 0 to disable. | ||
| // n_kv_max must bound the number of finite entries in every mask row. | ||
| GGML_API void ggml_flash_attn_ext_set_sparse( |
There was a problem hiding this comment.
I don't feel strongly about this but wouldn't ggml_flash_attn_ext_set_n_kv_max be the more appropriate name?
| const int32_t index = i < i_sup ? indices[i] : -1; | ||
| src = index >= 0 ? KV + int64_t(index)*stride_KV + k*h2_per_chunk : zero; |
There was a problem hiding this comment.
It's not great to have a condition here. It would likely be preferable to pad indices with safe values that result in some redundant work.
| // Skip unused kernel variants for faster compilation: | ||
| if (use_logit_softcap && !(DKQ == 128 || DKQ == 256 || DKQ == 512)) { | ||
| NO_DEVICE_CODE; | ||
| return; | ||
| } | ||
| if (DKQ == 192 && ncols2 != 8 && ncols2 != 16) { | ||
| NO_DEVICE_CODE; | ||
| return; | ||
| } | ||
| #ifdef VOLTA_MMA_AVAILABLE | ||
| if (ncols1*ncols2 < 32) { | ||
| NO_DEVICE_CODE; | ||
| return; | ||
| } | ||
| #endif // VOLTA_MMA_AVAILABLE | ||
|
|
||
| #if __CUDA_ARCH__ == GGML_CUDA_CC_TURING | ||
| if (ncols1*ncols2 > 32) { | ||
| NO_DEVICE_CODE; | ||
| return; | ||
| } | ||
| #endif // __CUDA_ARCH__ == GGML_CUDA_CC_TURING | ||
|
|
||
| #if defined(AMD_WMMA_AVAILABLE) | ||
| if (ncols1*ncols2 < 16 || ncols2 == 1 || DKQ > 128) { | ||
| NO_DEVICE_CODE; | ||
| return; | ||
| } | ||
| #endif // defined(AMD_WMMA_AVAILABLE) | ||
|
|
||
| #if defined(AMD_MFMA_AVAILABLE) | ||
| if (ncols1*ncols2 < 16 || DKQ > 256) { | ||
| NO_DEVICE_CODE; | ||
| return; | ||
| } | ||
| #endif // defined(AMD_MFMA_AVAILABLE) |
There was a problem hiding this comment.
Please skip unused kernel templates here, I would suggest you re-use ggml_cuda_flash_attn_ext_mma_f16_may_use_sparse.
There was a problem hiding this comment.
does shall_use_sparse etc already do this?
There was a problem hiding this comment.
No, because shall_use_sparse is a host function function intended for kernel selection logic.
There was a problem hiding this comment.
In terms of program logic shall_use_sparse must be strictly a subset of may_use_sparse.
There was a problem hiding this comment.
Okay it should be there
|
|
||
| const half2 * V_h2 = V_is_K_view ? K_h2 : (const half2 *) (V + nb23*sequence + nb22*z_KV); | ||
| const float * sinks_f = sinks ? (const float *) sinks + zt_Q : nullptr; | ||
| const int32_t * indices = use_sparse ? sparse_indices + (int64_t(sequence % ne33)*ne31 + jt)*ne11 : nullptr; |
There was a problem hiding this comment.
| const int32_t * indices = use_sparse ? sparse_indices + (int64_t(sequence % ne33)*ne31 + jt)*ne11 : nullptr; | |
| const int32_t * indices = use_sparse ? sparse_indices + (int64_t(sequence % ne33)*ne31 + jt*ncols1)*ne11 : nullptr; |
I think the way you're calculating the indices may be incorrect here though this defect would as of right now not manifest as a bug due to ncols1 == 1.
There was a problem hiding this comment.
Sorry, after thinking about it some more I'm not sure what I suggested here is correct. For > 1 tokens it's not clear to me what the best way to handle indices is.
|
Greedy output diverges from master, as expected since compacting the selected entries changes the accumulation order, so I checked perplexity instead. Same 153806 token corpus, run at c=32768, well past the 4096 KV length needed to arm the sparse path:
0.0006 apart against a 0.0045 confidence interval, same direction on all four chunks. No sign of anything dropped. Good for me, I will leave the rest to @JohannesGaessler. |
Merge feature/adaptive-kv-stream (61 commits, base ece963f) into master (95ef7fc = upstream ggml-org tip, 5 commits after b10786). Conflicts resolved (all additive, both features preserved): - ggml-cuda/fattn-common.cuh: launch_fattn gains both use_sparse (sparse-FA ggml-org#27970) and output_partial/partial_dst/partial_meta args - ggml-cuda/fattn-mma-f16.cuh: flash_attn_ext_f16 / _process_tile gain both use_sparse and output_partial template params; case() refactored into case_impl<output_partial> + case()/partial_case() wrappers - src/llama-kv-cache.h/.cpp: constructor gains both name_tag and kv_stream_stage_bytes (both defaulted) Feature adds: block KV streaming for Qwen3.5-style models, staged KV pool, kv-stream-bench tool, benchmarks, unit tests.
Upstream's sparse flash attention (ggml-org#27970) keeps n_kv_max in FLASH_ATTN_EXT op_params[4], which is where the per-row log-sum-exp tail output kept its own flag. With both in the tree every n_kv_max > 0 launch reads as "write an LSE tail" and stores past the end of the output tensor; the reverse also holds, an LSE call reads as n_kv_max=1. Both were visible on the merged tree as 13 sentinel mismatches on the new n_kv_max test cases. Slot map for this op is now: [0..2] scale, max_bias, logit_softcap, [3] prec, [4] n_kv_max, [5] LSE. The CUDA reader uses the accessor so a later move is one edit.
|
Would you take a sparse instance for the Qwen3.8-Flash-Next head shape? #28349 tried to enable the QSA sparse path in build_attn_qsa, and on CUDA it is a no-op for this model. I ran diagnostics today and my agent helped me check the source and break out the data points:
On CUDA the only QSA speedup at depth I can measure for this model right now is the get_rows gather in #28213, +7% decode at 35k and +19% at 101k on a 5 card PCIe box. |
FLASH_ATTN_EXT with n_kv_max > 0 (ggml_flash_attn_ext_set_n_kv_max, upstream ggml-org#27970, implemented for CUDA only) attends the finite entries of the mask only. The qwen4exp QSA layers select 2048 cells per query out of the whole cache, so the dense masked attention paid the full n_kv per query at every depth: the prefill's quadratic term and, per decode token, a read of the whole cache. Vulkan implementation, coopmat1 path: flash_attn_sparse_idx.comp: for every tile of Br mask rows, the cells whose entry is finite for at least one row of the tile (counts[n_lists] then lists[n_lists][list_stride] in prealloc_y, list_stride = min(KV, Br * n_kv_max) rounded to Bc). Two modes chosen by the host: one workgroup per tile writing the ascending list (prefill, tiles enough to fill the GPU) and chunks of 2048 columns appended through an atomic on the count (decode, one tile per token). flash_attn_cm1.comp with FA_SPARSE: the tile loop walks the list instead of the cache (KV := list length, split_k splits the list), stages the tile's cell indices in shared memory and gathers K, V and the mask entries by cell, dequantizing q8_0 in the shader; the Clamp bounds logic covers the list padding. Host: sparse pipeline states (flag 16, separate SPIR-V, 8 bindings, 136-byte push constants, gated on maxPushConstantsSize), the prepass dispatch and a dense fallback for every other case. The path starts at KV >= 8 x n_kv_max (~16k cells): below that the tiles' unions cover most of the cache and the prepass costs more than the gather saves. GGML_VK_DISABLE_SPARSE_FA=1 disables it in the backend. qwen4exp passes n_kv_max = n_top_k for the QSA attention (LLAMA_QSA_NO_SPARSE_FA=1 withholds it, for A/B runs). test-backend-ops: sparse-mask cases with the qwen4exp shapes (head 256, 2 KV heads, GQA 12, kv 4096-32768, f16 and q8_0), dense CPU reference. Measured on gfx1151 (Vulkan), production model, KV q8_0, MTP off, dense vs sparse with the same binary: test-backend-ops FLASH_ATTN_EXT 5176/5176, test-llama-archs NMSE 8.6e-8, wikitext PPL c=16384: 2.5342 dense, 2.5346 sparse (run-to-run band 0.3%) llama-bench pp2048 558 -> 557 t/s, pp8192 520 -> 520 (gated off), pp32768 385 -> 417 t/s (+8%), tg32 unchanged decode at 39.5k, 7 turns of 256: 26.61 -> 27.83 t/s (+4.6%), the 2048-token re-prefill 6.78 -> 6.56 s 126.5k prefill: 524.5 -> 370.8 s (241 -> 341 t/s, +41%); re-prefill of 2048 tokens at 126k 12.0 -> 7.6 s Measured tile unions (16 rows, 2051 cells each): 52% of the cache at 8k, 33% at 16k, 17% at 32k, 16% at 40k, so the FA work drops 6x at 40k and ~17x at 126k; the gains above are limited by the gathered rows' staging and in-shader dequantization, which a compact per-tile f16 scratch on the direct-load path would remove. Claude-Session: https://claude.ai/code/session_01KXWaojVXitKAUbG1LUNGr2
FLASH_ATTN_EXT with n_kv_max > 0 (ggml_flash_attn_ext_set_n_kv_max, upstream ggml-org#27970, implemented for CUDA only) attends the finite entries of the mask only. The qwen4exp QSA layers select 2048 cells per query out of the whole cache, so the dense masked attention paid the full n_kv per query at every depth: the prefill's quadratic term and, per decode token, a read of the whole cache. Vulkan implementation, coopmat1 path: flash_attn_sparse_idx.comp: for every tile of Br mask rows, the cells whose entry is finite for at least one row of the tile (counts[n_lists] then lists[n_lists][list_stride] in prealloc_y, list_stride = min(KV, Br * n_kv_max) rounded to Bc). Two modes chosen by the host: one workgroup per tile writing the ascending list (prefill, tiles enough to fill the GPU) and chunks of 2048 columns appended through an atomic on the count (decode, one tile per token). flash_attn_cm1.comp with FA_SPARSE: the tile loop walks the list instead of the cache (KV := list length, split_k splits the list), stages the tile's cell indices in shared memory and gathers K, V and the mask entries by cell, dequantizing q8_0 in the shader; the Clamp bounds logic covers the list padding. Host: sparse pipeline states (flag 16, separate SPIR-V, 8 bindings, 136-byte push constants, gated on maxPushConstantsSize), the prepass dispatch and a dense fallback for every other case. The path starts at KV >= 8 x n_kv_max (~16k cells): below that the tiles' unions cover most of the cache and the prepass costs more than the gather saves. GGML_VK_DISABLE_SPARSE_FA=1 disables it in the backend. qwen4exp passes n_kv_max = n_top_k for the QSA attention (LLAMA_QSA_NO_SPARSE_FA=1 withholds it, for A/B runs). test-backend-ops: sparse-mask cases with the qwen4exp shapes (head 256, 2 KV heads, GQA 12, kv 4096-32768, f16 and q8_0), dense CPU reference. Measured on gfx1151 (Vulkan), production model, KV q8_0, MTP off, dense vs sparse with the same binary: test-backend-ops FLASH_ATTN_EXT 5176/5176, test-llama-archs NMSE 8.6e-8, wikitext PPL c=16384: 2.5342 dense, 2.5346 sparse (run-to-run band 0.3%) llama-bench pp2048 558 -> 557 t/s, pp8192 520 -> 520 (gated off), pp32768 385 -> 417 t/s (+8%), tg32 unchanged decode at 39.5k, 7 turns of 256: 26.61 -> 27.83 t/s (+4.6%), the 2048-token re-prefill 6.78 -> 6.56 s 126.5k prefill: 524.5 -> 370.8 s (241 -> 341 t/s, +41%); re-prefill of 2048 tokens at 126k 12.0 -> 7.6 s Measured tile unions (16 rows, 2051 cells each): 52% of the cache at 8k, 33% at 16k, 17% at 32k, 16% at 40k, so the FA work drops 6x at 40k and ~17x at 126k; the gains above are limited by the gathered rows' staging and in-shader dequantization, which a compact per-tile f16 scratch on the direct-load path would remove.
Upstream's sparse flash attention (ggml-org#27970) keeps n_kv_max in FLASH_ATTN_EXT op_params[4], which is where the per-row log-sum-exp tail output kept its own flag. With both in the tree every n_kv_max > 0 launch reads as "write an LSE tail" and stores past the end of the output tensor; the reverse also holds, an LSE call reads as n_kv_max=1. Both were visible on the merged tree as 13 sentinel mismatches on the new n_kv_max test cases. Slot map for this op is now: [0..2] scale, max_bias, logit_softcap, [3] prec, [4] n_kv_max, [5] LSE. The CUDA reader uses the accessor so a later move is one edit.
Ports the test cases from the V2 patch that could not be applied here: mask_pattern 1-5 (masked interior tiles, a single visible element in an otherwise masked tile, fully masked queries with sinks, zero tiles and interior holes crossing packed-word boundaries) and mask_broadcast, which poisons the backing storage outside the logical mask. The n_kv_max / init_tensor_kq_mask_sparse branch is left out: this tree has no sparse FA (ggml-org#27970). Co-Authored-By: Claude Opus 5 <noreply@anthropic.com>
…elected (upstream ggml-org#27970) GLM-5.3-Flash's DSA layers mask the whole cache down to the cells their indexer picked (2048 of them and a tail of at most 3 on the proxy), and the MMA kernel read every cell under that mask: at 32K cached tokens the flash attention node took 595 / 667 us a decode token on the RTX 5080 / 5070 Ti under -sm tensor. Ported by hand from ggml-org/llama.cpp 8e93a97 (ggml-org#27970, Aman Gupta): FLASH_ATTN_EXT takes a bound n_kv_max on each mask row's finite cells (ggml_flash_attn_ext_set_n_kv_max, op_params[5], since this fork's [4] is mask_prefix), a kernel compacts each row's finite cells into an index list, and the MMA kernel for <512,512,1,8> and <576,512,1,16> loads K, V and the mask by index over n_kv_max cells, one query a tile, instead of the whole cache. Adapted to this fork: launch_fattn takes the sparse flag last, so the tile and vector kernels are unchanged; the sparse path needs an f16 mask (the packed I16 mask never reaches it: build_attn_sparse's mask is the cache's width), turns off kv_range, kv_live and the KV_max scans, and runs one stage. build_attn_sparse passes GGML_PAD(top_k + kpool - 1, 32) = 2080 on the proxy: the top-k pools' cells and the tail, positions [(q + 1)/r*r, q], at most r - 1 cells unless a sequence holds two cells at one position (rig-glm), which the 32-cell pad keeps at no cost. The inherited DSA path passes top_k as upstream does. deepseek4's compressed attention is not wired here (no DSV4 check past 4K on this box). When: an NVIDIA MMA device, n_kv_max set, the f16 mask of one head, no ALiBi or softcap, and K >= max(4096, n_gather) with n_gather = min(n_tokens, 64/ncols2) * n_kv_max, the cells the gather reads for the queries one dense pass covers. Upstream asks 2 n_gather; on these cards (test-backend-ops perf, GLM's shape with 32 heads on the latent, kv 8448, 16640, 33280 x batch 1 to 512) the gather ran at 0.25-0.99x the dense time wherever K >= n_gather and 1.13-1.27x wherever it is under, and the 2x margin gave up a 3-token verify at 8K cells (0.73x) and prefills at 16K (0.63-0.72x). The 44-layer proxy, both cards, -sm tensor, against GGML_CUDA_FATTN_SPARSE_LEGACY=1 (bands declared before each run): - the FLASH_ATTN_EXT node a decode token (nsys, NVTX, graphs off), card 0 / card 1: -d 2048 197/223 us against 195/224 (under the gate, the same kernels); -d 8192 260/259 against 407/446 (0.64/0.58x); -d 32768 337/335 against 595/667 (0.57/0.50x). The attention kernel itself is 17 us a layer at both depths (dense: 31-35 at 8K, 48-55 at 32K); the rest is the compaction, one block of 256 threads a row, 3.6 us at 8K and 11.2 at 32K (next). - tg32 @ d32768, graphs on: a process's reps 2-6 at 89.3 against 88.2 t/s (median; +1.2 %, where the node's 258-332 us of an 11.3 ms token predicted 2.3-2.9 %); its first rep is slower each way and more so here (73.9 against 78.2), so 3 interleaved pairs of -r 3 means came out 79.3/75.6/82.2 against 81.1/82.1/80.3 (the band, above in each pair, missed). - pp512 (3 interleaved pairs): -d 32768 1235.6 against 975.0 t/s (+26.7 %); -d 8192 1623.9 against 1603.3, dense each way under the gate (before it, the sparse kernel there cost 4.0 %). - PPL at -c 8192 with the sparse kernel taking every batch past 4096 cells (before the gate; one-token batches are the same either way): 348974.1045 against 349039.3242 (0.019 %); KLD mean 0.001611, max 0.002327, same top token 88.8 % (the proxy's order noise: ub 1 against ub 3 is 0.31 % and 0.0068). - test-backend-ops FLASH_ATTN_EXT: every case (3229/3229) and the sparse cases pass on both cards, each way. Upstream's cases, GLM's (512/512, 64 heads, 8192 cells, 2051-cell bound, 1 and 3 queries, 2 streams) and a dense fallback under its gather. Mutants that fail them: the gather reading row k_VKQ_0 + i instead of its index (7 of 7 sparse cases), the compaction dropping each row's last index (upstream's 4; the GLM cases' 1 in 2051 is under the error bound). - the crossover grid is in test-backend-ops perf (GLM's shape at 3 cache lengths and 8 batches). Co-authored-by: Aman Gupta <amangupta052@gmail.com>
Keep this even if PR ggml-org#27754 is dropped -- it fixes an upstream OMISSION, not our vendoring. build_attn_sparse passed n_kv_max=0 for glm5next, which disables upstream's sparse-FA gather (it requires n_kv_max > 0, ggml/src/ggml-cuda/fattn.cu:138). Flash attention therefore ran DENSE over the full n_kv: flat in memory, but O(n_kv) in time, so the model's top_k 2048 / kpool 4 sparse selection bought nothing at depth. The correct bound is already computed by the arch: glm5next_n_select() = indexer_top_k + indexer_kpool - 1 = 2051 (src/models/glm5next.cpp:9-17), and the top_k tensor is documented as [n_select, n_tokens/n_stream, n_stream] (src/llama-graph.h:1382), so top_k->ne[0] IS that bound. The glm-dsa path has always passed it (src/llama-graph.cpp:3150); this makes glm5next consistent. Cause of the omission: sparse-FA landed upstream in ggml-org#27970 (2026-09-02), after PR ggml-org#27754's graph code was written. The PR merged master in but never wired glm5next into it; its head 86ebfef still passes 0. Expected effect: no change below ~4k depth (the gather only engages at n_kv >= max(4096, 2*n_gather) = 4102), then progressively larger -- attention work per token goes from O(n_kv) to O(2051), i.e. ~4.5x less at 9k and ~64x less at 131072. NOT YET MEASURED: needs a rebuild on the target host. Builds clean locally.
Overview
Based on @fairydreaming's PR #25917, it is the same idea but instead just relies on a hint from the API about the max number of live kv-entries per token. Further sparse attention methods like QSA(qwen4) can be enrolled at a later stage as they use similar mechanisms. Performance of 2-bit DSV4 quant on DGX spark. More testing will be appreciated
Additional information
Requirements