Skip to content

cuda: smaller KV tile for the 8-column MMA flash attention config (head size 256) - #308

Open
sb32445 wants to merge 1 commit into
PrismML-Eng:prismfrom
sb32445:pr/fattn-mma-tile
Open

sb32445 wants to merge 1 commit into
PrismML-Eng:prismfrom
sb32445:pr/fattn-mma-tile

Conversation

@sb32445

@sb32445 sb32445 commented Oct 4, 2026 •

Copy link
Copy Markdown

Overview

For head size 256 with 8 columns (1 query, GQA 6), the Ampere/Ada MMA config uses nbatch_fa 64. The shared memory then allows only 2 blocks per SM (92 blocks on a 4070, 58 % of DRAM bandwidth at 120k keys). With nbatch_fa 32, 5 blocks per SM fit (228 blocks, 84 % DRAM).

RTX 4070, q4_0 K/V, 1 query, per layer: 32k keys 122 -> 87 us; 120k keys about 540 -> 370 us. Model, decode without speculative decoding (Ternary-Bonsai-2-27B): 32k 54.4 -> 56.4 tok/s, 120k 41.2 -> 46.0 tok/s. No change for 2-3 queries (other config) and none at depth 0.

Additional information

  • One-number change in ggml_cuda_fattn_mma_get_config_ampere. test-backend-ops FLASH_ATTN_EXT 3006/3006 (3006 includes 12 local extra cases, not in this PR).
  • Tried and dropped: nstages 1 (slower), smaller K/V tiles (not implemented with 2 stages), 64 threads for the 32-column config (within noise).
  • Builds on the kernel choice of cuda: use the MMA flash attention kernel for GQA above 4 with quantized K/V on Ada #307 (GQA/MMA PR, separate), but is independent of it in the code. Without cuda: use the MMA flash attention kernel for GQA above 4 with quantized K/V on Ada #307, quantized K/V with GQA 6 takes the vector kernel for 1 query, so this change has no effect there.
  • Not bit-identical. nbatch_fa is the number of KV rows per softmax rescaling, so a different value changes how the online softmax is grouped and rounded (the comment above the config says it should only affect speed, which holds up to rounding). Numbers below.
  • Ampere: this config function is shared by cc 8.0 and 8.9. I measured only on an RTX 4070 (Ada, 100 KB shared memory per SM); a GPU with more shared memory (e.g. A100) may fit more blocks with 64 and I have not checked that 32 is not slower there.
  • One run per cell, one GPU.

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, q4_0 K/V cache with mean-centering, one slot, greedy decode, no speculative decoding (1 query).
  • Per layer, one query, q4_0 K/V: 32k keys 122 -> 87 us, 120k keys about 540 -> 370 us; blocks per SM 2 -> 5, DRAM bandwidth at 120k keys 58 % -> 84 % (ncu).
  • Decode without speculative decoding, one run per cell: 32k keys 54.4 -> 56.4 tok/s, 120k keys 41.2 -> 46.0 tok/s (+12 %). No change with MTP n-max 2 (3 queries use a different config) and none at depth 0.
  • test-backend-ops test -b CUDA0 -o FLASH_ATTN_EXT: 2994/2994 passed on the rebased branch (CUDA0 against CPU); on the earlier base 3006/3006 including 12 local extra cases that are not part of this PR.
  • Numerics: the same server before and after (patches up to the GQA/MMA PR against all patches) gives different greedy text for 3 of 4 short prompts (4 prompts x 512 tokens, depth 0, first difference at characters 1257, 525 and 368; 1 prompt identical). Output hashes at depth 16384 differ, at 32768 they are equal. The "after" side also contains my other patches, which gave identical outputs in all my A/B runs, so I attribute the difference to this change but did not isolate it with a switch (there is none). I did not measure a task-level quality metric for the single-query path.
  • Not tested: other GPUs (see Ampere note), other head sizes (only the 256/256 8-column case changes), HIP and MUSA.

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.

…Ampere/Ada (head size 256)

With nbatch_fa 64 the kernel flash_attn_ext_f16<256,256,1,8> (one query, GQA 6)
fits only 2 blocks per SM because of shared memory and reaches 58% of the DRAM
bandwidth at 120k keys. nbatch_fa 32 fits 5 blocks per SM (84% DRAM).

RTX 4070, Bonsai 2 27B, q4_0 K/V, one query, per layer: 32k keys 122 -> 87 us,
120k keys about 540 -> 370 us. Decode without speculative decoding: 32k
54.4 -> 56.4 tok/s, 120k 41.2 -> 46.0 tok/s. No change with 2-3 queries
(different config), no change at depth 0.

test-backend-ops FLASH_ATTN_EXT: 3006/3006 passed.

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