CUDA: enable sparse fa for qwen4 - #28770
Conversation
JohannesGaessler
left a comment
There was a problem hiding this comment.
I don't really have anything to add. The implementation seems straightforwardly correct to me and is done well in terms of software architecture. One thing to keep in mind is that this is in essence the exact same infrastructure as would be needed for things like PagedAttention.
| return (DKQ == 512 && DV == 512 && ncols1 == 1 && ncols2 == 8) || | ||
| (DKQ == 576 && DV == 512 && ncols1 == 1 && ncols2 == 16); | ||
| (DKQ == 576 && DV == 512 && ncols1 == 1 && ncols2 == 16) || | ||
| (DKQ == 256 && DV == 256 && ncols1 == 1 && ncols2 == 8) || | ||
| (DKQ == 256 && DV == 256 && ncols1 == 8 && ncols2 == 8); | ||
| } |
There was a problem hiding this comment.
It is good to enable sparse attention for ncols1 > 1 since that will drastically improve the prefill performance. However, if it is only enabled for 1 and 8 it will cause trouble in combination with speculative methods. My opinion is that we should compile the template specializations for batch sizes 2 and 4; if the compilation becomes too bloated we should shave off template specializations somewhere else.
It's also not clear to me why the new code for sparse attention with ncols1 > 1 would work for Qwen 4 but not the other models.
There was a problem hiding this comment.
I only tested it with Qwen's headsizes, I guess it work anyway. I think I can expand the condition to be ncols==1, 2, 4, 8 and ncols2=8, 16.
|
One thing I forgot: the optimal |
Squash of ggml-org#28770 (am17an) at 41a4ad0. Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
|
As discussed offline, we can proceed to merge this and do optimizations in a follow up. |
ggerganov
left a comment
There was a problem hiding this comment.
at twice this value we enable the sparse FA. For Qwen4 this value is 32768 ctx.
Why do we observe improvement even at smaller depths (e.g. 10k)?
|
@ggerganov I expect this is because of the variation on my DGX spark because of the PLE table. This is with |
|
Since this PR, sparse decode on CUDA is up to 2x slower (reported in #29281, reproduced on my RTX PRO 6000): the runtime query loop in flash_attn_mask_to_sparse_indices prevents the scan from unrolling, and this small patch a0d1bb6 templates it on ncols1, which restores the performance while keeping the batched path intact. |
|
Preferably template the function by |
Upstream enables the sparse FA path for qwen4exp (ggml-org#28770). On Metal it runs the vec kernel, one query row and one head per threadgroup, which cost 12% of prefill at 4k context on M1 Max. kernel_flash_attn_ext_sparse_hr gives one threadgroup to each (query row, KV head): the heads that share the KV head are the rows of a simdgroup-matrix tile, so each selected K/V row is gathered once for all of them. It walks the index lists of kernel_flash_attn_ext_vec_idx. F16 K/V, DK == DV in {128, 256}, at most 16 heads per KV head, mask shared by the heads, 32+ query rows. For these GQA prefill shapes the backend runs the dense kernel until the KV range is 3x n_kv_max (12x with GGML_METAL_FA_SPARSE_HR=0, for the vec kernel): the sparse kernels pay for the gather. Qwen3.8-Flash-Next pp2048 at depth 8k / 16k / 32k / 64k, M1 Max: 253 / 247 / 232 / 204 tok/s, dense 241 / 213 / 177 / 174. Kernel: 10.1 vs 41.9 ms (4k), 11.8 vs 61.9 ms (32k) against the vec kernel. Not bit-identical to the vec kernel (summation order); matches the CPU reference.
…eam ggml-org#28770) Cherry-pick of ggml-org/llama.cpp 3cf0325 "CUDA: enable sparse fa for qwen4 (ggml-org#28770)": one index list per group of ncols1 queries, the union of their visible columns, with the live count of each list stored after the lists, and the (256, 256, 8, 8) sparse instantiation for the qwen4 head shape. Resolved against this fork's explicit index lists (src[5], the native QSA path): those stay one list per query with ncols1 == 1 and no counts; the kernel is told so with ne31 = -1. With ncols1 > 1 src[5] carries group lists followed by counts. One change to the kernel contract on top of upstream: a sparse tile's mask is compact, [n_kv_max, n_queries] indexed by position in the tile's index list rather than by KV column. launch_fattn gathers it that way from a dense mask (flash_attn_gather_sparse_mask); the QSA path, which has no dense mask (it would be n_kv x n_tokens at 262k context), builds it directly. The upstream (256,256,gqa 12) sparse tests pass against the CPU reference, as do two added q8_0-cache variants of them. The qwen4exp.cpp hunk (sparse attention in the non-native path) is left out; the native path is what this fork runs.
Upstream enables the sparse FA path for qwen4exp (ggml-org#28770). On Metal it runs the vec kernel, one query row and one head per threadgroup, which cost 12% of prefill at 4k context on M1 Max. kernel_flash_attn_ext_sparse_hr gives one threadgroup to each (query row, KV head): the heads that share the KV head are the rows of a simdgroup-matrix tile, so each selected K/V row is gathered once for all of them. It walks the index lists of kernel_flash_attn_ext_vec_idx. F16 K/V, DK == DV in {128, 256}, at most 16 heads per KV head, mask shared by the heads, 32+ query rows. For these GQA prefill shapes the backend runs the dense kernel until the KV range is 3x n_kv_max (12x with GGML_METAL_FA_SPARSE_HR=0, for the vec kernel): the sparse kernels pay for the gather. Qwen3.8-Flash-Next pp2048 at depth 8k / 16k / 32k / 64k, M1 Max: 253 / 247 / 232 / 204 tok/s, dense 241 / 213 / 177 / 174. Kernel: 10.1 vs 41.9 ms (4k), 11.8 vs 61.9 ms (32k) against the vec kernel. Not bit-identical to the vec kernel (summation order); matches the CPU reference.
Upstream enables the sparse FA path for qwen4exp (ggml-org#28770). On Metal it runs the vec kernel, one query row and one head per threadgroup, which cost 12% of prefill at 4k context on M1 Max. kernel_flash_attn_ext_sparse_hr gives one threadgroup to each (query row, KV head): the heads that share the KV head are the rows of a simdgroup-matrix tile, so each selected K/V row is gathered once for all of them. It walks the index lists of kernel_flash_attn_ext_vec_idx. F16 K/V, DK == DV in {128, 256}, at most 16 heads per KV head, mask shared by the heads, 32+ query rows. For these GQA prefill shapes the backend runs the dense kernel until the KV range is 3x n_kv_max (12x with GGML_METAL_FA_SPARSE_HR=0, for the vec kernel): the sparse kernels pay for the gather. Qwen3.8-Flash-Next pp2048 at depth 8k / 16k / 32k / 64k, M1 Max: 253 / 247 / 232 / 204 tok/s, dense 241 / 213 / 177 / 174. Kernel: 10.1 vs 41.9 ms (4k), 11.8 vs 61.9 ms (32k) against the vec kernel. Not bit-identical to the vec kernel (summation order); matches the CPU reference.
Overview
Cont #27970. Enable sparse-fa for Qwen4. This model's attention is un-optimized at the moment, the entire kv-cache is re-scored everytime, we should fix that as well.
How this works ncols1=8 is take all the union of the tokens being used, so at max it will use see
ncols1 * n_kv_maxtokens, at twice this value we enable the sparse FA. For Qwen4 this value is 32768 ctx.Additional information
Results on a DGX spark:
Requirements