Skip to content

ggml-hrx: attention sinks after FlashAttention (gpt-oss fully on HRX: 1 graph split) - #68

Merged
bong-water-water-bong merged 2 commits into
1bit/hrx-vulkan-patchedfrom
1bit/hrx-attention-sinks
Oct 2, 2026
Merged

bong-water-water-bong merged 2 commits into
1bit/hrx-vulkan-patchedfrom
1bit/hrx-attention-sinks

Conversation

@bong-water-water-bong

Copy link
Copy Markdown

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 of exp(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.
  • A fully masked row (S = 0) is written as zeros, never NaN.
  • 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):

  • The prefill matcher accepts the 5-input node.
  • match_flash_attention_f32_f16_dispatch appends the rescale after its own FA dispatch, in the same match, so it runs after FA.
  • The gate-fused variant and the decode-split next_q8 variant keep the 4-input form. next_q8 publishes a Q8 copy of the output for the next projection, and that copy would miss the rescale.
  • gpt-oss decode with sinks therefore uses the general FA kernel.
  • ALiBi and logit softcap stay declined, as the shared matcher already requires.

Gates

  • Codegen with sinks off: no AMD .loom file in the diff (flash_attention_f32_f16_wmma.loom and flash_attention_decode_split_f32_f16_wmma.loom are byte-identical). The 4-input path dispatches exactly as before.
  • Logits: compared against this branch's parent build (ggml-hrx: ADD_ID and SWIGLU_OAI on HRX (gpt-oss: 193 -> 49 graph splits) #67), wikitext c512, 4 chunks, one sequence per batch, HRX0:
    • Qwen3.8-27B UD-Q4_K_XL: before vs after mean KLD 0.000000, max 0.000048, same top 100%. After vs after gives the identical numbers, so that max is the run-to-run floor.
    • ZAYA1-8B Q4_K_M: mean 0.000000, max 0.000058, same top 100%, again identical to after vs after.
    • The 27B with 4 sequences per batch fails on both builds (graph_compute −1, the known HRX multi-sequence gap), hence one sequence per batch.
  • Speed, interleaved, balanced power mode, not under thermal-run:
before (#67) after
Qwen3.8-27B pp512 / tg128 347–351 / 10.5 345–350 / 10.5
ZAYA1-8B pp512 / tg128 1636–1658 (±216–297) / 90.8–91.3 1656–1926 (±181–300) / 90.8–93.3
  • check.cases (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:
    • prefill: 3 tokens × 300 keys;
    • decode: 1 × 1000;
    • a fully masked row, which gives zeros.
  • test-backend-ops: HRX reports every FLASH_ATTN_EXT case 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

@github-actions github-actions Bot added the ggml label Oct 2, 2026
@bong-water-water-bong

Copy link
Copy Markdown
Author

Review (PR-Agent duty): approve.

  • The method: an exact post-correction. O_sink = O_fa · S/(S + e^(sink−M)) = O_fa · σ(LSE − sink). AMD's FlashAttention kernels are byte-identical; the hook is about 11 lines in dispatch-flash-attention.cpp.
    • The base FA matcher accepts a 5-input node only when attention_sinks_supported holds. The gate variant still requires 4 inputs.
    • Our sink dispatch is appended after the FA dispatch in the same match, so it's ordered after FA's write.
  • What the sink path checks:
    • sinks F32 contiguous, one per query head;
    • mask F16 with rows key_count apart;
    • the scale is passed through.
    • ALiBi and softcap are already declined by the FA matcher (max_bias == 0 && logit_softcap == 0).
  • Gates:
    • Qwen3.8-27B and ZAYA1-8B logits are identical to the ggml-hrx: ADD_ID and SWIGLU_OAI on HRX (gpt-oss: 193 -> 49 graph splits) #67 build, within the run-to-run floor (max KLD 4.8e-5 / 5.8e-5), and pp/tg are within noise.
    • check.cases cover prefill, decode and fully masked (zeros), with 8:2 GQA and head 64, sanitizer clean.
    • Full suite 1061/1061.
    • test-backend-ops has no HRX FA coverage, the same as on the parent, so the check.cases and model logits are the evidence.
  • gpt-oss-20b: graph splits 49 -> 1, pp512 839, tg128 32.3, text correct. KLD 0.029 is tracked in engine bugfix: centos 7, gcc (GCC) 11.2.1 20220127 (Red Hat 11.2.1-9) ggml-org/llama.cpp#284.

Merging after #67, then retargeting.

@bong-water-water-bong
bong-water-water-bong changed the base branch from 1bit/hrx-gptoss-ops to 1bit/hrx-vulkan-patched October 2, 2026 13:03
}
// 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) {

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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);
}

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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.

@bong-water-water-bong

Copy link
Copy Markdown
Author

Review fix (aliasing guard): 4a0c1d0.

  • The in-place check in dispatch-attention-sink.cpp compared Value::id, so a different value (view) of the same allocation passed as disjoint. It now refuses when an input shares storage with the output (ValueMap::same_storage) and their byte ranges (storage_offset .. + byte_count) intersect.
  • New tests/test-hrx-attention-sink.cpp: a 5-input FLASH_ATTN_EXT whose output is aliased onto the query's storage, a distinct value with the same storage, must be refused, and a disjoint one accepted. It fails with the old id-only check (requirement failed: !append_for(true)) and passes with the fix.
  • No change to the dispatch path in normal graphs; FA's output is a fresh tensor there.

🤖 Generated with Claude Code

@bong-water-water-bong

Copy link
Copy Markdown
Author

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.

bong-water-water-bong and others added 2 commits October 2, 2026 14:51
… 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>
@bong-water-water-bong

Copy link
Copy Markdown
Author

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):

  • test-hrx-attention-sink: pass (disjoint output accepted, output aliasing the query refused)
  • test-hrx-fa-masked-v: 27 cases passed
  • test-hrx-moe-split: all cases OK
  • gpt-oss-20b MXFP4, KLD vs CPU (8 x 512): 0.0284, same top 87.7 % (0.0294 without sinks on 388290b)
  • gpt-oss-20b greedy x3 (sky / primary colors / haiku): correct and coherent, same text as before
  • Qwen3.8-27B UD-Q4_K_XL greedy x3: byte-identical to 388290b

@bong-water-water-bong
bong-water-water-bong merged commit 4485916 into 1bit/hrx-vulkan-patched Oct 2, 2026
10 of 24 checks passed
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant