Conversation
…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
There was a problem hiding this comment.
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
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.
| 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); |
There was a problem hiding this comment.
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.)
| // 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]; |
There was a problem hiding this comment.
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.
| // 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 |
There was a problem hiding this comment.
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.
There was a problem hiding this comment.
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).
There was a problem hiding this comment.
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) —prepareBlockArgTensorstill discardsaddBroadcastForBlockArg's failure, and the newmatch()check validates shapes but not theisa<BlockArgument>precondition the helper additionally requires, so the newassertis not guaranteed unreachable.TosaToRock.cpp:3342(Minor) —gemmGis recomputed inrewrite()from a different value than the onematch()validated against.tosa-to-rock-attention-kvcache-plain-q.mlir:4(Minor) — the newmatch()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-2tensor<Bx1xi32>argument. Worth confirming one of the existingtosa-to-rock-attention-*.mlirtests covers a rank-1tensor<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 parallelROCm/rocMLIRPR 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.
…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
|
@pfultz2 mixr-attention-flash-decoding-3d-output-fusion.mlir this test is failing on gfx950. |

What
TosaToRock's attention matcher rebuilt the per-group broadcast oflastValidKVIndex(andprefixOffset) from thetensor.collapse_shapefeeding Q into the first matmul. When Q is a plain kernel argument, the argument'sexpand_shapeand the matmul'scollapse_shapefold into a singleexpand_shape,addBroadcastForBlockArgfails, andprepareBlockArgTensorsilently passes the raw[B]tensor torock.attentionwithgemmG = B * headsgroups. The kernel loads the index atg_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:
addBroadcastForBlockArgnow takesgemmGand 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 dividegemmG(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_mlirkv-cache attention tests, where Q is a plain module parameter:kv_cache_attention_seq_len_mask,kv_cache_attention_seq_len_mask_batched, andkv_cache_attention_sinks_lsefail 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
tosa-to-rock-attention-kvcache-plain-q.mlirwith 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.tosa-to-rock-attention-*.mlirtests pass.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 changerocmlir-driver, which runs with the verifier on, rejects the lowering withBatch 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.librockCompiler.a: all 31test_gpu_mlirtests pass, including the three kv-cache tests above, on gfx1201.Assisted-by: Claude Code
🤖 Generated with Claude Code
https://claude.ai/code/session_01EtVq7zJgAEjLYv6dGRC9QB