Skip to content

[TosaToRock] Broadcast lastValidKVIndex over the attention groups without Q's collapse - #557

Open
pfultz2 wants to merge 4 commits into
developfrom
fix/attention-kvcache-index-plain-q
Open

pfultz2 wants to merge 4 commits into
developfrom
fix/attention-kvcache-index-plain-q

Conversation

@pfultz2

@pfultz2 pfultz2 commented Oct 1, 2026 •

Copy link
Copy Markdown

What

TosaToRock's attention matcher rebuilt the per-group broadcast of lastValidKVIndex (and prefixOffset) from the tensor.collapse_shape feeding Q into the first matmul. When Q is a plain kernel argument, the argument's expand_shape and the matmul's collapse_shape fold into a single expand_shape, addBroadcastForBlockArg fails, and prepareBlockArgTensor silently passes the raw [B] tensor to rock.attention with gemmG = B * heads groups. The kernel loads the index at g_block, so every group after the first reads past the end of the buffer and produces wrong attention outputs. The op verifier would reject this (Batch dimensions do not match (lastValidKVIndex and Output)), but the MIGraphX C API pipeline runs with verification disabled in release builds, so nothing catches it.

This change derives the broadcast from the attention batch instead of Q's producer:

  • addBroadcastForBlockArg now takes gemmG and broadcasts a [B]/[B, 1] block argument to [B, gemmG / B]. The groups are batch-major, so collapsing that indexes it by group exactly as the previous [batch, heads(, splitKV)] shape did.
  • match() rejects index tensors whose leading size does not divide gemmG (or whose [B, H] product does not equal it), since the rewrite cannot fail once it starts. The only inputs this newly rejects are ones that previously produced the out-of-bounds kernel.

Why

MIGraphX hits this with its test_gpu_mlir kv-cache attention tests, where Q is a plain module parameter: kv_cache_attention_seq_len_mask, kv_cache_attention_seq_len_mask_batched, and kv_cache_attention_sinks_lse fail verification against rocmlirTriton (RMS error ~0.3, head 0 correct, heads 1-3 garbage). Real models usually give Q a transpose or reshape producer inside the fused module and are unaffected, but a model whose Q producer has another consumer would silently get wrong results.

Testing

  • New lit test tosa-to-rock-attention-kvcache-plain-q.mlir with a plain-argument Q, single-batch and two-batch cases. It fails without this change (lastValidKVIndex = (%arg3 : tensor<1x1xi32>) passed through) and passes with it.
  • All 15 tosa-to-rock-attention-*.mlir tests pass.
  • New pr-e2e test mixr-attention-kvcache-plain-q.mlir: the GQA decode module MIGraphX emits, with a plain-argument Q and one last-valid index per batch for two batches, numerically verified against the clone CPU reference on gfx1201 (random indices in range, and a fixed index past the end). Without this change rocmlir-driver, which runs with the verifier on, rejects the lowering with Batch dimensions do not match (lastValidKVIndex and Output); through MIGraphX's release pipeline, where the verifier is off, the same module instead produces wrong results for heads 1-3.
  • Linked a MIGraphX build against the patched librockCompiler.a: all 31 test_gpu_mlir tests pass, including the three kv-cache tests above, on gfx1201.

Assisted-by: Claude Code

🤖 Generated with Claude Code

https://claude.ai/code/session_01EtVq7zJgAEjLYv6dGRC9QB

…hout Q's collapse

addBroadcastForBlockArg rebuilt the per-group broadcast of lastValidKVIndex
and prefixOffset from the tensor.collapse_shape feeding Q into the first
matmul. When Q is a plain kernel argument that collapse folds into the
argument's expand_shape, the broadcast fails, and prepareBlockArgTensor
silently hands rock.attention the raw [B] tensor for gemmG = B * heads
groups. The kernel indexes it by g_block, so every group after the first
reads out of bounds. The op verifier would catch this, but the MIGraphX C
API pipeline runs with verification off in release builds.

Derive the broadcast from the attention batch instead: broadcast a [B] or
[B, 1] block argument to [B, gemmG / B], which collapses to one entry per
group in the same batch-major order as the old [batch, heads(, splitKV)]
shape. Reject index tensors that cannot broadcast evenly across the groups
in match(), since the rewrite cannot fail once it starts.

Adds a lit test with a plain-argument Q that fails without this change.
MIGraphX's test_gpu_mlir kv-cache attention tests, which hit this, pass
against the patched library.

Assisted-by: Claude Code
Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01EtVq7zJgAEjLYv6dGRC9QB

@umangyadav umangyadav left a comment

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

Can you add e2e test ?

@umangyadav umangyadav added the claude-review Request a Claude PR review label Oct 1, 2026

Copilot AI left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

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

Copilot review overview

🟡 Changes recommended

Transpose-derived prefix offsets can bypass broadcasting and trigger the newly added assertion or produce an invalid attention operand.

Review effort: Balanced
Findings: 1 High severity

Open (1)
What changed in this PR

Updates attention lowering to broadcast per-batch KV indices across attention groups without relying on Q’s producer shape.

Changes:

  • Derives block-argument broadcasts from gemmG.
  • Validates index tensor shapes before rewriting.
  • Adds plain-Q single- and multi-batch regression tests.
File Description
mlir/​lib/​Conversion/​TosaToRock/​TosaToRock.cpp Revises index broadcasting and validation.
mlir/​test/​Conversion/​TosaToRock/​tosa-to-rock-attention-kvcache-plain-q.mlir Tests plain-argument Q attention lowering.

💡 Add a code-review agent skill or configure MCP servers for context-aware, tailored reviews. Learn more in the docs.

Comment on lines +3359 to +3360
assert(cast<ShapedType>(val.getType()).getShape()[0] == gemmG &&
"per-group tensor does not match the attention groups");
auto maybeNew =
addBroadcastForBlockArg(rewriter, val, firstMatMulOp.getA());
if (cast<ShapedType>(val.getType()).getShape()[0] != gemmG) {
auto maybeNew = addBroadcastForBlockArg(rewriter, val, gemmG);

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

The failure of addBroadcastForBlockArg is still silently discarded here, and the new comment above ("match() already checked the shapes are compatible") is not equivalent to what this helper requires. match()'s hasInvalidBatch validates only the shape, but addBroadcastForBlockArg additionally requires isa<BlockArgument>(blockArg) (line 2845). Those can diverge for prefixOffset: tryPrefixCausalPattern returns getValueSkipping(..., expandAndCollapse) (line 2076), which resolves only through expand/collapse, while it validates with isI32BlockArgument(..., blockArgSkip) whose skip set also includes tosa.transpose and tosa.mul (line 2546). So a derived [B, 1] value whose B divides gemmG passes hasInvalidBatch, fails the isa<BlockArgument> test here, leaves val untouched, and the rank-2 collapse then yields a [B] operand. In a debug build the new assert aborts; under NDEBUG — which is exactly the MIGraphX release configuration this PR exists to fix — the under-sized operand reaches rock.attention unchecked. rewrite() returns void and cannot bail, so the fix belongs in match(): extend hasInvalidBatch to also reject a value that is not a BlockArgument whenever its leading dimension differs from gemmG, which makes the assert genuinely unreachable. (Checklist Major: "LogicalResult returned but ignored" / "Missing assert with descriptive message on non-trivial preconditions" — here the precondition is asserted but never actually established.)

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.

Fixed.

// Helper to broadcast and reshape a block arg tensor to match output shape
// Helper to broadcast and reshape a block arg tensor to one entry per
// attention group; match() already checked the shapes are compatible.
int64_t gemmG = outputType.getShape()[0];

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

gemmG is recomputed here from the second matmul's outputType, whereas match() computed and validated against the first matmul's output (shapeC.size() == 3 ? shapeC[0] : 1, line 3226). The two agree for every shape TOSA can express today, but this re-derivation means the value hasInvalidBatch checked is not literally the value the rewrite uses — the same class of implicit-coupling bug this PR is fixing. Since AttentionMatcherValues is already the match-to-rewrite channel, consider adding an int64_t gemmG field to that struct (TosaToRock.cpp:1628), setting it alongside matched.lastKVIndex, and reading matched.gemmG here. Also note match() guards the rank with a ternary while this line indexes getShape()[0] unconditionally; routing both through one field removes that asymmetry too.

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.

Fixed.

// Copyright Advanced Micro Devices, Inc.
// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
//
// RUN: sed s/##TOKEN_ARCH##/%arch/g %s | rocmlir-opt -split-input-file --tosa-to-rock -verify-diagnostics -o -| FileCheck %s

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Both added cases exercise the accept path; the new rejection added to match() (TosaToRock.cpp:3243) has no coverage. That check is the safety net the Major finding above leans on, so it is worth pinning: add a third // ----- case whose lastValidKVIndex leading dimension does not divide gemmG (for example a tensor<3x1xi32> argument against the 4 attention groups used here) and assert with CHECK-NOT: rock.attention under a separate FileCheck prefix that the fusion is declined and the TOSA form is preserved. Checklist Major: "New optimization without a FileCheck test asserting the expected IR is produced" — the same reasoning applies to a newly added match-failure path.

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.

Tests added but the output is a diagnostic instead of preserved TOSA because the leftover softmax fails to legalize (tosa.reduce_max is marked illegal).

I instead added two // ----- cases to tosa-to-rock-attention-kvcache-plain-q.mlir using expected-error on the leftover tosa.reduce_max, which is a stronger pin (if attention fused, no diagnostic is emitted and -verify-diagnostics fails the case): a tensor<3x1xi32> index against 4 groups (declines via lastValidKVIndex does not broadcast across the 4 attention groups, confirmed by debug output) and the transposed prefix-offset case above (declines via the new non-block-arg check).

@rocmlir-pr-reviewer rocmlir-pr-reviewer Bot left a comment

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

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

Verdict: COMMENT  ·  Findings: 3 (0 Critical, 1 Major, 2 Minor)


Scope

Fixes TosaToRock's attention matcher so the per-group broadcast of lastValidKVIndex / prefixOffset is derived from the attention batch (gemmG) instead of being reconstructed from the tensor.collapse_shape feeding Q. When Q is a plain kernel argument that collapse does not exist, addBroadcastForBlockArg used to fail, and a raw [B] tensor reached rock.attention with B * heads groups — an out-of-bounds read that the op verifier would catch but the MIGraphX release pipeline does not. Two files: mlir/lib/Conversion/TosaToRock/TosaToRock.cpp and a new lit test tosa-to-rock-attention-kvcache-plain-q.mlir.

The core change is sound. addBroadcastForBlockArg now takes gemmG and emits a [B, gemmG / B] broadcast; because the groups are batch-major the subsequent collapse indexes by group exactly as the old [batch, heads(, splitKV)] shape did. Dropping the rank-3 collapse branch in prepareBlockArgTensor is safe — match() already rejects rank > 2 index tensors at TosaToRock.cpp:3197. The batch <= 0 || gemmG % batch != 0 guard correctly screens dynamic dims and division by zero.

Findings

  • TosaToRock.cpp:3348 (Major) — prepareBlockArgTensor still discards addBroadcastForBlockArg's failure, and the new match() check validates shapes but not the isa<BlockArgument> precondition the helper additionally requires, so the new assert is not guaranteed unreachable.
  • TosaToRock.cpp:3342 (Minor) — gemmG is recomputed in rewrite() from a different value than the one match() validated against.
  • tosa-to-rock-attention-kvcache-plain-q.mlir:4 (Minor) — the new match() rejection path has no negative lit coverage.

Notes

  • The new rank-1 expand path in addBroadcastForBlockArg (TosaToRock.cpp:2866-2871) is not exercised by either added case — both pass a rank-2 tensor<Bx1xi32> argument. Worth confirming one of the existing tosa-to-rock-attention-*.mlir tests covers a rank-1 tensor<Bxi32> index.
  • Back-port: mlir/lib/Conversion/TosaToRock/ is not on the checklist's enumerated "shared with rocMLIR" path list, so this is not a Major finding. That said, the attention matcher is substantively shared code and this is a real correctness fix, so please confirm whether a parallel ROCm/rocMLIR PR is needed or add a one-line note explaining why it is not.

CI status

No failing checks. Jenkins, Build and Test, MIGraphX, and Code coverage are still pending at c81b127; py-checks and detect pass.

@rocmlir-pr-reviewer rocmlir-pr-reviewer Bot removed the claude-review Request a Claude PR review label Oct 1, 2026
…ment Q

Numerically verify the plain-argument Q case against the clone harness:
the decode-shaped GQA kv-cache module MIGraphX emits, with one last-valid
index per batch for two batches. Without the previous commit the index is
passed through unbroadcast and heads 1-3 of each batch come out wrong.

Assisted-by: Claude Code
Co-Authored-By: Claude Fable 5.1 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01EtVq7zJgAEjLYv6dGRC9QB
@umangyadav

Copy link
Copy Markdown
Member

@pfultz2 mixr-attention-flash-decoding-3d-output-fusion.mlir this test is failing on gfx950.
Also check the AI review comments

@umangyadav
umangyadav requested a review from rachgupt93 October 2, 2026 12:23
@justinrosner

Copy link
Copy Markdown
Collaborator

@pfultz2 mixr-attention-flash-decoding-3d-output-fusion.mlir this test is failing on gfx950. Also check the AI review comments

That failure is unrelated. This is the test that was fixed by #558. Updating the branch should remove this failure

@pfultz2
pfultz2 requested a review from umangyadav October 2, 2026 22:58
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants