Skip to content

HRX: multi-pass KV-block reduction for flash-attention decode-split (engine#115) - #13

Merged
bong-water-water-bong merged 8 commits into
1bit/hrx-vulkan-patchedfrom
feat/hrx-decode-split-multipass
Sep 26, 2026
Merged

bong-water-water-bong merged 8 commits into
1bit/hrx-vulkan-patchedfrom
feat/hrx-decode-split-multipass

Conversation

@1bit-traffic-bot

Copy link
Copy Markdown

What

Implements the multi-pass / hierarchical KV-block reduction the ggml.flash_attention.decode_split.reduce_fused family was missing, so decode_split stays selected past key_value_token_capacity 2048 instead of falling back to flash_attention_f32_f16_wmma (the -32% decode cliff PR #9 introduced to remove all_rejected). Closes the remaining half of 1bit-MONSTER/engine#115.

Changes

  • loom (flash_attention_decode_split_f32_f16_wmma.loom): widen the compose kernels' %key_value_token_count assume range 1..2048 -> 1..262144; add reduce_completed.multipass (lane-strided over KV blocks, with the per-block normalisation scale written back into partial_max, removing the 32-block scale stage) and reduce_fused.multipass (capacity 2049..262144, priority 5), applied inside the same self-synchronising launch (atomic completion counter) exactly like reduce_fused.cooperative - no second dispatch, no cross-dispatch barrier.
  • dispatch (dispatch-flash-attention.cpp): raise kDecodeSplitMaxKeyValueTokenCapacity 2048 -> 32768.
  • dispatch: align the decode-split partial transients (partial_max / partial_sum / partial_output / q8_output) to 4096 instead of 256. At 256 a transient could end exactly on a page boundary, so a small overrun past its end touched an unmapped page and raised HSA_STATUS_ERROR_MEMORY_FAULT; at 4096 the same overrun stays inside the buffer's own allocated page.

Verification (gfx1151 Strix Halo, Qwen3-Coder-30B-A3B-Instruct-Q4_K_M)

Correctness - 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 HRX0-alone one-character case slip 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:

depth capacity blocks faults
1900 1920 30 0/5
2000 2048 32 0/5
2100 2112 33 0/5 (was 5/5)
3000 3008 47 1/5
4800 4864 76 0/5

<=2048 is unchanged. GGML_HRX_LOG_DISPATCH=1 confirms flash_attention_decode_split_next_q8 is selected above 2048, not the fallback.

Known residual

The 4096 alignment masks rather than fixes a small out-of-bounds access past a partial transient (its signature is the 1/5 still seen at d3000). The likeliest root cause is the very loose assume in the produce - %lane_output_channel ... [range(%lane_output_channel0, 0, 636)] against a channel dimension of only value_head_size = 128 - which stops the compiler proving the vector store in bounds. Tightening those bounds is a sensible follow-up; this PR keeps the alignment change as defence in depth.

agent added 3 commits September 26, 2026 07:04
…rced-cooperative is 0/5 at d2100 and d3000); divergent %lane loop bound is the suspect (engine#115)
… uniform-loop fixes A and B both fail (engine#115)
…er must apply reduce_completed.multipass (the cooperative reducer is invalid above 32 blocks and produced garbage)
@github-actions github-actions Bot added the documentation Improvements or additions to documentation label Sep 26, 2026
… reducer (engine#115)

The multipass reducer wrote the per-block normalization scale back into the
global partial_max transient and re-read it in the output pass. Replace that
round-trip with a recompute in the output pass (scale = exp(block_max - max)
from the intact partial_max) and add the unroll/schedule hints the reference
reduce_f32 output pass carries. Removes the prime suspect for the >2048 silent
corruption + intermittent HSA fault.
@bong-water-water-bong
bong-water-water-bong merged commit 8dd75eb into 1bit/hrx-vulkan-patched Sep 26, 2026
10 of 24 checks passed
@bong-water-water-bong

Copy link
Copy Markdown

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.

@bong-water-water-bong

Copy link
Copy Markdown

Residual >2048 fault root-caused and fixed (final commits b26f194, ca78206)

Two fixes beyond the SGPR one:

  1. 1616087 - 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. b26f194 - 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 b26f194): 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

Labels

documentation Improvements or additions to documentation ggml

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant