Skip to content

HRX: flash-attention decode-split template resolution fails ("all_rejected") at longer context (~3800+ tokens) #115

Description

@bong-water-water-bong

Found while long-context-verifying 1bit-MONSTER/engine#108's fix (1bit-MONSTER/llama.cpp PR #7, commit b4d3ec9, 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=Vulkan0 too), a single chat completion with a prompt around 3800 tokens (a ~3000-word synthetic document). Decode fails:

HRX Loom JIT Loom compilation failed failed: Loom compilation failed
  diagnostic[0] LOWERING/045: select-templates cannot resolve template.apply against template family <ggml.flash_attention.decode_split.reduce_fused>: all_rejected
compile_kernel: compile gfx1151|...|ggml_flash_attention_decode_split_f32_f16_wmma_next_q8|recipe=direct|key_value_token_count=4864|ggml.flash_attention.attention_scale=0.0883883461|ggml.flash_attention.decode.key_value_token_capacity=4864|ggml.flash_attention.key_value_head_count=4|ggml.flash_attention.qk_head_size=128|ggml.flash_attention.query_head_count=32|ggml.flash_attention.value_head_size=128: Loom compilation failed
graph_compute: failed to prepare command 3 kind=Kernel kernel_id=... bindings=10
llama_decode: failed to decode, ret = -3
srv        decode: Compute error. off = 0, n_batch = 2048, ret = -3

What's confirmed

  • Works at 1470 prompt tokens (key_value_token_count presumably under whatever triggers the split), fails somewhere between there and ~3800 tokens (key_value_token_count=4864 at the point of failure - the KV capacity is padded/rounded, so the true threshold is a bit lower than 4864).
  • Not related to the cross-device split or the MoE router (engine#108): reproduces identically with -dev HRX0 alone, no -ts/-ot at all.
  • Not model-specific in any narrow sense that's been checked - only tried on Qwen3-Coder-30B-A3B-Instruct-Q4_K_M so far, dense-attention (32 query heads / 4 KV heads / head_dim 128, GQA), f16 KV cache (server default).
  • Looks like a genuine coverage gap: "all_rejected" means the template-selection pass found zero candidate implementations for the ggml.flash_attention.decode_split.reduce_fused family at this parameter combination, not a compile error in one candidate.

Not yet done

  • Narrowing the exact key_value_token_count threshold.
  • Checking other models/GQA configurations, KV cache types (q8_0/q4_0 KV), or whether -fa off avoids it entirely (it would, but that's a workaround not a fix).
  • Reading 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

Activity

  1. bong-water-water-bong commented on Sep 26, 2026

    @bong-water-water-bong
    CollaboratorAuthor

    Root-caused. ggml/src/ggml-hrx/kernel-corpus/kernels/loom-libs/ops/flash_attention_decode_split_f32_f16_wmma.loom defines 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" at key_value_token_capacity=4864: every candidate's where guard 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_f32 dedicates 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 fixed 4x32x2xf32 alloca<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 like engine#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

  2. bong-water-water-bong commented on Sep 26, 2026

    @bong-water-water-bong
    CollaboratorAuthor

    Fixed - 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_split for any key_value_token_count up 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 HRX0 alone and the engine#108 HRX0/Vulkan0 split; 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

  3. bong-water-water-bong commented on Sep 26, 2026

    @bong-water-water-bong
    CollaboratorAuthor

    Speed 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 via GGML_HRX_LOG_DISPATCH=1) - -32% the instant depth crosses 2048
    3000 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 HRX0 past 2048 tokens.

    This is the concrete case for the multi-pass reduce_fused kernel 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

  4. 1bit-traffic-bot commented on Sep 26, 2026

    @1bit-traffic-bot

    Closing the remaining half of this issue: 1bit-MONSTER/llama.cpp#13 (branch feat/hrx-decode-split-multipass, one commit f59f6d37a on top of fa1a4563).

    PR #9 removed the all_rejected failure by capping the matcher at capacity 2048, but that handed decode to flash_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 so decode_split stays selected instead.

    Disposition of the "Not yet done" list

    • Narrowing the exact key_value_token_count threshold - 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 new multipass_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 only direct_f32 and cooperative_f32, and cooperative is 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 scratch 4x32x2xf32 for 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 off avoids it - DISPOSITIONED, not addressed here. All verification is Qwen3-Coder-30B-A3B-Instruct-Q4_K_M with 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 off still 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 HRX0 code word retrieved (one char case slip) 0
    4718-token repro (capacity 4864) -dev HRX0,Vulkan0 (engine#108 split) ZX-4718-QQ exact 0
    ~3587-token case (3569 prompt tokens) -dev HRX0,Vulkan0 ZX-4718-QQ exact 0

    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 HRX0 alone 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=1 confirms flash_attention_decode_split_next_q8 is 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.

  5. bong-water-water-bong commented on Sep 26, 2026

    @bong-water-water-bong
    CollaboratorAuthor

    Deep-context decode failures were JIT SGPR exhaustion (root-caused + fixed)

    Following the rejection, I found the real cause of the res = -3 failures 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 bytes
    

    The compose launch derived producer_block_count from the compile-time key_value_token_capacity, so the JIT constant-folded it and fully unrolled the reduce_completed block loops (~4.3 SGPRs/block). Deriving the count from the runtime token count (provably equal: the dispatch sets key_value_capacity = ceil_div(key_value_token_count, 64)*64) keeps the loops rolled. Commit 16160878b.

    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 HRX0 4700 tok -> ZX-4718-QQ; 3569 tok -> ZX-4718-QQQ (model one-char noise). HRX0/Vulkan0 split 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.
  6. bong-water-water-bong commented on Sep 27, 2026

    @bong-water-water-bong
    CollaboratorAuthor

    Residual >2048 fault root-caused and fixed (final commits b26f194e3, ca7820650)

    Two fixes beyond the SGPR one:

    1. 16160878b - deep-context res=-3 was 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.
    2. 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/Vulkan0 split (server -c 4864 -np 1): 4700 tok -> ZX-4718-QQ, 3569 tok -> ZX-4718-QQ (4/4 exact).
    • -dev HRX0 alone: greedy output renders the code word as ZX-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-multipass3 was deleted from the host mid-session.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

No one assigned

    Labels

    No labels
    No labels

    Type

    No type

    Projects

    No projects

      Milestone

      No milestone

      Relationships

      None yet

      Development

      No branches or pull requests

      Issue actions