Skip to content

CUDA: enable sparse fa for qwen4 - #28770

Merged
am17an merged 1 commit into
masterfrom
aman/qwen-sparse
Sep 20, 2026
Merged

am17an merged 1 commit into
masterfrom
aman/qwen-sparse

Conversation

@am17an

@am17an am17an commented Sep 11, 2026 •

Copy link
Copy Markdown
Contributor

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_max tokens, at twice this value we enable the sparse FA. For Qwen4 this value is 32768 ctx.

Additional information

Results on a DGX spark:

CPU Model Test t/s qsa-blk-cache t/s aman/qwen-sparse Speedup
CPU qwen4exp A3B IQ1_S - 1.5625 bpw pp2048@d10000 615.70 663.17 1.08
CPU qwen4exp A3B IQ1_S - 1.5625 bpw pp2048@d20000 500.49 544.11 1.09
CPU qwen4exp A3B IQ1_S - 1.5625 bpw pp2048@d50000 405.18 469.03 1.16
CPU qwen4exp A3B IQ1_S - 1.5625 bpw pp2048@d100000 252.81 318.28 1.26
CPU qwen4exp A3B IQ1_S - 1.5625 bpw tg32@d10000 22.98 23.56 1.03
CPU qwen4exp A3B IQ1_S - 1.5625 bpw tg32@d20000 20.76 23.58 1.14
CPU qwen4exp A3B IQ1_S - 1.5625 bpw tg32@d50000 15.88 18.72 1.18
CPU qwen4exp A3B IQ1_S - 1.5625 bpw tg32@d100000 12.00 14.19 1.18

Requirements

@am17an
am17an requested review from a team, CISC and ggerganov as code owners September 11, 2026 16:41
@github-actions github-actions Bot added model Model specific testing Everything test related ggml changes relating to the ggml tensor library for machine learning CUDA Related to the CUDA backend labels Sep 11, 2026

@JohannesGaessler JohannesGaessler left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Comment on lines 1760 to 1764
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);
}

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

@JohannesGaessler

Copy link
Copy Markdown
Contributor

One thing I forgot: the optimal ncols1 value for prefill may not be the highest one at large context depths. As long as you're I/O bound then increasing ncols1 is always beneficial since you would need to load the KV data anyways and it doesn't matter if you waste a bit of compute because that's not the bottleneck anyways. But at some point you will become compute bound so reducing wasted compute can become worthwhile. And at large context depths K/V data may be less likely to be shared across one "group" as you called it here. So in that case running the template specialization with ncols1 == 4 may be faster than the one with ncols1 == 8. My intuition is though that on modern NVIDIA GPUs ncols1 == 8 is optimal.

@JohannesGaessler JohannesGaessler self-assigned this Sep 12, 2026
abrisene pushed a commit to abrisene/llama.cpp that referenced this pull request Sep 18, 2026
Squash of ggml-org#28770 (am17an) at 41a4ad0.

Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
@am17an

am17an commented Sep 20, 2026

Copy link
Copy Markdown
Contributor Author

As discussed offline, we can proceed to merge this and do optimizations in a follow up.

@am17an am17an added the merge ready A maintainer can use this label to indicate that they consider the changes final and ready to merge. label Sep 20, 2026

@ggerganov ggerganov left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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

@am17an

am17an commented Sep 20, 2026

Copy link
Copy Markdown
Contributor Author

@ggerganov I expect this is because of the variation on my DGX spark because of the PLE table. This is with -lzm off but still there is some variation

@ServeurpersoCom

Copy link
Copy Markdown
Contributor

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.

@JohannesGaessler

Copy link
Copy Markdown
Contributor

Preferably template the function by ncols1 and check in host code whether an OOB check is actually needed. During prefill the number of tokens is very often a power of 2 so the OOB check is unnecessary.

mihailescu2m added a commit to mihailescu2m/llama.cpp that referenced this pull request Sep 23, 2026
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.
firelzrd added a commit to firelzrd/llama.cpp-2x4ever that referenced this pull request Sep 24, 2026
…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.
mihailescu2m added a commit to mihailescu2m/llama.cpp that referenced this pull request Sep 26, 2026
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.
Prototyped pushed a commit to Prototyped/llama.cpp that referenced this pull request Sep 26, 2026
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.
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

CUDA Related to the CUDA backend ggml changes relating to the ggml tensor library for machine learning merge ready A maintainer can use this label to indicate that they consider the changes final and ready to merge. model Model specific testing Everything test related

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants