Repository navigation
HRX: multi-pass KV-block reduction for flash-attention decode-split (engine#115) - #13
Conversation
…, OOB search narrowed (engine#115)
…nt size, token count and arena layout; aliasing is the best fit (engine#115)
…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)
… 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.
8dd75eb
into
1bit/hrx-vulkan-patched
Deep-context decode failures were JIT SGPR exhaustion (root-caused + fixed)Following the rejection, I found the real cause of the The compose launch derived Also clamped the output-channel stores to provable bounds ( Verification (build
|
| 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.
Residual >2048 fault root-caused and fixed (final commits b26f194, ca78206)Two fixes beyond the SGPR one:
Controlled split-vs-fallback comparison (llama-bench -p 0 -n 8 -r 3 -dev HRX0)
Depth sweep (committed b26f194): d1470 74, d1900 72, d2000 66, d2100 65, d3000 55, d4800 45, d8192 31 t/s. Repro
All commits are on |
What
Implements the multi-pass / hierarchical KV-block reduction the
ggml.flash_attention.decode_split.reduce_fusedfamily was missing, sodecode_splitstays selected pastkey_value_token_capacity2048 instead of falling back toflash_attention_f32_f16_wmma(the -32% decode cliff PR #9 introduced to removeall_rejected). Closes the remaining half of 1bit-MONSTER/engine#115.Changes
flash_attention_decode_split_f32_f16_wmma.loom): widen the compose kernels'%key_value_token_countassume range1..2048->1..262144; addreduce_completed.multipass(lane-strided over KV blocks, with the per-block normalisation scale written back intopartial_max, removing the 32-block scale stage) andreduce_fused.multipass(capacity 2049..262144, priority 5), applied inside the same self-synchronising launch (atomic completion counter) exactly likereduce_fused.cooperative- no second dispatch, no cross-dispatch barrier.dispatch-flash-attention.cpp): raisekDecodeSplitMaxKeyValueTokenCapacity2048 -> 32768.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 raisedHSA_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:-dev HRX0-dev HRX0,Vulkan0(engine#108 split)ZX-4718-QQexact-dev HRX0,Vulkan0ZX-4718-QQexactThe 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:<=2048 is unchanged.
GGML_HRX_LOG_DISPATCH=1confirmsflash_attention_decode_split_next_q8is 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 onlyvalue_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.