#include "mlir/Dialect/Arith/IR/Arith.h"
#include "mlir/Dialect/ArmNeon/ArmNeonDialect.h"
#include "mlir/Dialect/ArmNeon/Transforms.h"
#include "mlir/Dialect/Func/IR/FuncOps.h"
#include "mlir/Dialect/Utils/IndexingUtils.h"
#include "mlir/Dialect/Vector/IR/VectorOps.h"
#include "mlir/IR/AffineMap.h"
#include "mlir/IR/PatternMatch.h"
#define DEBUG_TYPE "lower-contract-to-arm-neon"
using namespace mlir;
using namespace mlir::arm_neon;
namespace {
template <typename Op>
std::optional<Value> getExtOperand(Value v) {
static_assert(llvm::is_one_of<Op, arith::ExtSIOp, arith::ExtUIOp>::value,
"Must be instantiated with either sign- or zero- extension op");
auto extOp = v.getDefiningOp<Op>();
if (!extOp) {
if constexpr (std::is_same<Op, arith::ExtSIOp>::value) {
auto eltTy = cast<VectorType>(v.getType()).getElementType();
if (!eltTy.isSignlessInteger() || eltTy.getIntOrFloatBitWidth() > 8)
return {};
return v;
}
return {};
}
auto inOp = extOp.getIn();
auto inTy = dyn_cast<VectorType>(inOp.getType());
if (!inTy)
return {};
auto inEltTy = inTy.getElementType();
if (!inEltTy.isSignlessInteger() || inEltTy.getIntOrFloatBitWidth() > 8)
return {};
auto outTy = dyn_cast<VectorType>(extOp.getType());
if (!(outTy && outTy.getElementType().isSignlessInteger(32)))
return {};
return inOp;
}
Value extendSmallIntVector(Location loc, VectorType srcTy, Value val,
bool signExt, PatternRewriter &rewriter) {
Type targetTy = srcTy.clone(rewriter.getI8Type());
return signExt ? rewriter.createOrFold<arith::ExtSIOp>(loc, targetTy, val)
: rewriter.createOrFold<arith::ExtUIOp>(loc, targetTy, val);
}
class VectorContractRewriter {
protected:
enum class MMLA {
Nop,
SignedInt,
UnsignedInt,
MixedInt,
Bfloat
};
MMLA mmlaOp = MMLA::Nop;
bool swapOperands = false;
Value lhs;
Value rhs;
Value acc;
int64_t dimM = 0;
int64_t dimN = 0;
int64_t dimK = 0;
SmallVector<int64_t> iterationBounds;
SmallVector<int64_t> subTileShape;
Value createMMLA(PatternRewriter &rewriter, Location loc, Value acc,
Value lhs, Value rhs) {
if (swapOperands)
std::swap(lhs, rhs);
switch (mmlaOp) {
case MMLA::SignedInt:
return rewriter.createOrFold<arm_neon::SmmlaOp>(loc, acc.getType(), acc,
lhs, rhs);
case MMLA::UnsignedInt:
return rewriter.createOrFold<arm_neon::UmmlaOp>(loc, acc.getType(), acc,
lhs, rhs);
case MMLA::MixedInt:
return rewriter.createOrFold<arm_neon::UsmmlaOp>(loc, acc.getType(), acc,
lhs, rhs);
case MMLA::Bfloat:
return arm_neon::BfmmlaOp::create(rewriter, loc, acc.getType(), acc, lhs,
rhs);
case MMLA::Nop:
llvm_unreachable("Uninitialized operation type");
}
}
LogicalResult matchAndInit(vector::ContractionOp op,
PatternRewriter &rewriter) {
SmallVector<vector::IteratorType> itTypes = op.getIteratorTypesArray();
if ((itTypes.size() != 3 || itTypes[0] != vector::IteratorType::parallel ||
itTypes[1] != vector::IteratorType::parallel ||
itTypes[2] != vector::IteratorType::reduction) &&
(itTypes.size() != 2 || itTypes[0] != vector::IteratorType::parallel ||
itTypes[1] != vector::IteratorType::reduction))
return rewriter.notifyMatchFailure(
op, "iterator types do not correspond to matrix multiplication");
VectorType lhsType = op.getLhsType();
VectorType rhsType = op.getRhsType();
if (!lhsType.hasRank() || !rhsType.hasRank() || lhsType.getRank() > 2 ||
rhsType.getRank() != 2)
return rewriter.notifyMatchFailure(op, "Invalid operand rank");
if (lhsType.isScalable() || rhsType.isScalable())
return rewriter.notifyMatchFailure(op,
"Not applicable to scalable vectors");
dimM = lhsType.getDimSize(0);
dimN = rhsType.getDimSize(0);
dimK = rhsType.getDimSize(1);
int64_t lhsDimK;
if (lhsType.getRank() == 1) {
dimM = 1;
lhsDimK = lhsType.getDimSize(0);
} else {
lhsDimK = lhsType.getDimSize(1);
}
if (lhsDimK != dimK)
return rewriter.notifyMatchFailure(op, "Dimensions mismatch");
return success();
}
public:
void lower(vector::ContractionOp op, PatternRewriter &rewriter) {
auto inputElementType = cast<ShapedType>(lhs.getType()).getElementType();
auto accElementType = cast<ShapedType>(acc.getType()).getElementType();
auto inputExpandedType =
VectorType::get({2, subTileShape.back()}, inputElementType);
auto outputExpandedType = VectorType::get({2, 2}, accElementType);
auto collapsedInputType =
VectorType::get(inputExpandedType.getNumElements(), inputElementType);
auto collapsedOutputType =
VectorType::get(outputExpandedType.getNumElements(), accElementType);
auto indexingMaps = op.getIndexingMapsArray();
AffineMap &lhsPermutationMap = indexingMaps[0];
AffineMap &rhsPermutationMap = indexingMaps[1];
AffineMap &accPermutationMap = indexingMaps[2];
Location loc = op.getLoc();
Value result =
arith::ConstantOp::create(rewriter, loc, op.getResultType(),
rewriter.getZeroAttr(op.getResultType()));
SmallVector<int64_t, 3> loopOrder = {0, 1};
if (iterationBounds.size() == 3)
loopOrder.push_back(2);
Value kAcc;
for (SmallVector<int64_t> offsets :
StaticTileOffsetRange(iterationBounds, subTileShape, loopOrder)) {
auto extractOperand = [&](Value operand, AffineMap permutationMap,
ArrayRef<int64_t> operandOffsets) {
SmallVector<int64_t> operandShape = applyPermutationMap(
permutationMap, ArrayRef<int64_t>(subTileShape));
SmallVector<int64_t> operandStrides(operandOffsets.size(), 1);
return rewriter.createOrFold<vector::ExtractStridedSliceOp>(
loc, operand, operandOffsets, operandShape, operandStrides);
};
SmallVector<int64_t> lhsOffsets =
applyPermutationMap(lhsPermutationMap, ArrayRef<int64_t>(offsets));
Value tiledLhs = extractOperand(lhs, lhsPermutationMap, lhsOffsets);
SmallVector<int64_t> rhsOffsets =
applyPermutationMap(rhsPermutationMap, ArrayRef<int64_t>(offsets));
Value tiledRhs = extractOperand(rhs, rhsPermutationMap, rhsOffsets);
SmallVector<int64_t> accOffsets =
applyPermutationMap(accPermutationMap, ArrayRef<int64_t>(offsets));
Value tiledAcc = extractOperand(acc, accPermutationMap, accOffsets);
if (dimM == 1) {
auto expandRowVector = [&](Value tiledOperand,
VectorType expandedTypeType) {
auto emptyOperand =
arith::ConstantOp::create(rewriter, loc, expandedTypeType,
rewriter.getZeroAttr(expandedTypeType));
SmallVector<int64_t> offsets(
cast<ShapedType>(emptyOperand.getType()).getRank(), 0);
SmallVector<int64_t> strides(
cast<ShapedType>(tiledOperand.getType()).getRank(), 1);
return rewriter.createOrFold<vector::InsertStridedSliceOp>(
loc, tiledOperand, emptyOperand, offsets, strides);
};
tiledLhs = expandRowVector(tiledLhs, inputExpandedType);
tiledAcc = expandRowVector(tiledAcc, outputExpandedType);
}
if (swapOperands)
tiledAcc = vector::TransposeOp::create(rewriter, loc, tiledAcc,
ArrayRef<int64_t>({1, 0}));
auto collapsedLhs = rewriter.createOrFold<vector::ShapeCastOp>(
tiledLhs.getLoc(), collapsedInputType, tiledLhs);
auto collapsedRhs = rewriter.createOrFold<vector::ShapeCastOp>(
tiledRhs.getLoc(), collapsedInputType, tiledRhs);
bool initialKAcc = offsets.back() == 0;
Value collapsedRes;
if (!initialKAcc) {
collapsedRes = kAcc;
} else {
collapsedRes = rewriter.createOrFold<vector::ShapeCastOp>(
tiledAcc.getLoc(), collapsedOutputType, tiledAcc);
}
kAcc =
createMMLA(rewriter, loc, collapsedRes, collapsedLhs, collapsedRhs);
Value tiledRes = rewriter.createOrFold<vector::ShapeCastOp>(
kAcc.getLoc(), tiledAcc.getType(), kAcc);
if (swapOperands)
tiledRes = vector::TransposeOp::create(rewriter, loc, tiledRes,
ArrayRef<int64_t>({1, 0}));
if (dimM == 1)
tiledRes = rewriter.createOrFold<vector::ExtractOp>(loc, tiledRes, 0);
SmallVector<int64_t> strides(
cast<ShapedType>(tiledRes.getType()).getRank(), 1);
result = rewriter.createOrFold<vector::InsertStridedSliceOp>(
loc, tiledRes, result, accOffsets, strides);
}
rewriter.replaceOp(op, result);
}
};
class VectorContractRewriterI8MM : public VectorContractRewriter {
public:
LogicalResult matchAndInit(vector::ContractionOp op,
PatternRewriter &rewriter) {
if (failed(VectorContractRewriter::matchAndInit(op, rewriter)))
return failure();
if ((dimM != 1 && dimM % 2 != 0) || dimN % 2 != 0 || dimK % 8 != 0)
return rewriter.notifyMatchFailure(op, "Unsupported operand shapes");
mmlaOp = MMLA::SignedInt;
auto maybeLhs = getExtOperand<arith::ExtSIOp>(op.getLhs());
if (!maybeLhs) {
mmlaOp = MMLA::UnsignedInt;
maybeLhs = getExtOperand<arith::ExtUIOp>(op.getLhs());
}
if (!maybeLhs)
return rewriter.notifyMatchFailure(
op, "LHS is not a sign- or zero- extended iN, N <= 8");
auto maybeRhs = getExtOperand<arith::ExtSIOp>(op.getRhs());
if (maybeRhs) {
if (mmlaOp == MMLA::UnsignedInt)
mmlaOp = MMLA::MixedInt;
} else {
if (mmlaOp == MMLA::SignedInt) {
mmlaOp = MMLA::MixedInt;
swapOperands = true;
}
maybeRhs = getExtOperand<arith::ExtUIOp>(op.getRhs());
}
if (!maybeRhs)
return rewriter.notifyMatchFailure(
op, "RHS is not a sign- or zero- extended iN, N <= 8");
lhs = *maybeLhs;
rhs = *maybeRhs;
acc = op.getAcc();
Location loc = op.getLoc();
auto lhsExtInType = cast<VectorType>(lhs.getType());
if (lhsExtInType.getElementTypeBitWidth() < 8)
lhs = extendSmallIntVector(loc, lhsExtInType, lhs,
(mmlaOp == MMLA::SignedInt ||
(mmlaOp == MMLA::MixedInt && !swapOperands)),
rewriter);
auto rhsExtInType = cast<VectorType>(rhs.getType());
if (rhsExtInType.getElementTypeBitWidth() < 8)
rhs = extendSmallIntVector(loc, rhsExtInType, rhs,
(mmlaOp == MMLA::SignedInt ||
(mmlaOp == MMLA::MixedInt && swapOperands)),
rewriter);
iterationBounds = *op.getShapeForUnroll();
if (iterationBounds.size() == 3)
subTileShape = SmallVector<int64_t>({dimM == 1 ? 1 : 2, 2, 8});
else
subTileShape = SmallVector<int64_t>({2, 8});
return success();
}
};
class VectorContractRewriterBFMMLA : public VectorContractRewriter {
public:
LogicalResult matchAndInit(vector::ContractionOp op,
PatternRewriter &rewriter) {
if (failed(VectorContractRewriter::matchAndInit(op, rewriter)))
return failure();
if ((dimM != 1 && dimM % 2 != 0) || dimN % 2 != 0 || dimK % 4 != 0)
return rewriter.notifyMatchFailure(op, "Unsupported operand shapes");
auto outTy = dyn_cast<VectorType>(op.getResultType());
if (!outTy || outTy.getElementType() != rewriter.getF32Type())
return rewriter.notifyMatchFailure(op,
"output type is not a vector of f32");
if (op.getLhsType().getElementType() != rewriter.getBF16Type())
return rewriter.notifyMatchFailure(op,
"input type is not a vector of bf16");
mmlaOp = MMLA::Bfloat;
swapOperands = false;
lhs = op.getLhs();
rhs = op.getRhs();
acc = op.getAcc();
iterationBounds = *op.getShapeForUnroll();
if (iterationBounds.size() == 3)
subTileShape = SmallVector<int64_t>({dimM == 1 ? 1 : 2, 2, 4});
else
subTileShape = SmallVector<int64_t>({2, 4});
return success();
}
};
class LowerContractionToNeonI8MMPattern
: public OpRewritePattern<vector::ContractionOp> {
public:
using OpRewritePattern::OpRewritePattern;
LogicalResult matchAndRewrite(vector::ContractionOp op,
PatternRewriter &rewriter) const override {
VectorContractRewriterI8MM vcr;
if (failed(vcr.matchAndInit(op, rewriter)))
return failure();
vcr.lower(op, rewriter);
return success();
}
};
class LowerContractionToNeonBFMMLAPattern
: public OpRewritePattern<vector::ContractionOp> {
public:
using OpRewritePattern::OpRewritePattern;
LogicalResult matchAndRewrite(vector::ContractionOp op,
PatternRewriter &rewriter) const override {
VectorContractRewriterBFMMLA vcr;
if (failed(vcr.matchAndInit(op, rewriter)))
return failure();
vcr.lower(op, rewriter);
return success();
}
};
}
void mlir::arm_neon::populateLowerContractionToNeonI8MMPatterns(
RewritePatternSet &patterns) {
MLIRContext *context = patterns.getContext();
patterns.add<LowerContractionToNeonI8MMPattern>(context, 2);
}
void mlir::arm_neon::populateLowerContractionToNeonBFMMLAPatterns(
RewritePatternSet &patterns) {
MLIRContext *context = patterns.getContext();
patterns.add<LowerContractionToNeonBFMMLAPattern>(context, 2);
}