26#define DEBUG_TYPE "vector-contract-lowering"
48 for (
const auto &it : llvm::enumerate(iteratorTypes)) {
52 results.push_back(it.value());
69 results.push_back(targetExpr);
94 return vector::ExtractOp::create(rewriter, loc, val, pos);
97 VectorType type = cast<VectorType>(val.
getType());
99 Value result = arith::ConstantOp::create(rewriter, loc, resType,
101 for (
int64_t d = 0, e = resType.getDimSize(0); d < e; d++) {
102 Value ext = vector::ExtractOp::create(rewriter, loc, val, d);
127 return vector::InsertOp::create(rewriter, loc, val,
result, pos);
130 VectorType type = cast<VectorType>(
result.getType());
131 for (
int64_t d = 0, e = type.getDimSize(0); d < e; d++) {
132 Value ext = vector::ExtractOp::create(rewriter, loc,
result, d);
133 Value ins = vector::ExtractOp::create(rewriter, loc, val, d);
135 result = vector::InsertOp::create(rewriter, loc, sto,
result, d);
141static std::optional<Value>
145 arith::FastMathFlagsAttr fmf = {}) {
146 using vector::CombiningKind;
150 if (kind == CombiningKind::MINNUMF || kind == CombiningKind::MAXNUMF ||
151 kind == CombiningKind::MINIMUMF || kind == CombiningKind::MAXIMUMF ||
152 kind == CombiningKind::MINIMUMNUMF ||
153 kind == CombiningKind::MAXIMUMNUMF)
156 mul = arith::MulIOp::create(rewriter, loc, x, y);
159 if (kind == CombiningKind::AND || kind == CombiningKind::MINUI ||
160 kind == CombiningKind::MINSI || kind == CombiningKind::MAXUI ||
161 kind == CombiningKind::MAXSI || kind == CombiningKind::OR ||
162 kind == CombiningKind::XOR)
166 if (
acc && isa<VectorType>(
acc.getType()) && kind == CombiningKind::ADD) {
167 Value fma = vector::FMAOp::create(rewriter, loc, x, y,
acc);
174 mul = arith::MulFOp::create(rewriter, loc, x, y, fmf);
178 return std::optional<Value>(
mul);
189 dimsIdx.push_back(i);
208 arith::FastMathFlagsAttr fmf = {}) {
210 return arith::AddIOp::create(rewriter, loc, x, y);
211 return arith::AddFOp::create(rewriter, loc, x, y, fmf);
218 arith::FastMathFlagsAttr fmf = {}) {
220 return arith::MulIOp::create(rewriter, loc, x, y);
221 return arith::MulFOp::create(rewriter, loc, x, y, fmf);
241class ContractionOpToOuterProductOpLowering
244 using MaskableOpRewritePattern::MaskableOpRewritePattern;
246 using FilterConstraintType =
247 std::function<LogicalResult(vector::ContractionOp op)>;
249 static LogicalResult defaultFilter(vector::ContractionOp op) {
253 ContractionOpToOuterProductOpLowering(
254 vector::VectorContractLowering vectorContractLowering,
255 MLIRContext *context, PatternBenefit benefit = 1,
256 FilterConstraintType constraint = defaultFilter)
257 : MaskableOpRewritePattern<vector::ContractionOp>(context, benefit),
258 vectorContractLowering(vectorContractLowering),
259 filter(std::move(constraint)) {}
263 PatternRewriter &rewriter)
const override;
267 vector::VectorContractLowering vectorContractLowering;
268 FilterConstraintType filter;
289class ContractionOpToDotLowering
292 using MaskableOpRewritePattern::MaskableOpRewritePattern;
294 using FilterConstraintType =
295 std::function<LogicalResult(vector::ContractionOp op)>;
297 static LogicalResult defaultFilter(vector::ContractionOp op) {
301 ContractionOpToDotLowering(
302 vector::VectorContractLowering vectorContractLowering,
303 MLIRContext *context, PatternBenefit benefit = 1,
304 const FilterConstraintType &constraint = defaultFilter)
305 : MaskableOpRewritePattern<vector::ContractionOp>(context, benefit),
306 vectorContractLowering(vectorContractLowering), filter(defaultFilter) {}
310 PatternRewriter &rewriter)
const override;
314 vector::VectorContractLowering vectorContractLowering;
315 FilterConstraintType filter;
332class ContractionOpLowering
335 using MaskableOpRewritePattern::MaskableOpRewritePattern;
336 using FilterConstraintType =
337 std::function<LogicalResult(vector::ContractionOp op)>;
339 static LogicalResult defaultFilter(vector::ContractionOp op) {
343 ContractionOpLowering(
344 vector::VectorContractLowering vectorContractLoweringOption,
345 MLIRContext *context, PatternBenefit benefit = 1,
346 FilterConstraintType constraint = defaultFilter)
347 : MaskableOpRewritePattern<vector::ContractionOp>(context, benefit),
348 vectorContractLoweringOption(vectorContractLoweringOption),
349 filter(std::move(constraint)) {}
353 PatternRewriter &rewriter)
const override;
357 vector::VectorContractLowering vectorContractLoweringOption;
358 FilterConstraintType filter;
360 FailureOr<Value> lowerParallel(PatternRewriter &rewriter,
361 vector::ContractionOp op, int64_t lhsIndex,
362 int64_t rhsIndex, Value mask)
const;
364 FailureOr<Value> lowerReduction(PatternRewriter &rewriter,
365 vector::ContractionOp op, Value mask)
const;
370struct UnrolledOuterProductGenerator
372 UnrolledOuterProductGenerator(RewriterBase &
b, vector::ContractionOp op)
373 : StructuredGenerator<vector::ContractionOp, vector::IteratorType>(
b, op),
374 kind(op.getKind()),
lhs(op.getLhs()),
rhs(op.getRhs()),
375 res(op.getAcc()), lhsType(op.getLhsType()) {
376 auto maskableOp = cast<MaskableOpInterface>(op.getOperation());
377 if (maskableOp.isMasked())
378 mask = maskableOp.getMaskingOp().getMask();
381 Value t(Value v, ArrayRef<int64_t> perm = {1, 0}) {
384 return vector::TransposeOp::create(rewriter, loc, v, perm);
387 Value
promote(Value v, Type dstElementType) {
388 Type elementType = v.
getType();
389 auto vecType = dyn_cast<VectorType>(elementType);
391 elementType = vecType.getElementType();
392 if (elementType == dstElementType)
394 Type promotedType = dstElementType;
396 promotedType = vecType.clone(promotedType);
397 if (isa<FloatType>(dstElementType))
398 return arith::ExtFOp::create(rewriter, loc, promotedType, v,
400 return arith::ExtSIOp::create(rewriter, loc, promotedType, v);
404 VectorType lhsType,
int reductionSize,
405 std::optional<Value> maybeMask = std::nullopt) {
407 if (mask && !maybeMask.has_value())
410 Type resElementType = cast<VectorType>(res.getType()).getElementType();
411 for (
int64_t k = 0; k < reductionSize; ++k) {
412 Value extractA = vector::ExtractOp::create(rewriter, loc, lhs, k);
413 Value extractB = vector::ExtractOp::create(rewriter, loc, rhs, k);
414 extractA = promote(extractA, resElementType);
415 extractB = promote(extractB, resElementType);
417 if (maybeMask.has_value() && maybeMask.value())
419 vector::ExtractOp::create(rewriter, loc, maybeMask.value(), k);
421 Operation *outerProdOp = vector::OuterProductOp::create(
422 rewriter, loc, res.getType(), extractA, extractB, res, kind);
434 if (vecType.getScalableDims()[reductionDim])
436 int64_t reductionSize = vecType.getDimSize(reductionDim);
437 assert(reductionSize > 0 &&
438 "Reduction dim must be a known static size to allow unrolling");
439 return reductionSize;
444 if (!
iters({Par(), Par(), Red()}))
451 if (
layout({{m, k}, {k, n}, {m, n}})) {
456 Value tMask = t(mask, {2, 0, 1});
457 return outerProd(tLhs, rhs, res, lhsType, *reductionSize, tMask);
461 if (
layout({{m, k}, {n, k}, {m, n}})) {
465 Value tMask = t(mask, {2, 0, 1});
466 return outerProd(tLhs, tRhs, res, lhsType, *reductionSize, tMask);
470 if (
layout({{k, m}, {k, n}, {m, n}})) {
472 Value tMask = t(mask, {2, 0, 1});
473 return outerProd(lhs, rhs, res, lhsType, *reductionSize, tMask);
477 if (
layout({{k, m}, {n, k}, {m, n}})) {
480 Value tMask = t(mask, {2, 0, 1});
481 return outerProd(lhs, tRhs, res, lhsType, *reductionSize, tMask);
486 if (
layout({{m, k}, {k, n}, {n, m}})) {
489 Value tMask = t(mask, {2, 0, 1});
490 return outerProd(rhs, tLhs, res, lhsType, *reductionSize, tMask);
494 if (
layout({{m, k}, {n, k}, {n, m}})) {
498 Value tMask = t(mask, {2, 0, 1});
499 return outerProd(tRhs, tLhs, res, lhsType, *reductionSize, tMask);
502 if (
layout({{k, m}, {k, n}, {n, m}})) {
504 Value tMask = t(mask, {2, 0, 1});
505 return outerProd(rhs, lhs, res, lhsType, *reductionSize, tMask);
508 if (
layout({{k, m}, {n, k}, {n, m}})) {
511 Value tMask = t(mask, {2, 0, 1});
512 return outerProd(tRhs, lhs, res, lhsType, *reductionSize, tMask);
524 if (!
iters({Par(), Red()}))
530 if (
layout({{m, k}, {k}, {m}})) {
533 Value tMask = t(mask);
534 return outerProd(tLhs, rhs, res, lhsType, *reductionSize, tMask);
538 if (
layout({{k, m}, {k}, {m}})) {
540 Value tMask = t(mask);
541 return outerProd(lhs, rhs, res, lhsType, *reductionSize, tMask);
545 if (
layout({{k}, {m, k}, {m}})) {
548 Value tMask = t(mask);
549 return outerProd(tRhs, lhs, res, lhsType, *reductionSize, tMask);
553 if (
layout({{k}, {k, m}, {m}})) {
555 Value tMask = t(mask);
556 return outerProd(rhs, lhs, res, lhsType, *reductionSize, tMask);
567 if (!
iters({Red(), Par()}))
573 if (
layout({{m, k}, {k}, {m}}))
575 return outerProd(t(lhs), rhs, res, lhsType, *reductionSize, mask);
577 if (
layout({{k, m}, {k}, {m}}))
579 return outerProd(lhs, rhs, res, lhsType, *reductionSize, mask);
581 if (
layout({{k}, {m, k}, {m}}))
583 return outerProd(t(rhs), lhs, res, lhsType, *reductionSize, mask);
585 if (
layout({{k}, {k, m}, {m}}))
587 return outerProd(rhs, lhs, res, lhsType, *reductionSize, mask);
592 vector::CombiningKind kind;
593 Value
lhs,
rhs, res, mask;
613ContractionOpToOuterProductOpLowering::matchAndRewriteMaskableOp(
614 vector::ContractionOp op, MaskingOpInterface maskOp,
616 if (vectorContractLowering != vector::VectorContractLowering::OuterProduct)
622 UnrolledOuterProductGenerator e(rewriter, op);
623 FailureOr<Value> matmatRes = e.matmat();
624 if (succeeded(matmatRes)) {
627 FailureOr<Value> matvecRes = e.matvec();
628 if (succeeded(matvecRes)) {
632 FailureOr<Value> tmatvecRes = e.tmatvec();
636FailureOr<Value> ContractionOpToDotLowering::matchAndRewriteMaskableOp(
637 vector::ContractionOp op, MaskingOpInterface maskOp,
646 if (vectorContractLowering != vector::VectorContractLowering::Dot)
649 auto iteratorTypes = op.getIteratorTypes().getValue();
650 static constexpr std::array<int64_t, 2> perm = {1, 0};
652 Value lhs = op.getLhs(), rhs = op.getRhs();
655 auto infer = [&](MapList m) {
671 if (maps == infer({{m, k}, {k, n}, {m, n}})) {
672 rhs = vector::TransposeOp::create(rewriter, loc, rhs, perm);
673 }
else if (maps == infer({{m, k}, {n, k}, {m, n}})) {
675 }
else if (maps == infer({{k, m}, {k, n}, {m, n}})) {
676 lhs = vector::TransposeOp::create(rewriter, loc, lhs, perm);
677 rhs = vector::TransposeOp::create(rewriter, loc, rhs, perm);
678 }
else if (maps == infer({{k, m}, {n, k}, {m, n}})) {
679 lhs = vector::TransposeOp::create(rewriter, loc, lhs, perm);
680 }
else if (maps == infer({{m, k}, {k, n}, {n, m}})) {
683 lhs = vector::TransposeOp::create(rewriter, loc, rhs, perm);
685 }
else if (maps == infer({{m, k}, {n, k}, {n, m}})) {
687 }
else if (maps == infer({{k, m}, {k, n}, {n, m}})) {
689 lhs = vector::TransposeOp::create(rewriter, loc, rhs, perm);
690 rhs = vector::TransposeOp::create(rewriter, loc, tmp, perm);
691 }
else if (maps == infer({{k, m}, {n, k}, {n, m}})) {
693 rhs = vector::TransposeOp::create(rewriter, loc, lhs, perm);
703 if (maps == infer({{m, n}, {n}, {m}})) {
705 }
else if (maps == infer({{n, m}, {n}, {m}})) {
706 lhs = vector::TransposeOp::create(rewriter, loc, lhs, perm);
707 }
else if (maps == infer({{n}, {m, n}, {m}})) {
709 }
else if (maps == infer({{n}, {n, m}, {m}})) {
711 lhs = vector::TransposeOp::create(rewriter, loc, lhs, perm);
719 VectorType dstType = cast<VectorType>(op.getResultType());
720 assert(dstType.getRank() >= 1 && dstType.getRank() <= 2 &&
721 "Expected dst type of rank 1 or 2");
723 unsigned rank = dstType.getRank();
724 unsigned dstRows = dstType.getShape()[0];
725 unsigned dstColumns = rank == 1 ? 1 : dstType.getShape()[1];
728 Value res = arith::ConstantOp::create(rewriter, loc, dstType,
730 bool isInt = isa<IntegerType>(dstType.getElementType());
731 arith::FastMathFlagsAttr fmf = op.getFastmathAttr();
733 extractedCols.reserve(dstColumns);
734 for (
unsigned r = 0; r < dstRows; ++r) {
735 Value rowLhs = vector::ExtractOp::create(rewriter, op.getLoc(), lhs, r);
736 for (
unsigned c = 0; c < dstColumns; ++c) {
743 : vector::ExtractOp::create(rewriter, op.getLoc(), rhs, c);
744 extractedCols.push_back(colRhs);
746 Value extractedColRhs = extractedCols[c];
748 createMul(op.getLoc(), rowLhs, extractedColRhs, isInt, rewriter, fmf);
749 Value sum = vector::ReductionOp::create(rewriter, op.getLoc(),
750 vector::CombiningKind::ADD,
755 res = vector::InsertOp::create(rewriter, op.getLoc(), sum, res, pos);
758 if (
auto acc = op.getAcc())
759 res =
createAdd(op.getLoc(), res,
acc, isInt, rewriter, fmf);
765struct ContractOpToElementwise
767 using MaskableOpRewritePattern::MaskableOpRewritePattern;
768 using FilterConstraintType =
769 std::function<LogicalResult(vector::ContractionOp op)>;
770 static LogicalResult defaultFilter(vector::ContractionOp op) {
773 ContractOpToElementwise(
774 vector::VectorContractLowering vectorContractLowering,
775 MLIRContext *context, PatternBenefit benefit = 1,
776 const FilterConstraintType &constraint = defaultFilter)
777 : MaskableOpRewritePattern<vector::ContractionOp>(context, benefit),
778 vectorContractLowering(vectorContractLowering), filter(defaultFilter) {}
781 matchAndRewriteMaskableOp(vector::ContractionOp contractOp,
782 MaskingOpInterface maskOp,
783 PatternRewriter &rewriter)
const override {
788 if (
failed(filter(contractOp)))
791 if (vectorContractLowering != vector::VectorContractLowering::ParallelArith)
794 ArrayRef<int64_t> lhsShape = contractOp.getLhsType().getShape();
795 ArrayRef<int64_t> rhsShape = contractOp.getRhsType().getShape();
796 AffineMap lhsMap = contractOp.getIndexingMapsArray()[0];
797 AffineMap rhsMap = contractOp.getIndexingMapsArray()[1];
798 SmallVector<int64_t> lhsReductionDims =
800 SmallVector<int64_t> rhsReductionDims =
803 for (int64_t dim : lhsReductionDims) {
804 if (lhsShape[dim] != 1)
807 for (int64_t dim : rhsReductionDims) {
808 if (rhsShape[dim] != 1)
811 AffineMap accMap = contractOp.getIndexingMapsArray()[2];
813 unsigned numLhsDimToBroadcast =
814 numParallelDims - (lhsMap.
getNumResults() - lhsReductionDims.size());
815 unsigned numRhsDimToBroadcast =
816 numParallelDims - (rhsMap.
getNumResults() - rhsReductionDims.size());
817 SmallVector<int64_t> lhsDims;
818 SmallVector<int64_t> lhsTranspose;
819 SmallVector<int64_t> rhsDims;
820 SmallVector<int64_t> rhsTranspose;
821 for (int64_t dim : lhsReductionDims)
822 lhsTranspose.push_back(numLhsDimToBroadcast + dim);
823 for (int64_t dim : rhsReductionDims)
824 rhsTranspose.push_back(numRhsDimToBroadcast + dim);
827 for (
unsigned i = 0; i < numParallelDims; i++) {
828 std::optional<unsigned> lhsDim =
831 lhsTranspose.push_back(numLhsDimToBroadcast + *lhsDim);
835 cast<VectorType>(contractOp.getResultType()).getDimSize(i));
836 lhsTranspose.push_back(lhsDims.size() - 1);
838 std::optional<unsigned> rhsDim =
841 rhsTranspose.push_back(numRhsDimToBroadcast + *rhsDim);
845 cast<VectorType>(contractOp.getResultType()).getDimSize(i));
846 rhsTranspose.push_back(rhsDims.size() - 1);
849 Value newLhs = contractOp.getLhs();
850 Value newRhs = contractOp.getRhs();
851 Location loc = contractOp.getLoc();
852 if (!lhsDims.empty()) {
853 lhsDims.append(lhsShape.begin(), lhsShape.end());
855 VectorType::get(lhsDims, contractOp.getLhsType().getElementType());
856 newLhs = vector::BroadcastOp::create(rewriter, loc, expandedType, newLhs);
858 if (!rhsDims.empty()) {
859 rhsDims.append(rhsShape.begin(), rhsShape.end());
861 VectorType::get(rhsDims, contractOp.getRhsType().getElementType());
862 newRhs = vector::BroadcastOp::create(rewriter, loc, expandedType, newRhs);
864 bool isInt = contractOp.getLhsType().getElementType().isIntOrIndex();
865 newLhs = vector::TransposeOp::create(rewriter, loc, newLhs, lhsTranspose);
866 newRhs = vector::TransposeOp::create(rewriter, loc, newRhs, rhsTranspose);
867 SmallVector<int64_t> lhsOffsets(lhsReductionDims.size(), 0);
868 SmallVector<int64_t> rhsOffsets(rhsReductionDims.size(), 0);
869 newLhs = vector::ExtractOp::create(rewriter, loc, newLhs, lhsOffsets);
870 newRhs = vector::ExtractOp::create(rewriter, loc, newRhs, rhsOffsets);
871 std::optional<Value>
result =
873 contractOp.getKind(), rewriter, isInt,
874 Value(), contractOp.getFastmathAttr());
883 vector::VectorContractLowering vectorContractLowering;
884 FilterConstraintType filter;
904FailureOr<Value> ContractionOpLowering::matchAndRewriteMaskableOp(
905 vector::ContractionOp op, MaskingOpInterface maskOp,
911 if (op.getLhsType().getElementType() !=
918 if (op.getKind() != vector::CombiningKind::ADD) {
920 op,
"contractions other than 'add' not supported");
926 ContractionOpToOuterProductOpLowering pat1(vectorContractLoweringOption, ctx);
927 FailureOr<Value> newVal1 =
928 pat1.matchAndRewriteMaskableOp(op, maskOp, rewriter);
932 ContractionOpToDotLowering pat2(vectorContractLoweringOption, ctx);
933 FailureOr<Value> newVal2 =
934 pat2.matchAndRewriteMaskableOp(op, maskOp, rewriter);
938 ContractOpToElementwise pat4(vectorContractLoweringOption, ctx);
939 FailureOr<Value> newVal4 =
940 pat4.matchAndRewriteMaskableOp(op, maskOp, rewriter);
948 mask = maskOp.getMask();
950 std::vector<std::pair<int64_t, int64_t>> batchDimMap = op.getBatchDimMap();
951 if (!batchDimMap.empty()) {
952 int64_t lhsIndex = batchDimMap[0].first;
953 int64_t rhsIndex = batchDimMap[0].second;
954 auto newOp = lowerParallel(rewriter, op, lhsIndex, rhsIndex, mask);
961 std::vector<std::pair<int64_t, int64_t>> contractingDimMap =
962 op.getContractingDimMap();
965 for (
auto &dimPair : contractingDimMap) {
966 lhsContractingDimSet.insert(dimPair.first);
967 rhsContractingDimSet.insert(dimPair.second);
971 VectorType lhsType = op.getLhsType();
972 for (
int64_t lhsIndex = 0, e = lhsType.getRank(); lhsIndex < e; ++lhsIndex) {
973 if (lhsContractingDimSet.count(lhsIndex) == 0) {
974 auto newOp = lowerParallel(rewriter, op, lhsIndex, -1, mask);
982 VectorType rhsType = op.getRhsType();
983 for (
int64_t rhsIndex = 0, e = rhsType.getRank(); rhsIndex < e; ++rhsIndex) {
984 if (rhsContractingDimSet.count(rhsIndex) == 0) {
985 auto newOp = lowerParallel(rewriter, op, -1, rhsIndex, mask);
993 if (!contractingDimMap.empty()) {
994 auto newOp = lowerReduction(rewriter, op, mask);
1006FailureOr<Value> ContractionOpLowering::lowerParallel(
PatternRewriter &rewriter,
1007 vector::ContractionOp op,
1011 VectorType lhsType = op.getLhsType();
1012 VectorType rhsType = op.getRhsType();
1013 VectorType resType = cast<VectorType>(op.getResultType());
1018 if (lhsIndex >= 0) {
1019 iterIndex = iMap[0].getDimPosition(lhsIndex);
1020 if (rhsIndex >= 0 && iterIndex != iMap[1].
getDimPosition(rhsIndex))
1022 diag <<
"expected lhsIndex=" << lhsIndex <<
" and rhsIndex=" << rhsIndex
1023 <<
" to map to the same dimension";
1025 if (lhsType.getScalableDims()[lhsIndex])
1027 diag <<
"Unrolling scalable dimension (lhsIndex=" << lhsIndex
1028 <<
") is not supported yet";
1030 dimSize = lhsType.getDimSize(lhsIndex);
1031 }
else if (rhsIndex >= 0) {
1032 iterIndex = iMap[1].getDimPosition(rhsIndex);
1033 if (rhsType.getScalableDims()[rhsIndex])
1035 diag <<
"Unrolling scalable dimension (rhsIndex=" << rhsIndex
1036 <<
") is not supported yet";
1038 dimSize = rhsType.getDimSize(rhsIndex);
1042 diag <<
"expected either lhsIndex=" << lhsIndex
1043 <<
" or rhsIndex=" << rhsIndex <<
" to be nonnegative";
1053 if (resIndex == -1 && dimSize != 1)
1055 diag <<
"expected the dimension for iterIndex=" << iterIndex
1056 <<
" to either appear in the result map, or to be a unit dimension";
1060 std::array<AffineMap, 3> lowIndexingMaps = {
1061 adjustMap(iMap[0], iterIndex, rewriter),
1062 adjustMap(iMap[1], iterIndex, rewriter),
1063 adjustMap(iMap[2], iterIndex, rewriter)};
1069 Value result = arith::ConstantOp::create(rewriter, loc, resType,
1072 for (
int64_t d = 0; d < dimSize; ++d) {
1073 auto lhs =
reshapeLoad(loc, op.getLhs(), lhsIndex, d, rewriter);
1074 auto rhs =
reshapeLoad(loc, op.getRhs(), rhsIndex, d, rewriter);
1079 lowMask =
reshapeLoad(loc, mask, iterIndex, d, rewriter);
1082 vector::ContractionOp::create(rewriter, loc, lhs, rhs,
acc, lowAffine,
1083 lowIter, op.getKind(), op.getFastmath());
1084 lowContract =
maskOperation(rewriter, lowContract, lowMask);
1092FailureOr<Value> ContractionOpLowering::lowerReduction(
1094 auto loc = op.getLoc();
1095 VectorType lhsType = op.getLhsType();
1096 VectorType rhsType = op.getRhsType();
1097 Type resType = op.getResultType();
1098 if (isa<VectorType>(resType))
1100 "did not expect a VectorType result");
1101 bool isInt = isa<IntegerType>(resType);
1105 std::optional<int64_t> lookupLhs =
getResultIndex(iMap[0], iterIndex);
1106 std::optional<int64_t> lookupRhs =
getResultIndex(iMap[1], iterIndex);
1107 if (!lookupLhs.has_value())
1109 diag <<
"expected iterIndex=" << iterIndex <<
"to map to a LHS dimension";
1111 if (!lookupRhs.has_value())
1113 diag <<
"expected iterIndex=" << iterIndex <<
"to map to a RHS dimension";
1115 int64_t lhsIndex = *lookupLhs;
1116 int64_t rhsIndex = *lookupRhs;
1117 int64_t dimSize = lhsType.getDimSize(lhsIndex);
1118 if (dimSize != rhsType.getDimSize(rhsIndex))
1120 diag <<
"expect LHS dimension " << lhsIndex
1121 <<
" to have the same size as RHS dimension " << rhsIndex;
1124 if (lhsType.getRank() == 1) {
1125 if (rhsType.getRank() != 1)
1127 op,
"When LHS has rank 1, expected also RHS to have rank 1");
1128 arith::FastMathFlagsAttr fmf = op.getFastmathAttr();
1129 Value m = createMul(loc, op.getLhs(), op.getRhs(), isInt, rewriter, fmf);
1130 auto kind = vector::CombiningKind::ADD;
1134 acc ? vector::ReductionOp::create(rewriter, loc, kind, m,
acc,
1136 :
vector::ReductionOp::create(rewriter, loc, kind, m,
1141 std::array<AffineMap, 3> lowIndexingMaps = {
1142 adjustMap(iMap[0], iterIndex, rewriter),
1143 adjustMap(iMap[1], iterIndex, rewriter),
1144 adjustMap(iMap[2], iterIndex, rewriter)};
1153 for (
int64_t d = 0; d < dimSize; ++d) {
1154 auto lhs =
reshapeLoad(loc, op.getLhs(), lhsIndex, d, rewriter);
1155 auto rhs =
reshapeLoad(loc, op.getRhs(), rhsIndex, d, rewriter);
1158 newMask =
reshapeLoad(loc, mask, iterIndex, d, rewriter);
1160 Operation *newContract = vector::ContractionOp::create(
1161 rewriter, loc, lhs, rhs,
result, lowAffine, lowIter, op.getKind(),
1181class OuterProductOpLowering :
public OpRewritePattern<vector::OuterProductOp> {
1185 LogicalResult matchAndRewrite(vector::OuterProductOp op,
1186 PatternRewriter &rewriter)
const override {
1187 VectorType resType = op.getResultVectorType();
1188 if ((resType.getShape().size() >= 2) && resType.allDimsScalable())
1191 auto loc = op.getLoc();
1193 VectorType lhsType = op.getOperandVectorTypeLHS();
1194 VectorType rhsType = dyn_cast<VectorType>(op.getOperandTypeRHS());
1195 Type eltType = resType.getElementType();
1196 bool isInt = isa<IntegerType, IndexType>(eltType);
1197 Value acc = op.getAcc();
1198 vector::CombiningKind kind = op.getKind();
1201 OpBuilder::InsertionGuard guard(rewriter);
1202 auto maskableOp = cast<vector::MaskableOpInterface>(op.getOperation());
1205 if (maskableOp.isMasked()) {
1207 rootOp = maskableOp.getMaskingOp();
1208 mask = maskableOp.getMaskingOp().getMask();
1216 vector::BroadcastOp::create(rewriter, loc, lhsType, op.getRhs());
1218 loc, op.getLhs(),
b, acc, kind, rewriter, isInt, mask);
1219 if (!mult.has_value())
1225 Value
result = arith::ConstantOp::create(rewriter, loc, resType,
1227 for (int64_t d = 0, e = resType.getDimSize(0); d < e; ++d) {
1228 Value x = vector::ExtractOp::create(rewriter, loc, op.getLhs(), d);
1229 Value a = vector::BroadcastOp::create(rewriter, loc, rhsType, x);
1232 r = vector::ExtractOp::create(rewriter, loc, acc, d);
1235 extrMask = vector::ExtractOp::create(rewriter, loc, mask, d);
1238 loc, a, op.getRhs(), r, kind, rewriter, isInt, extrMask);
1241 result = vector::InsertOp::create(rewriter, loc, *m,
result, d);
1253 VectorContractLowering vectorContractLoweringOption,
PatternBenefit benefit,
1254 bool disableOuterProductLowering) {
1255 if (!disableOuterProductLowering)
1256 patterns.
add<OuterProductOpLowering>(patterns.
getContext(), benefit);
1257 patterns.
add<ContractionOpLowering, ContractionOpToOuterProductOpLowering>(
1258 vectorContractLoweringOption, patterns.
getContext(), benefit);
1263 patterns.
add<OuterProductOpLowering>(patterns.
getContext(), benefit);
static int64_t product(ArrayRef< int64_t > vals)
static std::optional< int64_t > getResultIndex(AffineMap map, int64_t index)
static SmallVector< int64_t > getReductionIndex(AffineMap map, ArrayAttr iteratorTypes)
Return the positions of the reductions in the given map.
static std::optional< unsigned > getDimPosition(AffineMap map, unsigned dim)
Look for a given dimension in an affine map and return its position.
static Value reshapeStore(Location loc, Value val, Value result, int64_t index, int64_t pos, PatternRewriter &rewriter)
Inserts val into result at position pos along dimension index.
static SmallVector< Attribute > adjustIter(ArrayAttr iteratorTypes, int64_t index)
FailureOr< Value > tmatvec()
static Value createAdd(Location loc, Value x, Value y, bool isInt, PatternRewriter &rewriter, arith::FastMathFlagsAttr fmf={})
Creates an AddIOp if isInt is true otherwise create an arith::AddFOp using operands x and y.
FailureOr< Value > outerProd(Value lhs, Value rhs, Value res, VectorType lhsType, int reductionSize, std::optional< Value > maybeMask=std::nullopt)
FailureOr< Value > matvec()
static AffineMap adjustMap(AffineMap map, int64_t index, PatternRewriter &rewriter)
static Value reshapeLoad(Location loc, Value val, int64_t index, int64_t pos, PatternRewriter &rewriter)
Returns val with the dimension at position index dropped by indexing that dimension with pos.
FailureOr< Value > matmat()
Two outer parallel, one inner reduction (matmat flavor).
static std::optional< Value > createContractArithOp(Location loc, Value x, Value y, Value acc, vector::CombiningKind kind, PatternRewriter &rewriter, bool isInt, Value mask=Value(), arith::FastMathFlagsAttr fmf={})
Helper to create arithmetic operation associated with a kind of contraction.
std::optional< int64_t > getReductionSize(VectorType vecType, int64_t reductionDim)
Helper function for matmat, matvec, tmatvec. Returns the size of dimension reductionDim....
Base type for affine expression.
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
unsigned getDimPosition(unsigned idx) const
Extracts the position of the dimensional expression at the given result, when the caller knows it is ...
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
unsigned getNumDims() const
unsigned getNumResults() const
static SmallVector< AffineMap, 4 > inferFromExprList(ArrayRef< ArrayRef< AffineExpr > > exprsList, MLIRContext *context)
Returns a vector of AffineMaps; each with as many results as exprs.size(), as many dims as the larges...
TypedAttr getZeroAttr(Type type)
ArrayAttr getArrayAttr(ArrayRef< Attribute > value)
MLIRContext * getContext() const
ArrayAttr getAffineMapArrayAttr(ArrayRef< AffineMap > values)
This class contains all of the information necessary to report a diagnostic to the DiagnosticEngine.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
MLIRContext is the top-level object for a collection of MLIR operations.
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
Operation is the basic unit of execution within MLIR.
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
This class represents the benefit of a pattern match in a unitless scheme that ranges from 0 (very li...
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
MLIRContext * getContext() const
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
std::enable_if_t<!std::is_convertible< CallbackT, Twine >::value, LogicalResult > notifyMatchFailure(Location loc, CallbackT &&reasonCallback)
Used to notify the listener that the IR failed to be rewritten because of a match failure,...
Helper StructuredGenerator class to manipulate and rewrite ops with StructuredOpInterface.
bool iters(ArrayRef< IteratorType > its)
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
This is a builder type that keeps local references to arguments.
Builder & dropDim(unsigned pos)
Erase a dim from shape @pos.
void promote(RewriterBase &rewriter, scf::ForallOp forallOp)
Promotes the loop body of a scf::ForallOp to its containing block.
Value makeArithReduction(OpBuilder &b, Location loc, CombiningKind kind, Value v1, Value acc, arith::FastMathFlagsAttr fastmath=nullptr, Value mask=nullptr)
Returns the result value of reducing two scalar/vector values with the corresponding arith operation.
Operation * maskOperation(OpBuilder &builder, Operation *maskableOp, Value mask, Value passthru=Value())
Creates a vector.mask operation around a maskable operation.
bool isReductionIterator(Attribute attr)
Returns true if attr has "reduction" iterator type semantics.
Value selectPassthru(OpBuilder &builder, Value mask, Value newValue, Value passthru)
Creates a vector select operation that picks values from newValue or passthru for each result vector ...
bool isParallelIterator(Attribute attr)
Returns true if attr has "parallel" iterator type semantics.
void populateVectorOuterProductLoweringPatterns(RewritePatternSet &patterns, PatternBenefit benefit=1)
Populate the pattern set with the following patterns:
void populateVectorContractLoweringPatterns(RewritePatternSet &patterns, VectorContractLowering vectorContractLoweringOption, PatternBenefit benefit=1, bool disableOuterProductLowering=false)
Populate the pattern set with the following patterns:
Include the generated interface declarations.
void bindDims(MLIRContext *ctx, AffineExprTy &...exprs)
Bind a list of AffineExpr references to DimExpr at positions: [0 .
llvm::DenseSet< ValueT, ValueInfoT > DenseSet
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
AffineExpr getAffineDimExpr(unsigned position, MLIRContext *context)
These free functions allow clients of the API to not use classes in detail.
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
OpRewritePattern Base
Type alias to allow derived classes to inherit constructors with using Base::Base;.
A pattern for ops that implement MaskableOpInterface and that might be masked (i.e.
virtual FailureOr< Value > matchAndRewriteMaskableOp(SourceOp sourceOp, MaskingOpInterface maskingOp, PatternRewriter &rewriter) const =0