diff --git a/mlir/lib/Conversion/TosaToRock/TosaToRock.cpp b/mlir/lib/Conversion/TosaToRock/TosaToRock.cpp index ad08f18581..257e6a953a 100644 --- a/mlir/lib/Conversion/TosaToRock/TosaToRock.cpp +++ b/mlir/lib/Conversion/TosaToRock/TosaToRock.cpp @@ -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 lookBack; @@ -2833,10 +2836,12 @@ struct AttentionRewritePattern : public OpRewritePattern { } // 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 addBroadcastForBlockArg(PatternRewriter &rewriter, Value blockArg, - Value matrixQ) const { + int64_t gemmG) const { if (!blockArg) return failure(); @@ -2851,63 +2856,26 @@ struct AttentionRewritePattern : public OpRewritePattern { (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(matrixQ.getDefiningOp())) { - // If we didn't find a collapse op, we can't determine the original shape - return failure(); - } - - auto collapse = cast(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 expandedShape(blockArgShape); - expandedShape.append(numCollapsedDims - blockArgRank, 1); - - SmallVector reassoc; - for (size_t i = 0; i + 1 < blockArgRank; ++i) - reassoc.push_back({static_cast(i)}); - ReassociationIndices lastGroup; - for (size_t i = blockArgRank - 1; i < numCollapsedDims; ++i) - lastGroup.push_back(static_cast(i)); - reassoc.push_back(lastGroup); - - auto expandedType = RankedTensorType::get(expandedShape, elemTy); + if (blockArgShape.size() == 1) { + auto expandedType = RankedTensorType::get({batch, 1}, elemTy); + SmallVector 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(rewriter.getOneAttr(broadcastTy)); auto constOp = tosa::ConstOp::create(rewriter, loc, broadcastTy, oneElems); @@ -3253,6 +3221,39 @@ struct AttentionRewritePattern : public OpRewritePattern { TypedValue matC = maybeFirstMatMul.value().getOutput(); ArrayRef 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 shape = cast(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(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; @@ -3294,6 +3295,7 @@ struct AttentionRewritePattern : public OpRewritePattern { matched.lse = lse; matched.causalMaskInput = causalMaskInput; matched.lastKVIndex = lastKVIndex; + matched.gemmG = gemmG; matched.lookBack = lookBack; matched.lastKVClipMin = lastKVClipMin; matched.lastKVClipMax = lastKVClipMax; @@ -3346,30 +3348,28 @@ struct AttentionRewritePattern : public OpRewritePattern { 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(val.getType()).getShape()[0] != - outputType.getShape()[0]) { - auto maybeNew = - addBroadcastForBlockArg(rewriter, val, firstMatMulOp.getA()); + if (cast(val.getType()).getShape()[0] != gemmG) { + auto maybeNew = addBroadcastForBlockArg(rewriter, val, gemmG); if (succeeded(maybeNew)) val = maybeNew.value(); } - // Reshape {batch, numHeads} -> {batch * numHeads} + // Reshape {batch, groupsPerBatch} -> {batch * groupsPerBatch} int64_t rank = cast(val.getType()).getRank(); if (rank == 2) { SmallVector 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 reassocIndices = {{0, 1, 2}}; - val = tensor::CollapseShapeOp::create(rewriter, op.getLoc(), val, - reassocIndices); } + assert(cast(val.getType()).getShape()[0] == gemmG && + "per-group tensor does not match the attention groups"); }; prepareBlockArgTensor(lastKVIndex); diff --git a/mlir/test/Conversion/TosaToRock/tosa-to-rock-attention-kvcache-plain-q.mlir b/mlir/test/Conversion/TosaToRock/tosa-to-rock-attention-kvcache-plain-q.mlir new file mode 100644 index 0000000000..1cd54a976b --- /dev/null +++ b/mlir/test/Conversion/TosaToRock/tosa-to-rock-attention-kvcache-plain-q.mlir @@ -0,0 +1,170 @@ +// 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 + +// Q is a plain kernel argument, so its matmul operand is an expand_shape and +// not a collapse_shape. The per-batch lastValidKVIndex must still be broadcast +// across the 4 attention groups instead of being passed through unchanged. +// CHECK-LABEL: func @mlir_attention_plain_q +// CHECK: %[[HEAD_BROADCAST:.*]] = rock.transform %arg3 {{.*}} : tensor<1x1xi32> to tensor<1x4xi32> +// CHECK: %[[FLAT:.*]] = tensor.collapse_shape %[[HEAD_BROADCAST]] +// CHECK: rock.attention +// CHECK: lastValidKVIndex = (%[[FLAT]] : tensor<4xi32>) +func.func @mlir_attention_plain_q(%arg0: tensor<16xf16>, %arg1: tensor<128xf16>, %arg2: tensor<128xf16>, %arg3: tensor<1x1xi32>) -> (tensor<16xf16>) attributes {rock.kernel, rock.arch = "##TOKEN_ARCH##"} { + %q = tensor.expand_shape %arg0 [[0, 1, 2]] output_shape [4, 1, 4] : tensor<16xf16> into tensor<4x1x4xf16> + %k = tensor.expand_shape %arg1 [[0, 1, 2]] output_shape [4, 4, 8] : tensor<128xf16> into tensor<4x4x8xf16> + %v = tensor.expand_shape %arg2 [[0, 1, 2]] output_shape [4, 8, 4] : tensor<128xf16> into tensor<4x8x4xf16> + %a_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf16>}> : () -> tensor<1xf16> + %b_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf16>}> : () -> tensor<1xf16> + %0 = tosa.matmul %q, %k, %a_zp, %b_zp {acc_type = f32} : (tensor<4x1x4xf16>, tensor<4x4x8xf16>, tensor<1xf16>, tensor<1xf16>) -> tensor<4x1x8xf16> + %expanded = tensor.expand_shape %0 [[0, 1], [2], [3]] output_shape [1, 4, 1, 8] : tensor<4x1x8xf16> into tensor<1x4x1x8xf16> + %ninf = "tosa.const"() <{values = dense<0xFC00> : tensor<1x4x1x8xf16>}> : () -> tensor<1x4x1x8xf16> + %scale = "tosa.const"() <{values = dense<1.250000e-01> : tensor<1x4x1x8xf16>}> : () -> tensor<1x4x1x8xf16> + %range = arith.constant dense<[[[[0, 1, 2, 3, 4, 5, 6, 7]]]]> : tensor<1x1x1x8xi32> + %ones = "tosa.const"() <{values = dense<1> : tensor<1x4x1x8xi32>}> : () -> tensor<1x4x1x8xi32> + %shift = "tosa.const"() <{values = dense<0> : tensor<1xi8>}> : () -> tensor<1xi8> + %1 = tosa.mul %range, %ones, %shift : (tensor<1x1x1x8xi32>, tensor<1x4x1x8xi32>, tensor<1xi8>) -> tensor<1x4x1x8xi32> + %idx = tensor.expand_shape %arg3 [[0], [1, 2, 3]] output_shape [1, 1, 1, 1] : tensor<1x1xi32> into tensor<1x1x1x1xi32> + %2 = tosa.mul %idx, %ones, %shift : (tensor<1x1x1x1xi32>, tensor<1x4x1x8xi32>, tensor<1xi8>) -> tensor<1x4x1x8xi32> + %3 = tosa.greater %1, %2 : (tensor<1x4x1x8xi32>, tensor<1x4x1x8xi32>) -> tensor<1x4x1x8xi1> + %4 = tosa.cast %3 : (tensor<1x4x1x8xi1>) -> tensor<1x4x1x8xi32> + %5 = tosa.cast %4 : (tensor<1x4x1x8xi32>) -> tensor<1x4x1x8xi8> + %6 = tosa.mul %expanded, %scale, %shift : (tensor<1x4x1x8xf16>, tensor<1x4x1x8xf16>, tensor<1xi8>) -> tensor<1x4x1x8xf16> + %7 = tosa.cast %5 : (tensor<1x4x1x8xi8>) -> tensor<1x4x1x8xi1> + %8 = tosa.select %7, %ninf, %6 : (tensor<1x4x1x8xi1>, tensor<1x4x1x8xf16>, tensor<1x4x1x8xf16>) -> tensor<1x4x1x8xf16> + %9 = tosa.reduce_max %8 {axis = 3 : i32} : (tensor<1x4x1x8xf16>) -> tensor<1x4x1x1xf16> + %10 = tosa.sub %8, %9 : (tensor<1x4x1x8xf16>, tensor<1x4x1x1xf16>) -> tensor<1x4x1x8xf16> + %11 = tosa.exp %10 : (tensor<1x4x1x8xf16>) -> tensor<1x4x1x8xf16> + %12 = tosa.reduce_sum %11 {axis = 3 : i32} : (tensor<1x4x1x8xf16>) -> tensor<1x4x1x1xf16> + %13 = tosa.reciprocal %12 : (tensor<1x4x1x1xf16>) -> tensor<1x4x1x1xf16> + %14 = tosa.mul %11, %13, %shift : (tensor<1x4x1x8xf16>, tensor<1x4x1x1xf16>, tensor<1xi8>) -> tensor<1x4x1x8xf16> + %collapsed = tensor.collapse_shape %14 [[0, 1], [2], [3]] : tensor<1x4x1x8xf16> into tensor<4x1x8xf16> + %15 = tosa.matmul %collapsed, %v, %a_zp, %b_zp {acc_type = f32} : (tensor<4x1x8xf16>, tensor<4x8x4xf16>, tensor<1xf16>, tensor<1xf16>) -> tensor<4x1x4xf16> + %out = tensor.collapse_shape %15 [[0, 1, 2]] : tensor<4x1x4xf16> into tensor<16xf16> + return %out : tensor<16xf16> +} + +// ----- + +// A per-batch index with two batches is broadcast over the two groups of each +// batch, keeping a distinct value per batch. +// CHECK-LABEL: func @mlir_attention_plain_q_batched +// CHECK: %[[HEAD_BROADCAST:.*]] = rock.transform %arg3 {{.*}} : tensor<2x1xi32> to tensor<2x2xi32> +// CHECK: %[[FLAT:.*]] = tensor.collapse_shape %[[HEAD_BROADCAST]] +// CHECK: rock.attention +// CHECK: lastValidKVIndex = (%[[FLAT]] : tensor<4xi32>) +func.func @mlir_attention_plain_q_batched(%arg0: tensor<16xf16>, %arg1: tensor<128xf16>, %arg2: tensor<128xf16>, %arg3: tensor<2x1xi32>) -> (tensor<16xf16>) attributes {rock.kernel, rock.arch = "##TOKEN_ARCH##"} { + %q = tensor.expand_shape %arg0 [[0, 1, 2]] output_shape [4, 1, 4] : tensor<16xf16> into tensor<4x1x4xf16> + %k = tensor.expand_shape %arg1 [[0, 1, 2]] output_shape [4, 4, 8] : tensor<128xf16> into tensor<4x4x8xf16> + %v = tensor.expand_shape %arg2 [[0, 1, 2]] output_shape [4, 8, 4] : tensor<128xf16> into tensor<4x8x4xf16> + %a_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf16>}> : () -> tensor<1xf16> + %b_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf16>}> : () -> tensor<1xf16> + %0 = tosa.matmul %q, %k, %a_zp, %b_zp {acc_type = f32} : (tensor<4x1x4xf16>, tensor<4x4x8xf16>, tensor<1xf16>, tensor<1xf16>) -> tensor<4x1x8xf16> + %expanded = tensor.expand_shape %0 [[0, 1], [2], [3]] output_shape [2, 2, 1, 8] : tensor<4x1x8xf16> into tensor<2x2x1x8xf16> + %ninf = "tosa.const"() <{values = dense<0xFC00> : tensor<2x2x1x8xf16>}> : () -> tensor<2x2x1x8xf16> + %scale = "tosa.const"() <{values = dense<1.250000e-01> : tensor<2x2x1x8xf16>}> : () -> tensor<2x2x1x8xf16> + %range = arith.constant dense<[[[[0, 1, 2, 3, 4, 5, 6, 7]]]]> : tensor<1x1x1x8xi32> + %ones = "tosa.const"() <{values = dense<1> : tensor<2x2x1x8xi32>}> : () -> tensor<2x2x1x8xi32> + %shift = "tosa.const"() <{values = dense<0> : tensor<1xi8>}> : () -> tensor<1xi8> + %1 = tosa.mul %range, %ones, %shift : (tensor<1x1x1x8xi32>, tensor<2x2x1x8xi32>, tensor<1xi8>) -> tensor<2x2x1x8xi32> + %idx = tensor.expand_shape %arg3 [[0], [1, 2, 3]] output_shape [2, 1, 1, 1] : tensor<2x1xi32> into tensor<2x1x1x1xi32> + %2 = tosa.mul %idx, %ones, %shift : (tensor<2x1x1x1xi32>, tensor<2x2x1x8xi32>, tensor<1xi8>) -> tensor<2x2x1x8xi32> + %3 = tosa.greater %1, %2 : (tensor<2x2x1x8xi32>, tensor<2x2x1x8xi32>) -> tensor<2x2x1x8xi1> + %4 = tosa.cast %3 : (tensor<2x2x1x8xi1>) -> tensor<2x2x1x8xi32> + %5 = tosa.cast %4 : (tensor<2x2x1x8xi32>) -> tensor<2x2x1x8xi8> + %6 = tosa.mul %expanded, %scale, %shift : (tensor<2x2x1x8xf16>, tensor<2x2x1x8xf16>, tensor<1xi8>) -> tensor<2x2x1x8xf16> + %7 = tosa.cast %5 : (tensor<2x2x1x8xi8>) -> tensor<2x2x1x8xi1> + %8 = tosa.select %7, %ninf, %6 : (tensor<2x2x1x8xi1>, tensor<2x2x1x8xf16>, tensor<2x2x1x8xf16>) -> tensor<2x2x1x8xf16> + %9 = tosa.reduce_max %8 {axis = 3 : i32} : (tensor<2x2x1x8xf16>) -> tensor<2x2x1x1xf16> + %10 = tosa.sub %8, %9 : (tensor<2x2x1x8xf16>, tensor<2x2x1x1xf16>) -> tensor<2x2x1x8xf16> + %11 = tosa.exp %10 : (tensor<2x2x1x8xf16>) -> tensor<2x2x1x8xf16> + %12 = tosa.reduce_sum %11 {axis = 3 : i32} : (tensor<2x2x1x8xf16>) -> tensor<2x2x1x1xf16> + %13 = tosa.reciprocal %12 : (tensor<2x2x1x1xf16>) -> tensor<2x2x1x1xf16> + %14 = tosa.mul %11, %13, %shift : (tensor<2x2x1x8xf16>, tensor<2x2x1x1xf16>, tensor<1xi8>) -> tensor<2x2x1x8xf16> + %collapsed = tensor.collapse_shape %14 [[0, 1], [2], [3]] : tensor<2x2x1x8xf16> into tensor<4x1x8xf16> + %15 = tosa.matmul %collapsed, %v, %a_zp, %b_zp {acc_type = f32} : (tensor<4x1x8xf16>, tensor<4x8x4xf16>, tensor<1xf16>, tensor<1xf16>) -> tensor<4x1x4xf16> + %out = tensor.collapse_shape %15 [[0, 1, 2]] : tensor<4x1x4xf16> into tensor<16xf16> + return %out : tensor<16xf16> +} + +// ----- + +// An index whose leading dimension (3) does not divide the 4 attention groups +// cannot be broadcast per group, so match() must decline the fusion. The +// leftover softmax then fails to legalize, which is what pins the decline: if +// the attention had fused, no diagnostic would be emitted and +// -verify-diagnostics would fail this case. +func.func @mlir_attention_plain_q_indivisible(%arg0: tensor<16xf16>, %arg1: tensor<128xf16>, %arg2: tensor<384xf16>, %arg3: tensor<3x1xi32>) -> (tensor<48xf16>) attributes {rock.kernel, rock.arch = "##TOKEN_ARCH##"} { + %q = tensor.expand_shape %arg0 [[0, 1, 2]] output_shape [4, 1, 4] : tensor<16xf16> into tensor<4x1x4xf16> + %k = tensor.expand_shape %arg1 [[0, 1, 2]] output_shape [4, 4, 8] : tensor<128xf16> into tensor<4x4x8xf16> + %v = tensor.expand_shape %arg2 [[0, 1, 2]] output_shape [12, 8, 4] : tensor<384xf16> into tensor<12x8x4xf16> + %a_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf16>}> : () -> tensor<1xf16> + %b_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf16>}> : () -> tensor<1xf16> + %0 = tosa.matmul %q, %k, %a_zp, %b_zp {acc_type = f32} : (tensor<4x1x4xf16>, tensor<4x4x8xf16>, tensor<1xf16>, tensor<1xf16>) -> tensor<4x1x8xf16> + %expanded = tensor.expand_shape %0 [[0, 1], [2], [3]] output_shape [1, 4, 1, 8] : tensor<4x1x8xf16> into tensor<1x4x1x8xf16> + %ninf = "tosa.const"() <{values = dense<0xFC00> : tensor<3x4x1x8xf16>}> : () -> tensor<3x4x1x8xf16> + %scale = "tosa.const"() <{values = dense<1.250000e-01> : tensor<1x4x1x8xf16>}> : () -> tensor<1x4x1x8xf16> + %range = arith.constant dense<[[[[0, 1, 2, 3, 4, 5, 6, 7]]]]> : tensor<1x1x1x8xi32> + %ones = "tosa.const"() <{values = dense<1> : tensor<3x4x1x8xi32>}> : () -> tensor<3x4x1x8xi32> + %shift = "tosa.const"() <{values = dense<0> : tensor<1xi8>}> : () -> tensor<1xi8> + %1 = tosa.mul %range, %ones, %shift : (tensor<1x1x1x8xi32>, tensor<3x4x1x8xi32>, tensor<1xi8>) -> tensor<3x4x1x8xi32> + %idx = tensor.expand_shape %arg3 [[0], [1, 2, 3]] output_shape [3, 1, 1, 1] : tensor<3x1xi32> into tensor<3x1x1x1xi32> + %2 = tosa.mul %idx, %ones, %shift : (tensor<3x1x1x1xi32>, tensor<3x4x1x8xi32>, tensor<1xi8>) -> tensor<3x4x1x8xi32> + %3 = tosa.greater %1, %2 : (tensor<3x4x1x8xi32>, tensor<3x4x1x8xi32>) -> tensor<3x4x1x8xi1> + %4 = tosa.mul %expanded, %scale, %shift : (tensor<1x4x1x8xf16>, tensor<1x4x1x8xf16>, tensor<1xi8>) -> tensor<1x4x1x8xf16> + %5 = tosa.select %3, %ninf, %4 : (tensor<3x4x1x8xi1>, tensor<3x4x1x8xf16>, tensor<1x4x1x8xf16>) -> tensor<3x4x1x8xf16> + // expected-error@below {{failed to legalize operation 'tosa.reduce_max'}} + %6 = tosa.reduce_max %5 {axis = 3 : i32} : (tensor<3x4x1x8xf16>) -> tensor<3x4x1x1xf16> + %7 = tosa.sub %5, %6 : (tensor<3x4x1x8xf16>, tensor<3x4x1x1xf16>) -> tensor<3x4x1x8xf16> + %8 = tosa.exp %7 : (tensor<3x4x1x8xf16>) -> tensor<3x4x1x8xf16> + %9 = tosa.reduce_sum %8 {axis = 3 : i32} : (tensor<3x4x1x8xf16>) -> tensor<3x4x1x1xf16> + %10 = tosa.reciprocal %9 : (tensor<3x4x1x1xf16>) -> tensor<3x4x1x1xf16> + %11 = tosa.mul %8, %10, %shift : (tensor<3x4x1x8xf16>, tensor<3x4x1x1xf16>, tensor<1xi8>) -> tensor<3x4x1x8xf16> + %collapsed = tensor.collapse_shape %11 [[0, 1], [2], [3]] : tensor<3x4x1x8xf16> into tensor<12x1x8xf16> + %12 = tosa.matmul %collapsed, %v, %a_zp, %b_zp {acc_type = f32} : (tensor<12x1x8xf16>, tensor<12x8x4xf16>, tensor<1xf16>, tensor<1xf16>) -> tensor<12x1x4xf16> + %out = tensor.collapse_shape %12 [[0, 1, 2]] : tensor<12x1x4xf16> into tensor<48xf16> + return %out : tensor<48xf16> +} + +// ----- + +// A per-batch prefix offset that reaches the matcher through a tosa.transpose +// is validated against the underlying block argument, but the matcher returns +// the transpose result itself, which the rewrite cannot broadcast per group. +// match() must decline the fusion instead of emitting a 2-element prefixOffset +// against the 4 attention groups. +func.func @mlir_attention_prefix_transposed_offset(%arg0: tensor<64xf16>, %arg1: tensor<128xf16>, %arg2: tensor<128xf16>, %arg3: tensor<1x2xi32>) -> (tensor<64xf16>) attributes {rock.kernel, rock.arch = "##TOKEN_ARCH##"} { + %q = tensor.expand_shape %arg0 [[0, 1, 2]] output_shape [4, 4, 4] : tensor<64xf16> into tensor<4x4x4xf16> + %k = tensor.expand_shape %arg1 [[0, 1, 2]] output_shape [4, 4, 8] : tensor<128xf16> into tensor<4x4x8xf16> + %v = tensor.expand_shape %arg2 [[0, 1, 2]] output_shape [4, 8, 4] : tensor<128xf16> into tensor<4x8x4xf16> + %a_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf16>}> : () -> tensor<1xf16> + %b_zp = "tosa.const"() <{values = dense<0.0> : tensor<1xf16>}> : () -> tensor<1xf16> + %0 = tosa.matmul %q, %k, %a_zp, %b_zp {acc_type = f32} : (tensor<4x4x4xf16>, tensor<4x4x8xf16>, tensor<1xf16>, tensor<1xf16>) -> tensor<4x4x8xf16> + %expanded = tensor.expand_shape %0 [[0, 1], [2], [3]] output_shape [2, 2, 4, 8] : tensor<4x4x8xf16> into tensor<2x2x4x8xf16> + %ninf = "tosa.const"() <{values = dense<0xFC00> : tensor<2x2x4x8xf16>}> : () -> tensor<2x2x4x8xf16> + %scale = "tosa.const"() <{values = dense<1.250000e-01> : tensor<2x2x4x8xf16>}> : () -> tensor<2x2x4x8xf16> + %rows = arith.constant dense<[[[[0], [1], [2], [3]]]]> : tensor<1x1x4x1xi32> + %cols = arith.constant dense<[[[[0, 1, 2, 3, 4, 5, 6, 7]]]]> : tensor<1x1x1x8xi32> + %ones = "tosa.const"() <{values = dense<1> : tensor<2x2x4x8xi32>}> : () -> tensor<2x2x4x8xi32> + %shift = "tosa.const"() <{values = dense<0> : tensor<1xi8>}> : () -> tensor<1xi8> + %t = tosa.transpose %arg3 {perms = array} : (tensor<1x2xi32>) -> tensor<2x1xi32> + %off = tensor.expand_shape %t [[0], [1, 2, 3]] output_shape [2, 1, 1, 1] : tensor<2x1xi32> into tensor<2x1x1x1xi32> + %sum = tosa.add %rows, %off : (tensor<1x1x4x1xi32>, tensor<2x1x1x1xi32>) -> tensor<2x1x4x1xi32> + %1 = tosa.mul %sum, %ones, %shift : (tensor<2x1x4x1xi32>, tensor<2x2x4x8xi32>, tensor<1xi8>) -> tensor<2x2x4x8xi32> + %2 = tosa.mul %cols, %ones, %shift : (tensor<1x1x1x8xi32>, tensor<2x2x4x8xi32>, tensor<1xi8>) -> tensor<2x2x4x8xi32> + %3 = tosa.greater %2, %1 : (tensor<2x2x4x8xi32>, tensor<2x2x4x8xi32>) -> tensor<2x2x4x8xi1> + %4 = tosa.mul %expanded, %scale, %shift : (tensor<2x2x4x8xf16>, tensor<2x2x4x8xf16>, tensor<1xi8>) -> tensor<2x2x4x8xf16> + %5 = tosa.select %3, %ninf, %4 : (tensor<2x2x4x8xi1>, tensor<2x2x4x8xf16>, tensor<2x2x4x8xf16>) -> tensor<2x2x4x8xf16> + // expected-error@below {{failed to legalize operation 'tosa.reduce_max'}} + %6 = tosa.reduce_max %5 {axis = 3 : i32} : (tensor<2x2x4x8xf16>) -> tensor<2x2x4x1xf16> + %7 = tosa.sub %5, %6 : (tensor<2x2x4x8xf16>, tensor<2x2x4x1xf16>) -> tensor<2x2x4x8xf16> + %8 = tosa.exp %7 : (tensor<2x2x4x8xf16>) -> tensor<2x2x4x8xf16> + %9 = tosa.reduce_sum %8 {axis = 3 : i32} : (tensor<2x2x4x8xf16>) -> tensor<2x2x4x1xf16> + %10 = tosa.reciprocal %9 : (tensor<2x2x4x1xf16>) -> tensor<2x2x4x1xf16> + %11 = tosa.mul %8, %10, %shift : (tensor<2x2x4x8xf16>, tensor<2x2x4x1xf16>, tensor<1xi8>) -> tensor<2x2x4x8xf16> + %collapsed = tensor.collapse_shape %11 [[0, 1], [2], [3]] : tensor<2x2x4x8xf16> into tensor<4x4x8xf16> + %12 = tosa.matmul %collapsed, %v, %a_zp, %b_zp {acc_type = f32} : (tensor<4x4x8xf16>, tensor<4x8x4xf16>, tensor<1xf16>, tensor<1xf16>) -> tensor<4x4x4xf16> + %out = tensor.collapse_shape %12 [[0, 1, 2]] : tensor<4x4x4xf16> into tensor<64xf16> + return %out : tensor<64xf16> +} diff --git a/mlir/test/fusion/pr-e2e/attention/mixr-attention-kvcache-plain-q.mlir b/mlir/test/fusion/pr-e2e/attention/mixr-attention-kvcache-plain-q.mlir new file mode 100644 index 0000000000..c0ec22bcb4 --- /dev/null +++ b/mlir/test/fusion/pr-e2e/attention/mixr-attention-kvcache-plain-q.mlir @@ -0,0 +1,55 @@ +// Copyright Advanced Micro Devices, Inc. +// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception +// +// RUN: rocmlir-gen -fut mlir_attention --arch %arch --clone-harness %s | rocmlir-driver -kernel-pipeline=migraphx,highlevel -host-pipeline=migraphx,highlevel | rocmlir-gen -ph -rand_min_int 0 -rand_max_int 7 -rand_type_int_for_inputs=2 -rand 1 -rand_type float -fut mlir_attention --verifier clone - | rocmlir-driver -c | rocm-run | FileCheck %s +// RUN: rocmlir-gen -fut mlir_attention --arch %arch --clone-harness %s | rocmlir-driver -kernel-pipeline=migraphx,highlevel -host-pipeline=migraphx,highlevel | rocmlir-gen -ph -rand_min_int 8 -rand_max_int 8 -rand_type_int_for_inputs=2 -rand 1 -rand_type float -fut mlir_attention --verifier clone - | rocmlir-driver -c | rocm-run | FileCheck %s +// CHECK: [1 1 1] + +// Decode-shaped GQA kv-cache attention as MIGraphX emits it, with Q a plain +// kernel argument and one last-valid index per batch. Q's matmul operand is +// then an expand_shape rather than a collapse_shape, so the per-batch index +// must be broadcast across the batch * heads attention groups from the +// attention batch itself; before that fix the kernel read the index out of +// bounds for every group after the first and heads 1-3 came out wrong. + +module { + func.func @mlir_attention(%arg0: !migraphx.shaped<2x2x8x4xf16, 64x32x4x1>, %arg1: !migraphx.shaped<2x4x1x4xf16, 16x4x4x1>, %arg2: !migraphx.shaped<2x1xsi32, 1x1>, %arg3: !migraphx.shaped<2x2x8x4xf16, 64x32x4x1>) -> !migraphx.shaped<2x1x16xf16, 16x16x1> attributes {rock.kernel = "mixr"} { + %0 = migraphx.literal(dense<[0, 1, 2, 3, 4, 5, 6, 7]> : tensor<8xsi32>) : <8xsi32, 1> + %1 = migraphx.literal(dense<0xFC00> : tensor<1xf16>) : <1xf16, 1> + %2 = migraphx.literal(dense<1.250000e-01> : tensor<1xf16>) : <1xf16, 1> + %3 = migraphx.reshape %arg0 {dims = [2, 2, 1, 8, 4]} : <2x2x8x4xf16, 64x32x4x1> -> <2x2x1x8x4xf16, 64x32x32x4x1> + %4 = migraphx.transpose %3 {permutation = [0, 1, 2, 4, 3]} : <2x2x1x8x4xf16, 64x32x32x4x1> -> <2x2x1x4x8xf16, 64x32x32x1x4> + %5 = migraphx.multibroadcast %4 {out_dyn_dims = [], out_lens = [2, 2, 2, 4, 8]} : <2x2x1x4x8xf16, 64x32x32x1x4> -> <2x2x2x4x8xf16, 64x32x0x1x4> + %6 = migraphx.reshape %5 {dims = [2, 4, 4, 8]} : <2x2x2x4x8xf16, 64x32x0x1x4> -> <2x4x4x8xf16, 128x32x8x1> + %7 = migraphx.reshape %arg3 {dims = [2, 2, 1, 8, 4]} : <2x2x8x4xf16, 64x32x4x1> -> <2x2x1x8x4xf16, 64x32x32x4x1> + %8 = migraphx.multibroadcast %7 {out_dyn_dims = [], out_lens = [2, 2, 2, 8, 4]} : <2x2x1x8x4xf16, 64x32x32x4x1> -> <2x2x2x8x4xf16, 64x32x0x4x1> + %9 = migraphx.reshape %8 {dims = [2, 4, 8, 4]} : <2x2x2x8x4xf16, 64x32x0x4x1> -> <2x4x8x4xf16, 128x32x4x1> + %10 = migraphx.dot %arg1, %6 : <2x4x1x4xf16, 16x4x4x1>, <2x4x4x8xf16, 128x32x8x1> -> <2x4x1x8xf16, 32x8x8x1> + %11 = migraphx.multibroadcast %2 {out_dyn_dims = [], out_lens = [2, 4, 1, 8]} : <1xf16, 1> -> <2x4x1x8xf16, 0x0x0x0> + %12 = migraphx.mul %10, %11 : <2x4x1x8xf16, 32x8x8x1>, <2x4x1x8xf16, 0x0x0x0> -> <2x4x1x8xf16, 32x8x8x1> + %13 = migraphx.multibroadcast %0 {out_dyn_dims = [], out_lens = [2, 4, 1, 8]} : <8xsi32, 1> -> <2x4x1x8xsi32, 0x0x0x1> + %14 = migraphx.reshape %arg2 {dims = [2, 1]} : <2x1xsi32, 1x1> -> <2x1xsi32, 1x1> + %15 = migraphx.multibroadcast %14 {out_dyn_dims = [], out_lens = [2, 4]} : <2x1xsi32, 1x1> -> <2x4xsi32, 1x0> + %16 = migraphx.reshape %15 {dims = [2, 4, 1, 1]} : <2x4xsi32, 1x0> -> <2x4x1x1xsi32, 1x0x1x1> + %17 = migraphx.multibroadcast %16 {out_dyn_dims = [], out_lens = [2, 4, 1, 8]} : <2x4x1x1xsi32, 1x0x1x1> -> <2x4x1x8xsi32, 1x0x1x0> + %18 = migraphx.greater %13, %17 : <2x4x1x8xsi32, 0x0x0x1>, <2x4x1x8xsi32, 1x0x1x0> -> <2x4x1x8xsi32, 8x0x8x1> + %19 = migraphx.convert %18 {target_type = 0 : i64} : <2x4x1x8xsi32, 8x0x8x1> to <2x4x1x8xsi8, 8x0x8x1> + %20 = migraphx.multibroadcast %1 {out_dyn_dims = [], out_lens = [2, 4, 1, 8]} : <1xf16, 1> -> <2x4x1x8xf16, 0x0x0x0> + %21 = migraphx.where %19, %20, %12 : <2x4x1x8xsi8, 8x0x8x1>, <2x4x1x8xf16, 0x0x0x0>, <2x4x1x8xf16, 32x8x8x1> -> <2x4x1x8xf16, 32x8x8x1> + %22 = migraphx.reshape %21 {dims = [2, 4, 1, 8]} : <2x4x1x8xf16, 32x8x8x1> -> <2x4x1x8xf16, 32x8x8x1> + %23 = migraphx.reduce_max %22 {axes = [3]} : <2x4x1x8xf16, 32x8x8x1> -> <2x4x1x1xf16, 4x1x1x1> + %24 = migraphx.reshape %23 {dims = [2, 4, 1, 1]} : <2x4x1x1xf16, 4x1x1x1> -> <2x4x1x1xf16, 4x1x1x1> + %25 = migraphx.multibroadcast %24 {out_dyn_dims = [], out_lens = [2, 4, 1, 8]} : <2x4x1x1xf16, 4x1x1x1> -> <2x4x1x8xf16, 4x1x1x0> + %26 = migraphx.sub %21, %25 : <2x4x1x8xf16, 32x8x8x1>, <2x4x1x8xf16, 4x1x1x0> -> <2x4x1x8xf16, 32x8x8x1> + %27 = migraphx.exp %26 : <2x4x1x8xf16, 32x8x8x1> -> <2x4x1x8xf16, 32x8x8x1> + %28 = migraphx.reshape %27 {dims = [2, 4, 1, 8]} : <2x4x1x8xf16, 32x8x8x1> -> <2x4x1x8xf16, 32x8x8x1> + %29 = migraphx.reduce_sum %28 {axes = [3]} : <2x4x1x8xf16, 32x8x8x1> -> <2x4x1x1xf16, 4x1x1x1> + %30 = migraphx.reshape %29 {dims = [2, 4, 1, 1]} : <2x4x1x1xf16, 4x1x1x1> -> <2x4x1x1xf16, 4x1x1x1> + %31 = migraphx.multibroadcast %30 {out_dyn_dims = [], out_lens = [2, 4, 1, 8]} : <2x4x1x1xf16, 4x1x1x1> -> <2x4x1x8xf16, 4x1x1x0> + %32 = migraphx.div %27, %31 : <2x4x1x8xf16, 32x8x8x1>, <2x4x1x8xf16, 4x1x1x0> -> <2x4x1x8xf16, 32x8x8x1> + %33 = migraphx.dot %32, %9 : <2x4x1x8xf16, 32x8x8x1>, <2x4x8x4xf16, 128x32x4x1> -> <2x4x1x4xf16, 16x4x4x1> + %34 = migraphx.transpose %33 {permutation = [0, 2, 1, 3]} : <2x4x1x4xf16, 16x4x4x1> -> <2x1x4x4xf16, 16x4x4x1> + %35 = migraphx.reshape %34 {dims = [2, 1, 16]} : <2x1x4x4xf16, 16x4x4x1> -> <2x1x16xf16, 16x16x1> + return %35 : !migraphx.shaped<2x1x16xf16, 16x16x1> + } +}