hip : skip fully masked KV tiles in WMMA Flash Attention - #28943
moristeven477-ship-it wants to merge 2 commits into
Conversation
Assisted-by: Codex
|
There already is an auxiliary kernel |
This comment was marked as spam.
This comment was marked as spam.
|
According to the llama.cpp AI usage policy:
|
Assisted-by: Claude Opus 5
|
In my config ( |
Evidence: https://github.com/moristeven477-ship-it/llama.cpp/releases/tag/rdna3-masked-kv-v2-20260921 AI usage: Codex implemented the code and tests and prepared the evidence. The contributor supplied the PR body and authorized publication. Assisted-by: Codex
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>
|
As a happy user of the V1 patch, I tested the V2 patch on gfx1201 (Radeon AI PRO R9700, ROCm 10.0.0, Windows), so this is Prompt processing, t/s. Step 0 has the other slots empty, so there is nothing to skip; by step 2 all three slots hold content.
Two runs against the same no-skip build (V2 in one, V1 in the other); its control column agreed between them within 0.5%. Worth noting for this config: the separate Probe cost: step 0 is the case with the other slots empty, where there is least to gain, and V2 is not slower there in either column. Output is unchanged: perplexity is bit-identical and token IDs match over 640 tokens in the deterministic sequential-slot shape. test-backend-ops FLASH_ATTN_EXT passes 4599/4599 including the new mask-class cases. Important note: this tree has no sparse FA (#27970), so I dropped the use_sparse template parameter and the indices arguments when applying the patch. The numbers are for that adaptation, not for the patch verbatim. |
|
The way this is implemented still introduces too much technical debt. This needs a human to do the software design. |
Can you give me some hint or target please? |
|
To make it clear: my impression is that the design in this PR is mostly driven by a language model and significantly more human effort is needed to come up with a better design. This should be done by someone else. |
… merge, and HIP-only code this Vulkan tree never builds
What changed
When server slots share a unified KV cache, the AMD WMMA Flash Attention kernel can still compute interior KV tiles that are fully masked for the current request. Following the review suggestion, this revision moves mask scanning and reduction out of the attention loop and into an auxiliary preprocessing kernel.
The new flash_attn_mask_to_KV_blocks classifies each block as fully masked, all-zero, or general. The attention kernel skips runs of fully masked blocks, omits mask loading and mask addition for all-zero runs, and keeps normal mask evaluation for general blocks. Sixteen two-bit classes pack into a single 32-bit word and are reused across attention heads.
This path is limited to synchronous, dense AMD WMMA. It covers the measured -ub 512 --kv-unified configuration, in which the existing flash_attn_mask_to_KV_max helper is not launched. The patch also includes the broadcast-mask and query-tail fixes plus additional operator tests.
On RX 7900 XTX / Ubuntu 24.04 / ROCm 7.2.1, third-request time drops by approximately 33.3% versus upstream, and by a further 10% versus V1. The gain is primarily in prefill; this is not a general decode or concurrent-server throughput claim.
The model is Qwen3.8-27B-UD-Q4_K_M, fully offloaded, with Flash Attention enabled. The workload uses 98,304 context capacity, 40,960 input / 64 output tokens, and two unified-cache slots requested sequentially in 0/1/0 order. Ubatch is 512, with no prompt cache reuse and no speculative decoding.
Times are geometric means over six restarts per variant, covering all six U/V1/V2 execution orders on the common base c6824a9. Confidence intervals use paired log-ratios. The current source is ported to ec91ab5; the three-arm performance comparison has not been repeated on that base.
Validation
The V2 evidence release contains the patch, commands, hashes, raw data, and logs. The V1 artifacts do not substitute for this evidence.
Credit & inspiration
Thanks to JohannesGaessler for the auxiliary-preprocessing suggestion, to DanoPTT for the V1 testing and the helper-launch diagnosis, and to alankila together with the earlier reports in #28495 and #14924. The masked-block preprocessing idea follows other backends, including Vulkan (#17186); the ec91 port retains ynankani's swizzle work (#28536).