Repository navigation
ggml-hrx: attention sinks after FlashAttention (gpt-oss fully on HRX: 1 graph split) - #68
Conversation
|
Review (PR-Agent duty): approve.
Merging after #67, then retargeting. |
| } | ||
| // The output is rewritten in place after FlashAttention wrote it; no input may share its storage. | ||
| for (const Value * input : { query, key, mask, sinks }) { | ||
| if (input->id == output->id) { |
There was a problem hiding this comment.
Dead aliasing guard: compares id instead of storage_root
This loop is meant to implement the comment above it and the PR description's claim ("It refuses if any input shares storage with the output it rewrites"), but Value::id is each graph value's own unique identity — a freshly-produced output and any of its distinct input values will always have different ids by construction, so input->id == output->id can never be true for any real graph. The check is unreachable dead code that provides zero protection.
Every other aliasing guard in this backend for exactly this purpose compares storage_root instead, e.g. this same file's dispatch-flash-attention.cpp:
if (output->storage_root == match.query->storage_root || output->storage_root == match.key->storage_root ||
output->storage_root == match.value->storage_root || output->storage_root == match.mask->storage_root ||
output->storage_root == raw_gate->storage_root) {
return false;
}and dispatch-rmsnorm.cpp's pairwise_distinct_storage_roots. This looks like the intended comparison here was dropped/typo'd from storage_root to id.
Failure scenario: if ggml's allocator ever places this FLASH_ATTN_EXT node's output as a view/alias of one of query/key/mask/sinks (the exact case the comment is guarding against), this check will not catch it, and the row-by-row in-place rescale kernel will read already-overwritten or stale data instead of being declined.
|
|
||
| std::string index_config(int64_t value) { | ||
| return std::to_string(value); | ||
| } |
There was a problem hiding this comment.
Reimplements an existing helper already included in this file
index_config(int64_t) just does std::to_string(value) — identical to common_to_config_value(int64_t) declared in dispatch-mul-mat-common.h, which this file already #includes and already uses two lines below for the float overload (common_to_config_value(params->scale) at line 101). The four index_config(...) calls a few lines down (query_head_count, key_value_head_count, qk_head_size, value_head_size) could just call common_to_config_value directly, and this local helper can be deleted.
|
Review fix (aliasing guard): 4a0c1d0.
🤖 Generated with Claude Code |
4a0c1d0 to
17a73f8
Compare
|
Rebased onto #73 (guard head ba0be9a, on 388290b). The CMakeLists conflicts were list appends only. On this tree: test-hrx-attention-sink passes; gpt-oss-20b greedy x3 (sky / primary colors / haiku) gives the same text as without sinks; llama-bench pp512 990 ± 22, tg128 35.0 ± 5.5. #72 (FA masked V) is expected to merge first, which will need one more rebase: tests/CMakeLists.txt conflicts (both add a test block), plus a rerun. |
… FlashAttention A sink adds exp(sink - M) to the softmax denominator only, so with M and S the row maximum and sum of exp(scale q.k + mask - M): output_with_sink = output_without_sink * S / (S + exp(sink - M)). FlashAttention runs unchanged; ops/attention_sink_f32.loom + dispatch_registration/common/dispatch-attention-sink.cpp (ours) recompute M and S per (token, head) row with its conventions and rescale its output in place (a fully masked row is written as zeros). AMD kernels untouched; in dispatch-flash-attention.cpp the prefill matcher accepts the 5-input node and its dispatch appends the rescale, while the gate-fused and decode-split next_q8 variants keep the 4-input form (the latter publishes a Q8 copy of the output for the next projection). check.cases (prefill 3 x 300 keys, decode 1 x 1000, fully masked) vs a full-attention reference with and without the sink, 8 query heads over 2 KV heads, --sanitizer=access. gpt-oss-20b: graph splits 49 -> 1. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
…ut's storage The in-place guard compared value ids, so a different value (view) of the same allocation passed as disjoint (review). Compare storage (ValueMap::same_storage) and the byte ranges. tests/test-hrx-attention-sink.cpp: a 5-input FLASH_ATTN_EXT whose output is aliased onto the query's storage must be refused (fails with the old check); a disjoint one is accepted. Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com>
|
Rebased onto the fork tip 2bd7f58 (#73 guard + #72 FA masked-V). The only conflict was in tests/CMakeLists.txt, where both branches append a test block; all blocks are kept. Head is 5405484. Results on this tree (strixhalo):
|
17a73f8 to
5405484
Compare
4485916
into
1bit/hrx-vulkan-patched
Stacked on #67, which is stacked on #66; it retargets as those merge. gpt-oss's attention sinks now run on HRX, the last op it left on the CPU: graph splits 49 → 1.
Approach. AMD's FlashAttention kernels are unchanged. A sink adds
exp(sink − M)to the softmax denominator and nothing to the numerator. So with M and S the row max and sum ofexp(scale·q·k + mask − M):output_with_sink = output_without_sink · S / (S + exp(sink − M))This is exact.
ops/attention_sink_f32.loom(ours) recomputes M and S per (token, head) row with FlashAttention's conventions: the same scale, an additive mask, and −inf entries skipped. It then rescales FA's output in place.dispatch_registration/common/dispatch-attention-sink.cpp(ours) checks that the fifth input is F32[query_heads]sinks with the expected mask row layout. It refuses if any input shares storage with the output it rewrites.Hook in AMD's
dispatch-flash-attention.cpp(about 11 lines):match_flash_attention_f32_f16_dispatchappends the rescale after its own FA dispatch, in the same match, so it runs after FA.next_q8variant keep the 4-input form.next_q8publishes a Q8 copy of the output for the next projection, and that copy would miss the rescale.Gates
.loomfile in the diff (flash_attention_f32_f16_wmma.loomandflash_attention_decode_split_f32_f16_wmma.loomare byte-identical). The 4-input path dispatches exactly as before.graph_compute−1, the known HRX multi-sequence gap), hence one sequence per batch.iree-test-loom --sanitizer=access): 8 query heads over 2 KV heads (gpt-oss's 64:8 grouping in small), head size 64, scale 0.125. Each case compares the reference without a sink, then the rescale, against the reference with the sink:FLASH_ATTN_EXTcase as not supported, with or without sinks (5141/5141 on the parent too), so it has no coverage here. Full suite 1061/1061.gpt-oss-20b (MXFP4, HRX0):
Cost: the rescale recomputes q·k once per (token, head) row over the unmasked keys. In this measurement it's outweighed by removing the CPU round trips; at long context it grows with the number of keys.
🤖 Generated with Claude Code