Repository navigation
HRX: flash-attention decode-split template resolution fails ("all_rejected") at longer context (~3800+ tokens) #115
Description
Activity
bong-water-water-bong commented
on Sep 26, 2026 CollaboratorAuthorMore actionsRoot-caused.
ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/flash_attention_decode_split_f32_f16_wmma.loomdefines exactly two implementations of@ggml.flash_attention.decode_split.reduce_fused:template.def<@ggml....reduce_fused> priority(20) ..._direct_f32(...) where [range(key_value_token_capacity, 64, 256)] template.def<@ggml....reduce_fused> priority(10) ..._cooperative_f32(...) where [range(key_value_token_capacity, 257, 2048)]Nothing covers
key_value_token_capacity > 2048. That's exactly why selection reports "all_rejected" atkey_value_token_capacity=4864: every candidate'swhereguard excludes it, not a compile failure in a candidate that was picked.The 2048 ceiling is a real buffer-size limit, not an arbitrary cap.
reduce_fused_cooperative_f32dedicates one subgroup lane to each partial KV block (lane_has_block = lane < active_block_count, guarded to<= 32), and its per-workgroup scratch (scale_stage, a fixed4x32x2xf32alloca<workgroup>) is sized for exactly 32 blocks. 32 blocks x 64 tokens/block = 2048. Even in principle widening this to the wave's full 64 lanes only reaches 4096 - still short of the 4864 that failed here, and short of what any real long-context use (8k/16k/32k) needs. Fixing this properly means a multi-pass/hierarchical reduction over partial blocks, not raising a range guard - genuine kernel-design work, not a wiring bug likeengine#108's.Not attempting that here without checking in first - it's a different scope of work (new kernel implementation) than the bisected fixes in #108. Leaving this issue with the precise cause for whoever picks it up, or let me know if you want me to take a run at the multi-pass design.
🤖 Generated with Claude Code
bong-water-water-bong commented
on Sep 26, 2026 CollaboratorAuthorMore actionsFixed - PR 1bit-MONSTER/llama.cpp#9 (targets
1bit/hrx-vulkan-patched, 1 commit).Rather than the multi-pass reduction kernel I flagged as the "real" fix, I found a much narrower and safer one: the dispatch matcher was offering
decode_splitfor anykey_value_token_countup to 262144, far beyond what the kernel corpus actually implements (nothing above capacity 2048). Capping the matcher at 2048 makes it correctly decline past that point, and the registry's next-priority dispatch for the same op (flash_attention_f32_f16_wmma, already correct for arbitrary KV length, just not split-parallelized) takes over automatically - no new kernel, no cross-backend fallback needed.Verified end to end: the exact failing case (4718 prompt tokens, capacity 4864) now decodes correctly on both
-dev HRX0alone and theengine#108HRX0/Vulkan0split; 1470 tokens (below the boundary, still uses decode-split) is unchanged. Full detail in the PR body.Not done: confirming which dispatch actually serves the long-context case (would need
GGML_HRX_LOG_DISPATCH=1), and decode speed above 2048 tokens (expected slower than decode-split, since the fallback isn't split-parallelized - a genuine multi-pass reduction kernel would recover that, and remains a legitimate follow-up if the speed matters, but is separate scope from this correctness fix).🤖 Generated with Claude Code
bong-water-water-bong commented
on Sep 26, 2026 CollaboratorAuthorMore actionsSpeed cost of the fix, measured (Strix Halo,
Qwen3-Coder-30B-A3B-Instruct-Q4_K_M,-dev HRX0,llama-bench -p 0 -n 8 -d <depth>, 2 runs each):depth tok/s dispatch 1900 69.3 flash_attention_decode_split_next_q8(below the 2048 boundary)2000 68.3 same, right at the edge 2100 46.4 flash_attention_f32_f16_wmma(the fallback, confirmed viaGGML_HRX_LOG_DISPATCH=1) - -32% the instant depth crosses 20483000 40.5 fallback, continuing to decline 4800 32.5 fallback - this is the depth that crashed before this fix There's a sharp, real cliff exactly at the boundary (68 -> 46 tok/s), then a gradual further decline with depth - consistent with the fallback kernel not being split-parallelized across KV blocks (one workgroup processes the whole KV cache serially, instead of one workgroup per 64-token block). 32 tok/s at depth 4800 is perfectly usable (versus crashing before), but it is a real cost for anyone doing long-context decode on
-dev HRX0past 2048 tokens.This is the concrete case for the multi-pass
reduce_fusedkernel I described earlier, if HRX-native long-context decode speed matters enough to invest in it - it would keep decode split-parallelized past 2048 rather than falling back. Leaving that decision and that work to whoever picks this up next; not attempting it without checking in first, it's a different scope (new kernel algorithm) than the two dispatch/matcher fixes landed so far (engine#108, this issue).🤖 Generated with Claude Code
Closing the remaining half of this issue: 1bit-MONSTER/llama.cpp#13 (branch
feat/hrx-decode-split-multipass, one commitf59f6d37aon top offa1a4563).PR #9 removed the
all_rejectedfailure by capping the matcher at capacity 2048, but that handed decode toflash_attention_f32_f16_wmma, which is not split-parallelised - the -32% cliff measured above (68 -> 46 tok/s at d2100). #13 adds the missing multi-pass reduction sodecode_splitstays selected instead.Disposition of the "Not yet done" list
- Narrowing the exact
key_value_token_countthreshold - ADDRESSED. The threshold is 2048.match.key_value_capacity = ceil_div(mask->ne[0], 64) * 64, so the split covers capacities 64..262144 through three reducers:direct_f32(64..256),cooperative_f32(257..2048) and the newmultipass_f32(2049..262144). 2048 is exactly 32 producer blocks of 64 tokens. - Reading
ggml.flash_attention.decode_split.reduce_fused's template definitions - ADDRESSED. They were read: the family had onlydirect_f32andcooperative_f32, andcooperativeis limited to 32 blocks because it dedicates one subgroup lane to each partial KV block (lane_has_block = lane < active_block_count) and sizes its normalisation scratch4x32x2xf32for 32 blocks. The new multipass reducer lane-strides over the block dimension instead, so any block count works. - Other models / GQA configurations, KV cache types (q8_0 / q4_0), or whether
-fa offavoids it - DISPOSITIONED, not addressed here. All verification isQwen3-Coder-30B-A3B-Instruct-Q4_K_Mwith the default f16 KV cache (32 query heads / 4 KV heads / head_dim 128 GQA). The reducer is written against the config values (query_head_count,key_value_head_count,qk_head_size,value_head_size) rather than the model, but other GQA ratios and quantised KV caches have NOT been exercised.-fa offstill avoids the path entirely - a workaround, not a fix.
Verification (gfx1151 Strix Halo)
Buried code word
ZX-4718-QQ,llama-server, temperature 0, seed 42,cache_prompt:false:case device result GPU faults 4718-token repro (capacity 4864) -dev HRX0code word retrieved (one char case slip) 0 4718-token repro (capacity 4864) -dev HRX0,Vulkan0(engine#108 split)ZX-4718-QQexact0 ~3587-token case (3569 prompt tokens) -dev HRX0,Vulkan0ZX-4718-QQexact0 The fallback oracle (split disabled) on the same 4700-token case returns
ZX-4718QQ. The quarterly logistics review covers ...- it also misses the buried word - so the one-character case slip on-dev HRX0alone is a model/sampling artifact, not a regression.Cliff removed -
llama-bench -p 0 -n 8 -d D, five runs each, count of GPU faults: 1900 0/5, 2000 0/5, 2100 0/5 (was 5/5), 3000 1/5, 4800 0/5. <=2048 unchanged.GGML_HRX_LOG_DISPATCH=1confirmsflash_attention_decode_split_next_q8is selected above 2048, not the fallback.Residual: the 4096-alignment of the partial transients masks a small out-of-bounds access rather than fixing it; its signature is the 1/5 still seen at d3000. The follow-up (tighten the loose
range(%lane_output_channel0, 0, 636)assume against a 128-wide channel dim) is described in #13.- Narrowing the exact
bong-water-water-bong commented
on Sep 26, 2026 CollaboratorAuthorMore actionsDeep-context decode failures were JIT SGPR exhaustion (root-caused + fixed)
Following the rejection, I found the real cause of the
res = -3failures at deep contexts: a JIT register-allocation failure, not a memory fault:failed to allocate amdgpu.sgpr registers for @ggml_flash_attention_decode_split_f32_f16_wmma_next_q8 budget 106, peak 328 (d4800) / 552 (d8192), failure code spill-traffic-register-exhausted AMDGPU HSACO emission produced no executable bytesThe compose launch derived
producer_block_countfrom the compile-timekey_value_token_capacity, so the JIT constant-folded it and fully unrolled thereduce_completedblock loops (~4.3 SGPRs/block). Deriving the count from the runtime token count (provably equal: the dispatch setskey_value_capacity = ceil_div(key_value_token_count, 64)*64) keeps the loops rolled. Commit16160878b.Also clamped the output-channel stores to provable bounds (
092e444b3), addressing the loose assumes the auditor flagged.Verification (build
092e444b3,-dev HRX0,llama-bench -p 0 -n 8 -r 3)depth tok/s (faults) 1470 ~74 (0/3) 1900 72.5 / 72.3 / 72.9 (0/3) 2000 70.7 / 71.5 / 70.9 (0/3) 2100 64.7 / 64.6 (0/3) 3000 54.5 / 51.1 (1/5) 4800 44.7 / 44.0 / 44.6 (0/3) 8192 31.2 / 31.4 (0/2) - Cliff gone: d2000 70.7 -> d2100 64.7 (~9%); d4800 44.7 t/s (previously a JIT failure). Split d2100 64.5 vs fallback 48.8.
GGML_HRX_LOG_DISPATCH=1:common.flash_attention_decode_split_next_q8 (single)selected.- Correctness (
-c 4864 -np 1):-dev HRX04700 tok ->ZX-4718-QQ; 3569 tok ->ZX-4718-QQQ(model one-char noise).HRX0/Vulkan0split 4700 ->ZX-4718-QQ; 3569 ->ZX-4718-QQ. - Residual: a rare
HSA_STATUS_ERROR_MEMORY_FAULT(d3000 1/5), largely ambient - at d4800 the fallback faults 1/6 vs the split 2/6 under the same shared GPU.
bong-water-water-bong commented
on Sep 27, 2026 CollaboratorAuthorMore actionsResidual >2048 fault root-caused and fixed (final commits b26f194e3, ca7820650)
Two fixes beyond the SGPR one:
- 16160878b - deep-context
res=-3was JIT SGPR exhaustion (constant-folded block count -> full unroll -> peak 328/552 vs budget 106). Deriving the loop bound from the runtime token count fixed it. - b26f194e3 - the residual >2048 fault was that the same change made the partial-view STRIDE runtime. Keeping the stride compile-time (constant capacity-derived count) while only the reduce LOOP BOUND is the runtime count removes it.
Controlled split-vs-fallback comparison (llama-bench -p 0 -n 8 -r 3 -dev HRX0)
depth split faults fallback faults 2100 0/3 0/3 3000 0/3 .. 2/5 (noisy, shared GPU) 0/3 .. 0/5 4800 0/5 0/5 8192 0/5 0/5 Depth sweep (committed b26f194e3): d1470 74, d1900 72, d2000 66, d2100 65, d3000 55, d4800 45, d8192 31 t/s.
Repro
HRX0/Vulkan0split (server -c 4864 -np 1): 4700 tok ->ZX-4718-QQ, 3569 tok ->ZX-4718-QQ(4/4 exact).-dev HRX0alone: greedy output renders the code word asZX-4718QQ/ZX-4718-QQQ(the oracle itself varies the hyphen/case); the HRX0-alone server additionally hits an independent HRX runtime bug during batched prefill (HRX host staging buffer copy failed ... command buffer is not in a recording state, hrx-main/runtime), unrelated to this kernel.
All commits are on
origin/feat/hrx-decode-split-multipass. NOTE: the build worktree~/wt/hrx-fa-multipass3was deleted from the host mid-session.- 16160878b - deep-context
Found while long-context-verifying
1bit-MONSTER/engine#108's fix (1bit-MONSTER/llama.cppPR #7, commitb4d3ec9, HRX+Vulkan build on Strix Halo,gfx1151).Repro
llama-server -m Qwen3-Coder-30B-A3B-Instruct-Q4_K_M.gguf -dev HRX0 -ngl 99 -c 8192(no cross-device split; reproduces the same on-dev HRX0,Vulkan0 -ts 1,0 -ot exps=Vulkan0too), a single chat completion with a prompt around 3800 tokens (a ~3000-word synthetic document). Decode fails:What's confirmed
key_value_token_countpresumably under whatever triggers the split), fails somewhere between there and ~3800 tokens (key_value_token_count=4864at the point of failure - the KV capacity is padded/rounded, so the true threshold is a bit lower than 4864).engine#108): reproduces identically with-dev HRX0alone, no-ts/-otat all.Qwen3-Coder-30B-A3B-Instruct-Q4_K_Mso far, dense-attention (32 query heads / 4 KV heads / head_dim 128, GQA),f16KV cache (server default).ggml.flash_attention.decode_split.reduce_fusedfamily at this parameter combination, not a compile error in one candidate.Not yet done
key_value_token_countthreshold.-fa offavoids it entirely (it would, but that's a workaround not a fix).ggml.flash_attention.decode_split.reduce_fused's template definitions in the Loom kernel corpus to see which parameter range they cover and why this one falls outside it.🤖 Generated with Claude Code