Skip to content

cuda: use the MMA flash attention kernel for GQA above 4 with quantized K/V on Ada - #307

Open
sb32445 wants to merge 2 commits into
PrismML-Eng:prismfrom
sb32445:pr/fattn-gqa-mma
Open

sb32445 wants to merge 2 commits into
PrismML-Eng:prismfrom
sb32445:pr/fattn-gqa-mma

Conversation

@sb32445

@sb32445 sb32445 commented Oct 4, 2026

Copy link
Copy Markdown

Overview

On Ada, flash attention with quantized K/V takes the vector kernel for 1-2 queries regardless of the GQA ratio. With GQA above 4 the MMA kernel is faster (same rule that already exists for F16 K/V a few lines above). This makes the vector kernel apply only for GQA <= 4.

Ternary-Bonsai-2-27B (GQA 6, head size 256, q4_0 K/V), RTX 4070, 120k context, decode without speculative decoding: 25.3 -> 41.2 tok/s (+63 %). With MTP speculative decoding (n-max 2) 79.7 tok/s at 120k; depth 0 unchanged.

Additional information

  • The branch has two commits: the first contains the environment switches GGML_CUDA_FATTN_GQA_MMA (=0 restores the old choice) and GGML_CUDA_FATTN_VEC_MAXQ (query limit of the vector kernel, default 2) that were used for the measurements below, the last one removes both (the vector kernel keeps its original limit of 2 queries). To reproduce a measurement, build the first commit.
  • test-backend-ops FLASH_ATTN_EXT 2994/2994 passed. I also ran 12 local extra cases (q4_0/q8_0, GQA 6, head size 256, 1-512 queries, kv 1024-16384), 3006/3006; they are not part of this PR.
  • Unlike the F16 branch above, the new rule has no key-count threshold (the F16 rule also requires K->ne[1] >= 8192). The MMA kernel was not slower than the vector kernel in llama-bench pp1/pp2 from 2k keys on, and decode at depth 0 is also slightly faster (+0.62 %), so I did not add one.
  • This is not bit-identical. The two kernels accumulate in a different order, so single-query decode with quantized K/V gives results that differ at rounding level; greedy text can diverge after some tokens (numbers below). With MTP and 3 queries nothing changes (that case already used the MMA kernel).
  • Only measured on one Ada GPU and one model shape.

Test results

  • Hardware / software: RTX 4070 12 GB (AD104, cc 8.9, 504 GB/s, 48 MB L2, 100 KB shared memory per SM), Linux 6.18, NVIDIA driver 615.71, CUDA 13.4, GCC 16.2; Release build, -DGGML_CUDA=ON -DCMAKE_CUDA_ARCHITECTURES=89.
  • Base: speed numbers were measured on prism at 88c4bc60b; the four commits since (SYCL, WebGPU and cuda: fused FWHT quantizer for 64-wide warps (#303)) do not touch this code path. The branch is rebased on 2459f68b5, builds, and test-backend-ops was repeated on it.
  • Model: Ternary-Bonsai-2-27B (PTQ1_0), head size 256, GQA 6 (24 query heads, 4 KV heads), q4_0 K/V cache with mean-centering, one slot, greedy decode.
  • Speed, decode without speculative decoding (1 query), by filled context: 32768 keys 44.8 -> 54.2 tok/s, 120000 keys 25.3 -> 41.2 tok/s (+63 %); depth 0: +0.62 % (95 % CI [+0.56, +0.68] %, 2 pairs of the same binary, GGML_CUDA_FATTN_GQA_MMA=0 against =1). With draft-mtp n-max 1 (2 queries): 32768 keys 53.3 -> 84.0, 120000 keys 23.2 -> 68.5 tok/s. With n-max 2 (3 queries, already on the MMA kernel) and at depth 0 nothing changes.
  • llama-bench pp1/pp2 at depth 2048 to 16384: the MMA kernel is not slower than the vector kernel from 2k keys on.
  • test-backend-ops test -b CUDA0 -o FLASH_ATTN_EXT: 2994/2994 passed on the rebased branch (CUDA0 against CPU). On the earlier base I also ran 12 extra local cases (q4_0/q8_0, GQA 6, head size 256, 1-512 queries, 1024-16384 keys), 3006/3006 on both paths; they are not part of this PR.
  • Numerics (single query, quantized K/V): greedy outputs of the same server with the switch off and on are not identical. Output hashes of a 192-token greedy answer at depth 16384 are equal between the old rule (vector kernel) and this change (MMA kernel with the existing tile size), at depth 32768 they differ. 4 short prompts x 512 greedy tokens at depth 0, old rule (switch =0) against this PR (build without the tile change of the other PR): 3 of 4 texts diverge, the first difference at 13 %, 40 % and 77 % of the text, 1 of 4 is identical. I did not measure a task-level quality metric for the single-query path (perplexity does not exercise it; it runs batches of 512 queries).
  • With MTP (n-max 2, 3 queries) the output hashes are identical with and without this change at depth 16384 and 32768.
  • Not tested: other GPUs, other head sizes or GQA ratios. The branch condition is cc >= GGML_CUDA_CC_ADA_LOVELACE, so the new rule applies to every NVIDIA GPU from cc 8.9 up (Ada, Hopper, Blackwell), not only to Ada; I could only test Ada and can restrict it to cc 8.9 if you prefer. HIP and MUSA are not affected by this condition.

Requirements

  • I have read and agree with the contributing guidelines
  • AI usage disclosure: The patches were developed with Claude Code (Anthropic's coding agent): it wrote the code, the measurement scripts and the first drafts of the commit messages and PR texts. I decided what to work on (which kernels and host paths to optimise, based on profiles of my own decode setup). The measurements and checks listed in the PR texts were run in the Claude Code sessions; I did not re-run them independently. I will maintain the changes. Commits where Claude Code was used carry a Co-Authored-By trailer.

sb32445 and others added 2 commits October 4, 2026 14:21
…ed K/V on Ada

With quantized K/V on Ada the vector kernel was picked for 1 and 2 queries at every
context length. At long context it is much slower than the tensor core kernel. The branch
for F16 K/V already avoids the vector kernel when the GQA ratio is above 4; this applies
the same ratio rule to the quantized branch.

Ternary Bonsai 2 27B PTQ1_0 (head size 256, GQA 6, q4_0 K/V, mean-centered), RTX 4070
12 GB, greedy decode, one slot, tokens per second by filled context:

  depth    no MTP: before -> after    draft-mtp n-max 1: before -> after
  32768    44.8 -> 54.2               53.3 -> 84.0
  120000   25.3 -> 41.2               23.2 -> 68.5

draft-mtp n-max 2 (3 queries, already on the MMA kernel) and depth 0 are unchanged.
llama-bench pp1/pp2 at depth 2048..16384 shows the MMA kernel is not slower from 2k keys on.

test-backend-ops FLASH_ATTN_EXT passes (2994/2994). With 12 extra local cases for q4_0
and q8_0 at GQA 6, head size 256, 1 and 2 queries it passes on both paths (3006/3006).

GGML_CUDA_FATTN_GQA_MMA=0 restores the old choice, GGML_CUDA_FATTN_VEC_MAXQ changes the
query limit of the vector kernel (default 2).

Co-Authored-By: Claude Sonnet 5.5 <noreply@anthropic.com>
GGML_CUDA_FATTN_GQA_MMA and GGML_CUDA_FATTN_VEC_MAXQ of the previous commit were
only there to measure the change. The vector kernel keeps the original limit of
2 queries and is used for GQA <= 4 only.

Co-Authored-By: Claude Sonnet 5.5 <noreply@anthropic.com>
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant