Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
118 changes: 59 additions & 59 deletions mlir/lib/Conversion/TosaToRock/TosaToRock.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1630,6 +1630,9 @@ struct AttentionMatcherValues {
Value lse;
Value causalMaskInput;
Value lastKVIndex;
// Number of attention groups (batch x heads, x splitKV for flash decoding)
// that lastKVIndex and prefixOffset were validated against in match().
int64_t gemmG = 1;
bool isCausal;
Value prefixOffset;
std::optional<int64_t> lookBack;
Expand Down Expand Up @@ -2833,10 +2836,12 @@ struct AttentionRewritePattern : public OpRewritePattern<tosa::MatMulOp> {
}

// Broadcast a shared or per-batch block argument shaped [1], [1, 1], [B], or
// [B, 1] across the query heads.
// [B, 1] across the gemmG attention groups (batch x heads, x splitKV for
// flash decoding), producing a [B, gemmG / B] tensor. The groups are
// batch-major, so collapsing the result indexes it by group.
FailureOr<Value> addBroadcastForBlockArg(PatternRewriter &rewriter,
Value blockArg,
Value matrixQ) const {
int64_t gemmG) const {
if (!blockArg)
return failure();

Expand All @@ -2851,63 +2856,26 @@ struct AttentionRewritePattern : public OpRewritePattern<tosa::MatMulOp> {
(blockArgShape.size() != 2 || blockArgShape[1] != 1))
return failure();

// Find the original shape of matrixQ (before reshaping) to get the batch
// and numHeads values
if (!isa<tensor::CollapseShapeOp>(matrixQ.getDefiningOp())) {
// If we didn't find a collapse op, we can't determine the original shape
return failure();
}

auto collapse = cast<tensor::CollapseShapeOp>(matrixQ.getDefiningOp());
auto reassocIndices = collapse.getReassociationIndices();

// Check if the first reassociation merges two or three dimensions
// 2D case: [batch, numHeads] for the 4D attention layout
// 3D case: [batch, numHeads, splitKV] for the 5D flash-decoding layout
if (reassocIndices.empty() ||
(reassocIndices[0].size() != 2 && reassocIndices[0].size() != 3))
return failure();

// Get the original shape before collapse
auto srcShape = collapse.getSrcType().getShape();
size_t numCollapsedDims = reassocIndices[0].size();

if (srcShape.size() < numCollapsedDims)
return failure();

int64_t batch = srcShape[0];
if (blockArgShape[0] != 1 && blockArgShape[0] != batch)
int64_t batch = blockArgShape[0];
if (batch <= 0 || gemmG % batch != 0)
return failure();

auto loc = blockArg.getLoc();
Type elemTy = blockArgType.getElementType();

// Pad the block argument with unit dimensions up to the collapsed rank so
// it broadcasts against the heads, and against splitKV when present. The
// trailing new dimensions all hang off the last existing one.
// Pad the block argument to [batch, 1] so it broadcasts against the
// groups of each batch.
Value expanded = blockArg;
size_t blockArgRank = blockArgShape.size();
if (blockArgRank < numCollapsedDims) {
SmallVector<int64_t> expandedShape(blockArgShape);
expandedShape.append(numCollapsedDims - blockArgRank, 1);

SmallVector<ReassociationIndices, 2> reassoc;
for (size_t i = 0; i + 1 < blockArgRank; ++i)
reassoc.push_back({static_cast<int64_t>(i)});
ReassociationIndices lastGroup;
for (size_t i = blockArgRank - 1; i < numCollapsedDims; ++i)
lastGroup.push_back(static_cast<int64_t>(i));
reassoc.push_back(lastGroup);

auto expandedType = RankedTensorType::get(expandedShape, elemTy);
if (blockArgShape.size() == 1) {
auto expandedType = RankedTensorType::get({batch, 1}, elemTy);
SmallVector<ReassociationIndices, 1> reassoc = {{0, 1}};
expanded = tensor::ExpandShapeOp::create(rewriter, loc, expandedType,
blockArg, reassoc);
}

// Create a tosa.const that is all ones in our desired broadcast shape of
// batch x numHeads (x splitKV)
auto broadcastTy =
RankedTensorType::get(srcShape.take_front(numCollapsedDims), elemTy);
// batch x groupsPerBatch
auto broadcastTy = RankedTensorType::get({batch, gemmG / batch}, elemTy);
auto oneElems = cast<ElementsAttr>(rewriter.getOneAttr(broadcastTy));
auto constOp = tosa::ConstOp::create(rewriter, loc, broadcastTy, oneElems);

Expand Down Expand Up @@ -3253,6 +3221,39 @@ struct AttentionRewritePattern : public OpRewritePattern<tosa::MatMulOp> {

TypedValue<TensorType> matC = maybeFirstMatMul.value().getOutput();
ArrayRef<int64_t> shapeC = matC.getType().getShape();

// The kernel reads lastValidKVIndex and prefixOffset once per attention
// group, so a shared or per-batch value must broadcast evenly across the
// groups and a per-group value must already match them. Reject anything
// else here, since the rewrite cannot fail once it starts.
int64_t gemmG = shapeC.size() == 3 ? shapeC[0] : 1;
auto hasInvalidBatch = [&](Value v, StringRef name) {
if (!v)
return false;
ArrayRef<int64_t> shape = cast<ShapedType>(v.getType()).getShape();
bool valid = false;
if (shape.size() == 1 || (shape.size() == 2 && shape[1] == 1)) {
valid = shape[0] > 0 && gemmG % shape[0] == 0;
// Reconstructing the per-group broadcast is only implemented for
// plain block arguments, but the mask matchers can hand back a
// derived value here (e.g. the prefix offset resolved up to a
// tosa.transpose of the block argument).
if (shape[0] != gemmG)
valid &= isa<BlockArgument>(v);
} else if (shape.size() == 2) {
valid = shape[0] * shape[1] == gemmG;
}
if (!valid) {
LLVM_DEBUG(llvm::dbgs() << name << " does not broadcast across the "
<< gemmG << " attention groups\n");
return true;
}
return false;
};
if (hasInvalidBatch(lastKVIndex, "lastValidKVIndex") ||
hasInvalidBatch(prefixOffset, "prefixOffset"))
return failure();

bool isDotProduct = *(std::prev(shapeC.end(), 1)) == 1;
isDotProduct &= *(std::prev(shapeC.end(), 2)) == 1;

Expand Down Expand Up @@ -3294,6 +3295,7 @@ struct AttentionRewritePattern : public OpRewritePattern<tosa::MatMulOp> {
matched.lse = lse;
matched.causalMaskInput = causalMaskInput;
matched.lastKVIndex = lastKVIndex;
matched.gemmG = gemmG;
matched.lookBack = lookBack;
matched.lastKVClipMin = lastKVClipMin;
matched.lastKVClipMax = lastKVClipMax;
Expand Down Expand Up @@ -3346,30 +3348,28 @@ struct AttentionRewritePattern : public OpRewritePattern<tosa::MatMulOp> {
bool isCausal = matched.isCausal;
TypeAttr softmaxTypeAttr = TypeAttr::get(matched.softmaxType);

// 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
// against this same group count.
int64_t gemmG = matched.gemmG;
auto prepareBlockArgTensor = [&](Value &val) {
if (!val)
return;
// Broadcast if dimension doesn't match output
if (cast<ShapedType>(val.getType()).getShape()[0] !=
outputType.getShape()[0]) {
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.

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.

Resolved -- addressed in this revision.

if (succeeded(maybeNew))
val = maybeNew.value();
}
// Reshape {batch, numHeads} -> {batch * numHeads}
// Reshape {batch, groupsPerBatch} -> {batch * groupsPerBatch}
int64_t rank = cast<ShapedType>(val.getType()).getRank();
if (rank == 2) {
SmallVector<ReassociationIndices> reassocIndices = {{0, 1}};
val = tensor::CollapseShapeOp::create(rewriter, op.getLoc(), val,
reassocIndices);
} else if (rank == 3) {
// We will only have rank == 3 when we have flash decoding.
SmallVector<ReassociationIndices> reassocIndices = {{0, 1, 2}};
val = tensor::CollapseShapeOp::create(rewriter, op.getLoc(), val,
reassocIndices);
}
assert(cast<ShapedType>(val.getType()).getShape()[0] == gemmG &&
"per-group tensor does not match the attention groups");
Comment on lines +3371 to +3372
};

prepareBlockArgTensor(lastKVIndex);
Expand Down
Loading
Loading