30#include "llvm/ADT/STLExtras.h"
31#include "llvm/ADT/SmallVectorExtras.h"
32#include "llvm/Support/FormatVariadic.h"
39#define DEBUG_TYPE "vector-to-vector"
49 return builder.
create(state);
81struct MultiReduceToContract
85 LogicalResult matchAndRewrite(vector::MultiDimReductionOp reduceOp,
86 PatternRewriter &rewriter)
const override {
87 if (reduceOp.getKind() != vector::CombiningKind::ADD)
89 Operation *mulOp = reduceOp.getSource().getDefiningOp();
90 if (!mulOp || !isa<arith::MulIOp, arith::MulFOp>(mulOp))
92 SmallVector<bool> reductionMask = reduceOp.getReductionMask();
94 SmallVector<AffineExpr> exprs;
95 SmallVector<vector::IteratorType> iteratorTypes;
96 for (
const auto &isReduceDim : llvm::enumerate(reductionMask)) {
97 if (!isReduceDim.value()) {
98 iteratorTypes.push_back(vector::IteratorType::parallel);
101 iteratorTypes.push_back(vector::IteratorType::reduction);
106 0, exprs, reduceOp.getContext());
111 iteratorTypes, [&](IteratorType t) -> mlir::Attribute {
112 return IteratorTypeAttr::get(rewriter.getContext(), t);
141struct CombineContractABTranspose final
145 LogicalResult matchAndRewrite(vector::ContractionOp contractOp,
146 PatternRewriter &rewriter)
const override {
147 SmallVector<AffineMap> maps =
148 llvm::to_vector<4>(contractOp.getIndexingMapsArray());
149 Value
lhs = contractOp.getLhs();
150 Value
rhs = contractOp.getRhs();
152 bool changed =
false;
153 for (Value *operand : {&
lhs, &
rhs}) {
154 AffineMap &map = maps[index++];
155 auto transposeOp = operand->getDefiningOp<vector::TransposeOp>();
159 transposeOp.getPermutation(), contractOp.getContext());
161 *operand = transposeOp.getVector();
167 contractOp,
lhs,
rhs, contractOp.getAcc(),
205struct CombineContractResultTranspose final
209 LogicalResult matchAndRewrite(vector::TransposeOp resTOp,
210 PatternRewriter &rewriter)
const override {
211 auto contractOp = resTOp.getVector().getDefiningOp<vector::ContractionOp>();
212 if (!contractOp || !contractOp->hasOneUse())
215 auto accTOp = contractOp.getAcc().getDefiningOp<vector::TransposeOp>();
219 MLIRContext *context = contractOp.getContext();
220 auto maps = llvm::to_vector<3>(contractOp.getIndexingMapsArray());
221 AffineMap contractMap = maps.back();
232 auto combinedResMap = resTMap.compose(contractMap);
239 maps.back() = combinedResMap;
242 resTOp, contractOp.getLhs(), contractOp.getRhs(), accTOp.getVector(),
275FailureOr<Value> combineContractAndBroadcast(vector::ContractionOp contractOp,
276 MaskingOpInterface maskingOp,
279 llvm::to_vector<4>(contractOp.getIndexingMapsArray());
280 Value lhs = contractOp.getLhs();
281 Value rhs = contractOp.getRhs();
283 bool changed =
false;
284 for (
Value *operand : {&lhs, &rhs}) {
289 auto sc = operand->getDefiningOp<vector::ShapeCastOp>();
290 auto broadcast = operand->getDefiningOp<vector::BroadcastOp>();
294 if (sc && !sc.isBroadcastLike())
296 contractOp,
"Operand defined via vector.shape_cast that has "
297 "non-broadcast semantics");
300 VectorType srcType = sc ? sc.getSourceVectorType()
301 : dyn_cast<VectorType>(
broadcast.getSourceType());
303 sc ? sc.getResultVectorType() :
broadcast.getResultVectorType();
307 if (!srcType || srcType.getRank() >= resType.getRank())
309 int64_t rankDiff = resType.getRank() - srcType.getRank();
310 bool innerDimBroadcast =
false;
312 for (
const auto &dim : llvm::enumerate(srcType.getShape())) {
313 if (dim.value() != resType.getDimSize(rankDiff + dim.index())) {
314 innerDimBroadcast =
true;
321 if (innerDimBroadcast)
326 bool nonUnitDimReductionBroadcast =
false;
327 for (
int64_t i = 0; i < rankDiff; ++i) {
328 if (resType.getDimSize(i) != 1 &&
331 nonUnitDimReductionBroadcast =
true;
335 if (nonUnitDimReductionBroadcast)
339 contractOp.getContext());
340 map = broadcastMap.
compose(map);
356 for (
unsigned i = 0, e = unusedDimsBitVector.size(); i < e; ++i) {
357 if (!unusedDimsBitVector.test(i))
358 iterators.push_back(contractOp.getIteratorTypes().getValue()[i]);
364 VectorType oldMaskType;
365 bool isAnyUnusedDimNonUnit =
false;
367 oldMaskType = cast<VectorType>(maskingOp.getMask().getType());
368 for (
unsigned i = 0, e = unusedDimsBitVector.size(); i < e; ++i) {
369 if (unusedDimsBitVector.test(i) && oldMaskType.getShape()[i] != 1) {
370 isAnyUnusedDimNonUnit =
true;
381 bool hasReductionIteratorApplyingOnBothSides =
false;
382 for (
unsigned i = 0; i < iterators.size(); ++i) {
386 hasReductionIteratorApplyingOnBothSides =
true;
390 if (!hasReductionIteratorApplyingOnBothSides)
398 Operation *newOp = vector::ContractionOp::create(
399 rewriter, contractOp.getLoc(), lhs, rhs, contractOp.getAcc(),
404 if (isAnyUnusedDimNonUnit)
406 "Cannont drop non-unit mask dim.");
407 assert(unusedDimsBitVector.size() ==
408 static_cast<size_t>(oldMaskType.getRank()) &&
409 "The mask rank is incorrect!");
413 Value mask = maskingOp.getMask();
414 if (unusedDimsBitVector.count() != 0) {
422 oldMaskType.getShape().drop_front(unusedDimsBitVector.count());
423 auto newShapeScalableDims =
424 oldMaskType.getScalableDims().drop_front(unusedDimsBitVector.count());
425 VectorType maskOpType =
426 VectorType::get(newShape, rewriter.
getI1Type(), newShapeScalableDims);
427 mask = vector::ShapeCastOp::create(rewriter, contractOp.getLoc(),
428 maskOpType, maskingOp.getMask())
437struct CombineContractBroadcastMask
439 using MaskableOpRewritePattern::MaskableOpRewritePattern;
442 matchAndRewriteMaskableOp(vector::ContractionOp contractOp,
443 MaskingOpInterface maskingOp,
444 PatternRewriter &rewriter)
const override {
445 return combineContractAndBroadcast(contractOp, maskingOp, rewriter);
462struct ReorderCastOpsOnBroadcast
464 using OpInterfaceRewritePattern<CastOpInterface>::OpInterfaceRewritePattern;
466 LogicalResult matchAndRewrite(CastOpInterface op,
467 PatternRewriter &rewriter)
const override {
468 if (op->getNumOperands() != 1)
470 if (!isa<VectorType>(op->getResult(0).getType()))
472 auto bcastOp = op->getOperand(0).getDefiningOp<vector::BroadcastOp>();
477 if (
auto vecTy = dyn_cast<VectorType>(bcastOp.getSourceType()))
478 castResTy = vecTy.clone(castResTy);
482 op, op->getResult(0).
getType(), castOp->getResult(0));
501struct ReorderElementwiseOpsOnTranspose final
504 LogicalResult matchAndRewrite(Operation *op,
505 PatternRewriter &rewriter)
const override {
511 SmallVector<ArrayRef<int64_t>> transposeMaps;
517 auto transposeOp = operand.getDefiningOp<vector::TransposeOp>();
519 transposeMaps.push_back(transposeOp.getPermutation());
520 srcType = transposeOp.getSourceVectorType();
525 if (transposeMaps.empty())
530 if (!llvm::all_equal(transposeMaps))
533 SmallVector<Value> srcValues;
538 auto order = transposeMaps.front();
539 SmallVector<int64_t> invOrder(order.size());
540 for (
int i = 0, e = order.size(); i < e; ++i)
541 invOrder[order[i]] = i;
544 auto transposeOp = operand.getDefiningOp<vector::TransposeOp>();
546 srcValues.push_back(transposeOp.getVector());
550 srcType.clone(cast<VectorType>(operand.getType()).getElementType());
551 srcValues.push_back(vector::TransposeOp::create(
552 rewriter, operand.getLoc(), vectorType, operand, invOrder));
556 auto vectorType = srcType.clone(
558 Operation *elementwiseOp =
562 transposeMaps.front());
569 return llvm::map_to_vector<4>(arrayAttr.getAsRange<IntegerAttr>(),
570 [](IntegerAttr attr) { return attr.getInt(); });
582struct BubbleDownVectorBitCastForExtract
586 LogicalResult matchAndRewrite(vector::ExtractOp extractOp,
587 PatternRewriter &rewriter)
const override {
589 if (extractOp.getSourceVectorType().getRank() != 1)
592 auto castOp = extractOp.getSource().getDefiningOp<vector::BitCastOp>();
596 VectorType castSrcType = castOp.getSourceVectorType();
597 VectorType castDstType = castOp.getResultVectorType();
598 assert(castSrcType.getRank() == castDstType.getRank());
603 if (castSrcType.getNumElements() == 1)
608 if (castSrcType.getNumElements() > castDstType.getNumElements())
611 unsigned expandRatio =
612 castDstType.getNumElements() / castSrcType.getNumElements();
615 auto mixedPos = extractOp.getMixedPosition();
616 if (!mixedPos.empty() && !isa<Attribute>(mixedPos[0]))
618 uint64_t index = cast<IntegerAttr>(cast<Attribute>(mixedPos[0])).getInt();
622 Location loc = extractOp.getLoc();
623 Value packedValue = vector::ExtractOp::create(
624 rewriter, loc, castOp.getSource(), index / expandRatio);
625 Type packedVecType = VectorType::get({1}, packedValue.
getType());
626 Value zero = arith::ConstantOp::create(rewriter, loc, packedVecType,
628 packedValue = vector::InsertOp::create(rewriter, loc, packedValue, zero,
633 VectorType packedType =
634 VectorType::get({expandRatio}, castDstType.getElementType());
636 vector::BitCastOp::create(rewriter, loc, packedType, packedValue);
640 index % expandRatio);
657struct BubbleDownBitCastForStridedSliceExtract
661 LogicalResult matchAndRewrite(vector::ExtractStridedSliceOp extractOp,
662 PatternRewriter &rewriter)
const override {
663 auto castOp = extractOp.getSource().getDefiningOp<vector::BitCastOp>();
667 VectorType castSrcType = castOp.getSourceVectorType();
668 VectorType castDstType = castOp.getResultVectorType();
669 assert(castSrcType.getRank() == castDstType.getRank());
671 int64_t castSrcLastDim = castSrcType.getShape().back();
672 int64_t castDstLastDim = castDstType.getShape().back();
674 if (castSrcLastDim > castDstLastDim)
678 if (llvm::any_of(extractOp.getStrides().getAsValueRange<IntegerAttr>(),
679 [](
const APInt &val) { return !val.isOne(); }))
682 unsigned rank = extractOp.getSourceVectorType().getRank();
683 assert(castDstLastDim % castSrcLastDim == 0);
684 int64_t expandRatio = castDstLastDim / castSrcLastDim;
690 ArrayAttr newOffsets = extractOp.getOffsets();
691 if (newOffsets.size() == rank) {
692 SmallVector<int64_t> offsets = getIntValueVector(newOffsets);
693 if (offsets.back() % expandRatio != 0)
695 offsets.back() = offsets.back() / expandRatio;
700 ArrayAttr newSizes = extractOp.getSizes();
701 if (newSizes.size() == rank) {
702 SmallVector<int64_t> sizes = getIntValueVector(newSizes);
703 if (sizes.back() % expandRatio != 0)
705 sizes.back() = sizes.back() / expandRatio;
709 SmallVector<int64_t> dims =
710 llvm::to_vector<4>(cast<VectorType>(extractOp.getType()).getShape());
711 dims.back() = dims.back() / expandRatio;
712 VectorType newExtractType =
713 VectorType::get(dims, castSrcType.getElementType());
715 auto newExtractOp = vector::ExtractStridedSliceOp::create(
716 rewriter, extractOp.getLoc(), newExtractType, castOp.getSource(),
717 newOffsets, newSizes, extractOp.getStrides());
720 extractOp, extractOp.getType(), newExtractOp);
736struct BubbleUpBitCastForInsert :
public OpRewritePattern<vector::BitCastOp> {
739 LogicalResult matchAndRewrite(vector::BitCastOp bitcastOp,
740 PatternRewriter &rewriter)
const override {
741 VectorType castSrcType = bitcastOp.getSourceVectorType();
742 VectorType castDstType = bitcastOp.getResultVectorType();
745 if (castSrcType.getRank() == 0 || castSrcType.isScalable() ||
746 castDstType.isScalable())
749 int64_t castSrcLastDim = castSrcType.getShape().back();
750 int64_t castDstLastDim = castDstType.getShape().back();
751 bool isNumElemsShrink = castSrcLastDim >= castDstLastDim;
753 if (isNumElemsShrink) {
754 assert(castSrcLastDim % castDstLastDim == 0);
755 ratio = castSrcLastDim / castDstLastDim;
757 assert(castDstLastDim % castSrcLastDim == 0);
758 ratio = castDstLastDim / castSrcLastDim;
761 auto insertOp = bitcastOp.getSource().getDefiningOp<vector::InsertOp>();
766 auto insertSrcType = dyn_cast<VectorType>(insertOp.getValueToStoreType());
771 SmallVector<int64_t> srcDims(insertSrcType.getShape());
773 isNumElemsShrink ? srcDims.back() / ratio : srcDims.back() * ratio;
774 VectorType newCastSrcType =
775 VectorType::get(srcDims, castDstType.getElementType());
777 vector::BitCastOp::create(rewriter, bitcastOp.getLoc(), newCastSrcType,
778 insertOp.getValueToStore());
780 SmallVector<int64_t> dstDims(insertOp.getDestVectorType().getShape());
782 isNumElemsShrink ? dstDims.back() / ratio : dstDims.back() * ratio;
783 VectorType newCastDstType =
784 VectorType::get(dstDims, castDstType.getElementType());
787 auto newCastDstOp = vector::BitCastOp::create(
788 rewriter, bitcastOp.getLoc(), newCastDstType, insertOp.getDest());
792 bitcastOp, newCastSrcOp, newCastDstOp, insertOp.getMixedPosition());
808struct BubbleUpBitCastForStridedSliceInsert
812 LogicalResult matchAndRewrite(vector::BitCastOp bitcastOp,
813 PatternRewriter &rewriter)
const override {
814 VectorType castSrcType = bitcastOp.getSourceVectorType();
815 VectorType castDstType = bitcastOp.getResultVectorType();
816 assert(castSrcType.getRank() == castDstType.getRank());
818 if (castSrcType.getRank() == 0)
821 int64_t castSrcLastDim = castSrcType.getShape().back();
822 int64_t castDstLastDim = castDstType.getShape().back();
824 if (castSrcLastDim < castDstLastDim)
827 assert(castSrcLastDim % castDstLastDim == 0);
828 int64_t shrinkRatio = castSrcLastDim / castDstLastDim;
831 bitcastOp.getSource().getDefiningOp<vector::InsertStridedSliceOp>();
836 if (llvm::any_of(insertOp.getStrides().getAsValueRange<IntegerAttr>(),
837 [](
const APInt &val) { return !val.isOne(); }))
840 unsigned rank = insertOp.getSourceVectorType().getRank();
843 if (rank != insertOp.getDestVectorType().getRank())
847 unsigned sourceWidth = castSrcType.getElementType().getIntOrFloatBitWidth();
848 unsigned destinationWidth =
849 castDstType.getElementType().getIntOrFloatBitWidth();
850 unsigned numElements = destinationWidth / sourceWidth;
851 if (insertOp.getSourceVectorType().getNumElements() % numElements != 0)
854 ArrayAttr newOffsets = insertOp.getOffsets();
855 assert(newOffsets.size() == rank);
856 SmallVector<int64_t> offsets = getIntValueVector(newOffsets);
857 if (offsets.back() % shrinkRatio != 0)
859 offsets.back() = offsets.back() / shrinkRatio;
862 SmallVector<int64_t> srcDims =
863 llvm::to_vector<4>(insertOp.getSourceVectorType().getShape());
864 srcDims.back() = srcDims.back() / shrinkRatio;
865 VectorType newCastSrcType =
866 VectorType::get(srcDims, castDstType.getElementType());
869 vector::BitCastOp::create(rewriter, bitcastOp.getLoc(), newCastSrcType,
870 insertOp.getValueToStore());
872 SmallVector<int64_t> dstDims =
873 llvm::to_vector<4>(insertOp.getDestVectorType().getShape());
874 dstDims.back() = dstDims.back() / shrinkRatio;
875 VectorType newCastDstType =
876 VectorType::get(dstDims, castDstType.getElementType());
878 auto newCastDstOp = vector::BitCastOp::create(
879 rewriter, bitcastOp.getLoc(), newCastDstType, insertOp.getDest());
882 bitcastOp, bitcastOp.getType(), newCastSrcOp, newCastDstOp, newOffsets,
883 insertOp.getStrides());
911 BreakDownVectorBitCast(MLIRContext *context,
912 std::function<
bool(vector::BitCastOp)> controlFn,
913 PatternBenefit benefit)
914 : OpRewritePattern(context, benefit), controlFn(std::move(controlFn)) {}
916 LogicalResult matchAndRewrite(vector::BitCastOp bitcastOp,
917 PatternRewriter &rewriter)
const override {
919 if (controlFn && !controlFn(bitcastOp))
922 VectorType castSrcType = bitcastOp.getSourceVectorType();
923 VectorType castDstType = bitcastOp.getResultVectorType();
924 assert(castSrcType.getRank() == castDstType.getRank());
929 if (castSrcType.isScalable())
931 "Scalable vectors are not supported");
934 if (castSrcType.getRank() != 1)
937 int64_t castSrcLastDim = castSrcType.getShape().back();
938 int64_t castDstLastDim = castDstType.getShape().back();
940 if (castSrcLastDim < castDstLastDim)
943 assert(castSrcLastDim % castDstLastDim == 0);
944 int64_t shrinkRatio = castSrcLastDim / castDstLastDim;
946 if (castSrcLastDim == shrinkRatio)
949 Location loc = bitcastOp.getLoc();
950 Type elemType = castDstType.getElementType();
953 Value zero = arith::ConstantOp::create(rewriter, loc, elemType,
955 Value res = BroadcastOp::create(rewriter, loc, castDstType, zero);
957 SmallVector<int64_t> sliceShape = {castDstLastDim};
958 SmallVector<int64_t> strides = {1};
959 VectorType newCastDstType =
960 VectorType::get(SmallVector<int64_t>{castDstLastDim / shrinkRatio},
961 castDstType.getElementType());
963 for (
int i = 0, e = shrinkRatio; i < e; ++i) {
964 Value extracted = ExtractStridedSliceOp::create(
965 rewriter, loc, bitcastOp.getSource(),
966 ArrayRef<int64_t>{i * castDstLastDim}, sliceShape, strides);
968 BitCastOp::create(rewriter, loc, newCastDstType, extracted);
969 res = InsertStridedSliceOp::create(
970 rewriter, loc, bitcast, res,
971 ArrayRef<int64_t>{i * castDstLastDim / shrinkRatio}, strides);
978 std::function<bool(BitCastOp)> controlFn;
981static bool haveSameShapeAndScaling(
Type t,
Type u) {
982 auto tVec = dyn_cast<VectorType>(t);
983 auto uVec = dyn_cast<VectorType>(u);
990 return tVec.getShape() == uVec.getShape() &&
991 tVec.getScalableDims() == uVec.getScalableDims();
996static Type cloneOrReplace(
Type type,
Type newElementType) {
997 if (
auto shapedType = dyn_cast<ShapedType>(type)) {
998 return shapedType.clone(newElementType);
1000 return newElementType;
1005static Value getBroadcastLikeSource(
Value value) {
1011 if (
auto broadcast = dyn_cast<vector::BroadcastOp>(op))
1030struct ReorderElementwiseOpsOnBroadcast final
1033 LogicalResult matchAndRewrite(Operation *op,
1034 PatternRewriter &rewriter)
const override {
1042 op,
"Op doesn't have ElementwiseMappableTraits");
1046 Type resultElemType = resultType.getElementType();
1052 Value broadcastSource;
1053 Value firstBroadcastSource;
1055 Operation *definingOp = operand.getDefiningOp();
1058 if (definingOp->
hasTrait<OpTrait::ConstantLike>())
1060 Value source = getBroadcastLikeSource(operand);
1063 if (!firstBroadcastSource)
1064 firstBroadcastSource = source;
1065 if (isa<VectorType>(source.
getType())) {
1066 broadcastSource = source;
1071 if (!broadcastSource)
1072 broadcastSource = firstBroadcastSource;
1073 if (!broadcastSource)
1075 Type unbroadcastResultType =
1076 cloneOrReplace(broadcastSource.
getType(), resultElemType);
1083 if (isa<vector::FMAOp>(op) && !isa<VectorType>(unbroadcastResultType)) {
1085 op,
"Op only accepts vector types, but the broadcast source is a "
1092 if (!llvm::all_of(op->
getOperands(), [broadcastSource](Value val) {
1093 if (auto source = getBroadcastLikeSource(val))
1094 return haveSameShapeAndScaling(source.getType(),
1095 broadcastSource.getType()) ||
1096 (isa<VectorType>(broadcastSource.getType()) &&
1097 !isa<VectorType>(source.getType()));
1098 SplatElementsAttr splatConst;
1099 return matchPattern(val, m_Constant(&splatConst));
1103 "not all operands are constants or broadcasts from the same type");
1107 SmallVector<Value> srcValues;
1110 SplatElementsAttr splatConst;
1114 Type newType = cloneOrReplace(unbroadcastResultType, elementType);
1115 if (
auto newTypeShaped = dyn_cast<ShapedType>(newType)) {
1120 Operation *newConstOp =
1121 operand.getDefiningOp()->getDialect()->materializeConstant(
1122 rewriter, newConst, newType, operand.getLoc());
1123 srcValues.push_back(newConstOp->
getResult(0));
1126 if (isa<VectorType>(broadcastSource.
getType()) &&
1127 !isa<VectorType>(source.
getType()))
1128 source = vector::BroadcastOp::create(
1129 rewriter, operand.getLoc(),
1132 srcValues.push_back(source);
1137 Operation *elementwiseOp =
1142 op, resultType, elementwiseOp->
getResults());
1164class ExtractOpFromElementwise final
1169 LogicalResult matchAndRewrite(vector::ExtractOp op,
1170 PatternRewriter &rewriter)
const override {
1171 Operation *eltwise = op.getSource().getDefiningOp();
1176 isa<vector::FMAOp>(eltwise))
1190 if (!op.getDynamicPosition().empty())
1192 op,
"dynamic position not yet implemented");
1194 Type dstType = op.getType();
1196 OpBuilder::InsertionGuard g(rewriter);
1200 Location loc = eltwise->
getLoc();
1201 SmallVector<OpFoldResult> pos = op.getMixedPosition();
1203 Value newArg = vector::ExtractOp::create(rewriter, loc, arg, pos);
1204 mapping.
map(arg, newArg);
1207 Operation *newEltwise = rewriter.
clone(*eltwise, mapping);
1218static bool isSupportedMemSinkElementType(
Type type) {
1219 if (isa<IndexType>(type))
1240class ExtractOpFromLoad final :
public OpRewritePattern<vector::ExtractOp> {
1245 PatternRewriter &rewriter)
const override {
1246 auto loadOp = op.getSource().getDefiningOp<vector::LoadOp>();
1251 if (!loadOp->hasOneUse())
1254 VectorType loadVecType = loadOp.getVectorType();
1255 if (loadVecType.isScalable())
1257 "scalable vectors are not supported");
1259 MemRefType memType = loadOp.getMemRefType();
1263 if (!isSupportedMemSinkElementType(memType.getElementType()))
1266 int64_t rankOffset = memType.getRank() - loadVecType.getRank();
1270 auto extractVecType = dyn_cast<VectorType>(op.getResult().getType());
1271 int64_t finalRank = 0;
1273 finalRank = extractVecType.getRank();
1275 SmallVector<Value>
indices = loadOp.getIndices();
1276 SmallVector<OpFoldResult> extractPos = op.getMixedPosition();
1281 OpBuilder::InsertionGuard g(rewriter);
1283 Location loc = loadOp.getLoc();
1284 ArithIndexingBuilder idxBuilderf(rewriter, loc);
1285 for (
auto i : llvm::seq<int64_t>(rankOffset,
indices.size() - finalRank)) {
1286 OpFoldResult pos = extractPos[i - rankOffset];
1294 Value base = loadOp.getBase();
1295 if (extractVecType) {
1318class StoreOpFromBroadcast final :
public OpRewritePattern<vector::StoreOp> {
1323 PatternRewriter &rewriter)
const override {
1324 VectorType vecType = op.getVectorType();
1325 if (vecType.isScalable())
1327 "scalable vectors are not supported");
1329 if (isa<VectorType>(op.getMemRefType().getElementType()))
1331 op,
"memrefs of vectors are not supported");
1333 if (vecType.getNumElements() != 1)
1335 op,
"only 1-element vectors are supported");
1337 Value toStore = op.getValueToStore();
1338 Value source = getBroadcastLikeSource(toStore);
1341 op,
"value to store is not from a broadcast");
1348 Value base = op.getBase();
1351 if (isa<VectorType>(source.
getType())) {
1371 bool force32BitVectorIndices,
int64_t dim,
1380 if (dim == 0 && force32BitVectorIndices) {
1383 }
else if (dim == 0) {
1386 }
else if (force32BitVectorIndices) {
1388 llvm::to_vector<4>(llvm::seq<int32_t>(0, dim)));
1391 llvm::to_vector<4>(llvm::seq<int64_t>(0, dim)));
1393 Value indices = arith::ConstantOp::create(rewriter, loc, indicesAttr);
1397 Value ov = vector::BroadcastOp::create(rewriter, loc,
indices.getType(), o);
1406 if (force32BitVectorIndices) {
1409 b = arith::MinSIOp::create(rewriter, loc,
b, maxBound);
1413 vector::BroadcastOp::create(rewriter, loc,
indices.getType(), bound);
1414 return arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::slt,
1418template <
typename ConcreteOp>
1421 explicit MaterializeTransferMask(MLIRContext *context,
bool enableIndexOpt,
1422 PatternBenefit benefit = 1)
1423 : mlir::OpRewritePattern<ConcreteOp>(context, benefit),
1424 force32BitVectorIndices(enableIndexOpt) {}
1427 PatternRewriter &rewriter)
const override {
1428 if (!xferOp.hasOutOfBoundsDim())
1431 if (xferOp.getVectorType().getRank() > 1 || xferOp.getIndices().empty())
1434 Location loc = xferOp->getLoc();
1435 VectorType vtp = xferOp.getVectorType();
1442 unsigned lastIndex = llvm::size(xferOp.getIndices()) - 1;
1443 Value off = xferOp.getIndices()[lastIndex];
1446 Value
b = arith::SubIOp::create(rewriter, loc, dim.
getType(), dim, off);
1447 Value mask = vector::CreateMaskOp::create(
1449 VectorType::get(vtp.getShape(), rewriter.
getI1Type(),
1450 vtp.getScalableDims()),
1452 if (xferOp.getMask()) {
1454 mask = arith::AndIOp::create(rewriter, loc, mask, xferOp.getMask());
1458 xferOp.getMaskMutable().assign(mask);
1466 const bool force32BitVectorIndices;
1470class VectorCreateMaskOpConversion
1473 explicit VectorCreateMaskOpConversion(MLIRContext *context,
1474 bool enableIndexOpt,
1475 PatternBenefit benefit = 1)
1476 : mlir::OpRewritePattern<vector::CreateMaskOp>(context, benefit),
1477 force32BitVectorIndices(enableIndexOpt) {}
1479 LogicalResult matchAndRewrite(vector::CreateMaskOp op,
1480 PatternRewriter &rewriter)
const override {
1481 auto dstType = op.getType();
1482 if (cast<VectorType>(dstType).isScalable())
1484 int64_t rank = dstType.getRank();
1488 op, buildVectorComparison(rewriter, op, force32BitVectorIndices,
1489 rank == 0 ? 0 : dstType.getDimSize(0),
1495 const bool force32BitVectorIndices;
1499static bool allI1ConstantValuesSetTo(arith::ConstantOp constantOp,
bool value) {
1500 auto denseAttr = dyn_cast<DenseIntElementsAttr>(constantOp.getValue());
1505 assert(denseAttr.getElementType().isInteger(1) &&
"Unexpected type");
1506 return denseAttr.isSplat() && denseAttr.getSplatValue<
bool>() == value;
1524 PatternRewriter &rewriter)
const override {
1525 auto vecType = dyn_cast<VectorType>(selectOp.getType());
1526 if (!vecType || !vecType.getElementType().isInteger(1))
1530 Value cond = selectOp.getCondition();
1531 if (isa<VectorType>(cond.
getType()))
1535 if (vecType.getRank() != 1 || vecType.isScalable())
1539 if (vecType.getShape()[0] != 1)
1542 auto trueConst = selectOp.getTrueValue().getDefiningOp<arith::ConstantOp>();
1543 if (!trueConst || !allI1ConstantValuesSetTo(trueConst,
true))
1547 selectOp.getFalseValue().getDefiningOp<arith::ConstantOp>();
1548 if (!falseConst || !allI1ConstantValuesSetTo(falseConst,
false))
1552 auto elemType = rewriter.
getIntegerType(vecType.getNumElements());
1553 auto bcastType = VectorType::get({1}, elemType);
1574static FailureOr<size_t>
1578 if (failed(srcType.getStridesAndOffset(srcStrides, srcOffset)))
1581 auto isUnitDim = [](VectorType type,
int dim) {
1582 return type.getDimSize(dim) == 1 && !type.getScalableDims()[dim];
1589 int rankDiff = srcType.getRank() - vectorType.getRank();
1590 for (
int64_t i = 0, e = vectorType.getRank(); i < e; ++i) {
1593 int dim = vectorType.getRank() - i - 1;
1594 if (srcStrides[dim + rankDiff] != 1 ||
1595 srcType.getDimSize(dim + rankDiff) != 1 || !isUnitDim(vectorType, dim))
1607 LogicalResult matchAndRewrite(vector::TransferReadOp readOp,
1610 if (readOp.getTransferRank() == 0)
1613 auto srcType = dyn_cast<MemRefType>(readOp.getBase().getType());
1617 if (!readOp.getPermutationMap().isMinorIdentity())
1620 auto targetType = readOp.getVectorType();
1621 if (targetType.getRank() <= 1)
1624 FailureOr<size_t> maybeDimsToDrop =
1626 if (failed(maybeDimsToDrop))
1629 size_t dimsToDrop = maybeDimsToDrop.value();
1630 if (dimsToDrop == 0)
1633 auto inBounds = readOp.getInBoundsValues();
1634 auto droppedInBounds =
ArrayRef<bool>(inBounds).take_back(dimsToDrop);
1635 if (llvm::is_contained(droppedInBounds,
false))
1638 auto resultTargetVecType =
1639 VectorType::get(targetType.getShape().drop_back(dimsToDrop),
1640 targetType.getElementType(),
1641 targetType.getScalableDims().drop_back(dimsToDrop));
1643 auto loc = readOp.getLoc();
1650 MemRefType resultMemrefType = memref::SubViewOp::inferRankReducedResultType(
1651 srcType.getShape().drop_back(dimsToDrop), srcType, offsets, sizes,
1654 readOp.getInBoundsAttr().getValue().drop_back(dimsToDrop));
1655 Value rankedReducedView =
1656 memref::SubViewOp::create(rewriter, loc, resultMemrefType,
1657 readOp.getBase(), offsets, sizes, strides);
1659 cast<ShapedType>(rankedReducedView.
getType()), resultTargetVecType);
1662 Value mask = readOp.getMask();
1664 auto maskType = cast<VectorType>(mask.getType());
1665 auto reducedMaskType = VectorType::get(
1666 maskType.getShape().drop_back(dimsToDrop), maskType.getElementType(),
1667 maskType.getScalableDims().drop_back(dimsToDrop));
1668 mask = rewriter.
createOrFold<vector::ShapeCastOp>(loc, reducedMaskType,
1673 rewriter, loc, resultTargetVecType, rankedReducedView,
1674 readOp.getIndices().drop_back(dimsToDrop), AffineMapAttr::get(permMap),
1675 readOp.getPadding(), mask, inBoundsAttr);
1704 LogicalResult matchAndRewrite(vector::TransferWriteOp writeOp,
1707 if (writeOp.getTransferRank() == 0)
1710 auto srcType = dyn_cast<MemRefType>(writeOp.getBase().getType());
1714 if (!writeOp.getPermutationMap().isMinorIdentity())
1717 auto targetType = writeOp.getVectorType();
1718 if (targetType.getRank() <= 1)
1721 FailureOr<size_t> maybeDimsToDrop =
1723 if (failed(maybeDimsToDrop))
1726 size_t dimsToDrop = maybeDimsToDrop.value();
1727 if (dimsToDrop == 0)
1730 auto inBounds = writeOp.getInBoundsValues();
1731 auto droppedInBounds =
ArrayRef<bool>(inBounds).take_back(dimsToDrop);
1732 if (llvm::is_contained(droppedInBounds,
false))
1735 auto resultTargetVecType =
1736 VectorType::get(targetType.getShape().drop_back(dimsToDrop),
1737 targetType.getElementType(),
1738 targetType.getScalableDims().drop_back(dimsToDrop));
1747 MemRefType resultMemrefType = memref::SubViewOp::inferRankReducedResultType(
1748 srcType.getShape().drop_back(dimsToDrop), srcType, offsets, sizes,
1751 writeOp.getInBoundsAttr().getValue().drop_back(dimsToDrop));
1753 Value rankedReducedView =
1754 memref::SubViewOp::create(rewriter, loc, resultMemrefType,
1755 writeOp.getBase(), offsets, sizes, strides);
1757 cast<ShapedType>(rankedReducedView.
getType()), resultTargetVecType);
1759 auto shapeCast = rewriter.
createOrFold<vector::ShapeCastOp>(
1760 loc, resultTargetVecType, writeOp.getVector());
1763 Value mask = writeOp.getMask();
1765 auto maskType = cast<VectorType>(mask.getType());
1766 auto reducedMaskType = VectorType::get(
1767 maskType.getShape().drop_back(dimsToDrop), maskType.getElementType(),
1768 maskType.getScalableDims().drop_back(dimsToDrop));
1769 mask = rewriter.
createOrFold<vector::ShapeCastOp>(loc, reducedMaskType,
1774 writeOp, shapeCast, rankedReducedView,
1775 writeOp.getIndices().drop_back(dimsToDrop), AffineMapAttr::get(permMap),
1776 mask, inBoundsAttr);
1789 std::function<LogicalResult(vector::ContractionOp op)>;
1794 filter(std::move(constraint)) {}
1798 if (failed(filter(op)))
1802 Value lhs = op.getLhs();
1804 Value res = op.getAcc();
1808 auto infer = [&](MapList m) {
1815 static constexpr std::array<int64_t, 2> perm = {1, 0};
1816 auto iteratorTypes = op.getIteratorTypes().getValue();
1818 if (iteratorTypes.size() != 3 ||
1825 const auto canonicalForm = infer({{m, k}, {n, k}, {m, n}});
1826 if (maps == canonicalForm)
1831 auto createTranspose = [&rewriter, loc](
Value mat) ->
Value {
1832 if (
auto sext = mat.getDefiningOp<arith::ExtSIOp>()) {
1834 vector::TransposeOp::create(rewriter, loc, sext.getIn(), perm);
1835 VectorType newType =
1836 cast<VectorType>(trans.
getType())
1837 .clone(cast<VectorType>(mat.getType()).getElementType());
1838 return arith::ExtSIOp::create(rewriter, loc, newType, trans);
1840 if (
auto zext = mat.getDefiningOp<arith::ExtUIOp>()) {
1842 vector::TransposeOp::create(rewriter, loc,
zext.getIn(), perm);
1843 VectorType newType =
1844 VectorType::get(cast<VectorType>(trans.
getType()).getShape(),
1845 cast<VectorType>(mat.getType()).getElementType());
1846 return arith::ExtUIOp::create(rewriter, loc, newType, trans);
1848 return vector::TransposeOp::create(rewriter, loc, mat, perm);
1851 if (maps == infer({{m, k}, {k, n}, {m, n}})) {
1852 rhs = createTranspose(
rhs);
1853 }
else if (maps == infer({{k, m}, {n, k}, {m, n}})) {
1854 lhs = createTranspose(lhs);
1855 }
else if (maps == infer({{k, m}, {k, n}, {m, n}})) {
1856 rhs = createTranspose(
rhs);
1857 lhs = createTranspose(lhs);
1858 }
else if (maps == infer({{k, m}, {k, n}, {n, m}})) {
1859 std::swap(
rhs, lhs);
1860 rhs = createTranspose(
rhs);
1861 lhs = createTranspose(lhs);
1862 }
else if (maps == infer({{k, m}, {n, k}, {n, m}})) {
1863 std::swap(
rhs, lhs);
1864 rhs = createTranspose(
rhs);
1865 }
else if (maps == infer({{m, k}, {k, n}, {n, m}})) {
1866 std::swap(lhs,
rhs);
1867 lhs = createTranspose(lhs);
1868 }
else if (maps == infer({{m, k}, {n, k}, {n, m}})) {
1869 std::swap(lhs,
rhs);
1875 op.getIteratorTypes());
1900template <
typename ExtOp>
1908 auto lhsDefOp = contractOp.getLhs().getDefiningOp<ExtOp>();
1909 auto rhsDefOp = contractOp.getRhs().getDefiningOp<ExtOp>();
1911 if (!lhsDefOp || !rhsDefOp) {
1913 "no defining op on contract operands");
1917 contractOp, lhsDefOp->getOperand(0), rhsDefOp->getOperand(0),
1918 contractOp.getAcc(), contractOp.getIndexingMapsAttr(),
1919 contractOp.getIteratorTypesAttr());
1941 if (op.getKind() != vector::CombiningKind::ADD)
1949 if (!
acc.getType().isIntOrFloat())
1952 auto parentReduction =
acc.getDefiningOp<vector::ReductionOp>();
1953 if (!parentReduction)
1958 if (isa<IntegerType>(
acc.getType())) {
1960 loc, parentReduction.getVector(), op.getVector());
1962 vAdd = arith::AddFOp::create(rewriter, loc, parentReduction.getVector(),
1966 parentReduction.getAcc());
1977 auto inVecShape = inVecTy.getShape();
1980 for (
auto [dim, isScalable] :
1981 llvm::zip_equal(inVecShape, inVecTy.getScalableDims())) {
1982 if (dim == 1 && !isScalable)
1985 newShape.push_back(dim);
1986 newScalableDims.push_back(isScalable);
1989 if (newShape.empty()) {
1990 newShape.push_back(1);
1991 newScalableDims.push_back(
false);
1994 return VectorType::get(newShape, inVecTy.getElementType(), newScalableDims);
2031 if (!resultVectorType)
2038 if (!sourceVectorType)
2040 if (sourceVectorType.getRank() < 2)
2046 auto opVectorType = cast<VectorType>(operand.getType());
2048 if (newVType == opVectorType)
2051 auto opSC = vector::ShapeCastOp::create(rewriter, loc, newVType, operand);
2052 newOperands.push_back(opSC);
2055 VectorType newResultVectorType =
2094 VectorType sourceType = op.getSourceVectorType();
2095 VectorType sourceTypeWithoutUnitDims =
2098 if (sourceType == sourceTypeWithoutUnitDims)
2105 for (
auto [i, dim] : llvm::enumerate(sourceDims)) {
2106 droppedDimsBefore[i] = droppedDims;
2107 if (dim == std::make_tuple(1,
false))
2115 if (sourceDims[idx] == std::make_tuple(1,
false))
2117 newPerm.push_back(idx - droppedDimsBefore[idx]);
2123 if (newPerm.empty()) {
2124 newPerm.push_back(0);
2129 auto dropDimsShapeCast = vector::ShapeCastOp::create(
2130 rewriter, loc, sourceTypeWithoutUnitDims, op.getVector());
2132 auto transposeWithoutUnitDims =
2133 vector::TransposeOp::create(rewriter, loc, dropDimsShapeCast, newPerm);
2136 op, op.getResultVectorType(), transposeWithoutUnitDims);
2173 for (
OpOperand &operand : forOp.getInitArgsMutable()) {
2174 auto vectorType = dyn_cast<VectorType>(operand.get().getType());
2179 if (vectorType == newVectorType)
2184 return vector::ShapeCastOp::create(
b, loc, type, source);
2188 castFn(rewriter, forOp.getLoc(), newVectorType, operand.get());
2190 replaceAndCastForOpIterArg(rewriter, forOp, operand,
2217 if (op.getKind() != vector::CombiningKind::ADD)
2220 Type elemType = op.getSourceVectorType().getElementType();
2223 if (!isa<FloatType>(elemType))
2226 auto vAdd = op.getVector().getDefiningOp<arith::AddFOp>();
2229 auto addLhs = vAdd.getLhs().getDefiningOp<arith::AddFOp>();
2236 auto newAdd = arith::AddFOp::create(rewriter, vAdd.getLoc(),
2237 addLhs.getLhs(), vAdd.getRhs());
2256 unsigned maxNumElementsToExtract,
2259 maxNumElementsToExtract(maxNumElementsToExtract) {}
2263 VectorType type = op.getSourceVectorType();
2264 if (type.isScalable() || op.isMasked())
2266 assert(type.getRank() == 1 &&
"Expected a 1-d vector");
2268 int64_t numElems = type.getNumElements();
2269 if (numElems > maxNumElementsToExtract) {
2271 op, llvm::formatv(
"has too many vector elements ({0}) to break down "
2272 "(max allowed: {1})",
2273 numElems, maxNumElementsToExtract));
2278 for (
auto [idx, extractedElem] : llvm::enumerate(extracted))
2279 extractedElem = vector::ExtractOp::create(rewriter, loc, op.getVector(),
2282 Value res = extracted.front();
2283 for (
auto extractedElem : llvm::drop_begin(extracted))
2285 extractedElem, op.getFastmathAttr());
2288 op.getFastmathAttr());
2295 unsigned maxNumElementsToExtract = 0;
2314template <
typename MulOpType>
2319 bool isValidBroadcastSource(vector::BroadcastOp broadcastOp)
const {
2322 if (!broadcastOp.computeBroadcastedUnitDims().empty())
2325 auto srcType = dyn_cast<VectorType>(broadcastOp.getSourceType());
2326 return srcType && srcType.getRank() != 2;
2331 auto resType = llvm::dyn_cast<VectorType>(mulOp.getResult().getType());
2334 if (resType.getRank() != 2)
2339 auto matchOuterProduct =
2341 Value operandB) -> FailureOr<vector::OuterProductOp> {
2342 auto transposedLhs = operandA.
getDefiningOp<vector::TransposeOp>();
2347 if (permutation.size() != 2 || permutation[0] != 1 || permutation[1] != 0)
2350 auto broadcastedLhs =
2351 transposedLhs.getVector().getDefiningOp<vector::BroadcastOp>();
2352 if (!broadcastedLhs || !isValidBroadcastSource(broadcastedLhs))
2355 auto broadcastedRhs = operandB.getDefiningOp<vector::BroadcastOp>();
2356 if (!broadcastedRhs || !isValidBroadcastSource(broadcastedRhs))
2359 return vector::OuterProductOp::create(
2360 rewriter, mulOp->getLoc(), resType, broadcastedLhs.getSource(),
2361 broadcastedRhs.getSource(),
Value(), vector::CombiningKind::ADD);
2364 Value lhs = mulOp->getOperand(0), rhs = mulOp->getOperand(1);
2365 auto maybeOuterP = matchOuterProduct(lhs, rhs);
2367 if (failed(maybeOuterP))
2368 maybeOuterP = matchOuterProduct(rhs, lhs);
2369 if (failed(maybeOuterP))
2371 rewriter.
replaceOp(mulOp, maybeOuterP->getResult());
2385void mlir::vector::populateVectorMaskMaterializationPatterns(
2388 patterns.
add<VectorCreateMaskOpConversion,
2389 MaterializeTransferMask<vector::TransferReadOp>,
2390 MaterializeTransferMask<vector::TransferWriteOp>>(
2391 patterns.
getContext(), force32BitVectorIndices, benefit);
2395void mlir::vector::populateDropUnitDimWithShapeCastPatterns(
2401void mlir::vector::populateBubbleVectorBitCastOpPatterns(
2403 patterns.
add<BubbleDownVectorBitCastForExtract,
2404 BubbleDownBitCastForStridedSliceExtract,
2405 BubbleUpBitCastForInsert, BubbleUpBitCastForStridedSliceInsert>(
2409void mlir::vector::populateBreakDownVectorBitCastOpPatterns(
2411 std::function<
bool(vector::BitCastOp)> controlFn,
PatternBenefit benefit) {
2413 std::move(controlFn), benefit);
2418 std::function<LogicalResult(vector::ContractionOp)> constraint,
2421 std::move(constraint));
2426 patterns.
add<MultiReduceToContract, CombineContractBroadcastMask,
2427 CombineContractABTranspose, CombineContractResultTranspose>(
2440 patterns.
add<ReorderElementwiseOpsOnTranspose, ReorderCastOpsOnBroadcast,
2441 ReorderElementwiseOpsOnBroadcast, ExtractOpFromElementwise>(
2448 patterns.
add<ExtractOpFromLoad, StoreOpFromBroadcast>(patterns.
getContext(),
2452void mlir::vector::populateChainedVectorReductionFoldingPatterns(
2459void mlir::vector::populateBreakDownVectorReductionPatterns(
2463 maxNumElementsToExtract, benefit);
2468 patterns.
add<FoldArithToVectorOuterProduct<arith::MulFOp>,
2469 FoldArithToVectorOuterProduct<arith::MulIOp>>(
2477#include "mlir/Dialect/Vector/Transforms/VectorTransformsEnums.cpp.inc"
static uint64_t zext(uint32_t arg)
*if copies could not be generated due to yet unimplemented cases *copyInPlacementStart and copyOutPlacementStart in copyPlacementBlock *specify the insertion points where the incoming copies and outgoing should be the output argument nBegin is set to its * replacement(set to `begin` if no invalidation happens). Since outgoing *copies could have been inserted at `end`
static Value broadcast(Location loc, Value toBroadcast, unsigned numElements, const TypeConverter &typeConverter, ConversionPatternRewriter &rewriter)
Broadcasts the value to vector with numElements number of elements.
Drop inner most contiguous unit dimensions from transfer_read operand.
Drop inner most contiguous unit dimensions from transfer_write operand. E.g., vector....
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 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...
static AffineMap getPermutationMap(ArrayRef< unsigned > permutation, MLIRContext *context)
Returns an AffineMap representing a permutation.
AffineMap compose(AffineMap map) const
Returns the AffineMap resulting from composing this with map.
IntegerAttr getIndexAttr(int64_t value)
AffineMap getMultiDimIdentityMap(unsigned rank)
IntegerType getIntegerType(unsigned width)
TypedAttr getZeroAttr(Type type)
AffineExpr getAffineDimExpr(unsigned position)
DenseIntElementsAttr getI32VectorAttr(ArrayRef< int32_t > values)
DenseIntElementsAttr getI64VectorAttr(ArrayRef< int64_t > values)
ArrayAttr getArrayAttr(ArrayRef< Attribute > value)
MLIRContext * getContext() const
ArrayAttr getI64ArrayAttr(ArrayRef< int64_t > values)
ArrayAttr getBoolArrayAttr(ArrayRef< bool > values)
ArrayAttr getAffineMapArrayAttr(ArrayRef< AffineMap > values)
DenseElementsAttr resizeSplat(ShapedType newType)
Return a new DenseElementsAttr that has the same data as the current attribute, but with a different ...
std::enable_if_t<!std::is_base_of< Attribute, T >::value||std::is_same< Attribute, T >::value, T > getSplatValue() const
Return the splat value for this attribute.
An attribute that represents a reference to a dense integer vector or tensor object.
static DenseIntElementsAttr get(const ShapedType &type, Arg &&arg)
Get an instance of a DenseIntElementsAttr with the given arguments.
void map(Value from, Value to)
Inserts a new mapping for 'from' to 'to'.
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.
This class helps build Operations.
Operation * clone(Operation &op, IRMapping &mapper)
Creates a deep copy of the specified operation, remapping any operands that use values outside of the...
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
void createOrFold(SmallVectorImpl< Value > &results, Location location, Args &&...args)
Create an operation of specific op type at the current insertion point, and immediately try to fold i...
Operation * create(const OperationState &state)
Creates an operation given the fields represented as an OperationState.
This class represents an operand of an operation.
OpTraitRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting again...
OpTraitRewritePattern(MLIRContext *context, PatternBenefit benefit=1)
Operation is the basic unit of execution within MLIR.
Value getOperand(unsigned idx)
bool hasTrait()
Returns true if the operation was registered with a particular trait, e.g.
bool hasOneUse()
Returns true if this operation has exactly one use.
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
unsigned getNumRegions()
Returns the number of regions held by this operation.
Location getLoc()
The source location the operation was defined or derived from.
unsigned getNumOperands()
Attribute getPropertiesAsAttribute()
Return the properties converted to an attribute.
OperationName getName()
The name of an operation is the key identifier for it.
DictionaryAttr getDiscardableAttrDictionary()
Return all of the discardable attributes on this operation as a DictionaryAttr.
operand_type_range getOperandTypes()
result_type_range getResultTypes()
operand_range getOperands()
Returns an iterator on the underlying Value's.
result_range getResults()
unsigned getNumResults()
Return the number of results held by this operation.
This class represents the benefit of a pattern match in a unitless scheme that ranges from 0 (very li...
unsigned short getBenefit() const
If the corresponding pattern can match, return its benefit. If the.
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...
virtual void eraseOp(Operation *op)
This method erases an operation that is known to have no uses.
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,...
void modifyOpInPlace(Operation *root, CallableT &&callable)
This method is a utility wrapper around an in-place modification of an operation.
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
This class provides an abstraction over the various different ranges of value types.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
bool isIntOrFloat() const
Return true if this is an integer (of any signedness) or a float type.
unsigned getIntOrFloatBitWidth() const
Return the bit width of an integer or a float type, assert failure on other types.
bool isSignlessIntOrIndexOrFloat() const
Return true if this is a signless integer, index, or float type.
This class provides an abstraction over the different types of ranges over Values.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
void setType(Type newType)
Mutate the type of this Value to be of the specified type.
Type getType() const
Return the type of this value.
bool hasOneUse() const
Returns true if this value has exactly one use.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
bool hasElementwiseMappableTraits(Operation *op)
Together, Elementwise, Scalarizable, Vectorizable, and Tensorizable provide an easy way for scalar op...
SmallVector< OpFoldResult > getMixedSizes(OpBuilder &builder, Location loc, Value value)
Return the dimensions of the given memref value.
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.
auto getDims(VectorType vType)
Returns a range over the dims (size and scalability) of a VectorType.
void populateElementwiseToVectorOpsPatterns(RewritePatternSet &patterns)
Collect a set of patterns that fold elementwise op on vectors to the vector dialect.
AffineMap getTransferMinorIdentityMap(ShapedType shapedType, VectorType vectorType)
Build the default minor identity map suitable for a vector transfer.
void populateDropInnerMostUnitDimsXferOpPatterns(RewritePatternSet &patterns, PatternBenefit benefit=1)
Collect a set of patterns to collapse the most inner unit dims in xfer Ops.
bool isParallelIterator(Attribute attr)
Returns true if attr has "parallel" iterator type semantics.
void populateFoldArithExtensionPatterns(RewritePatternSet &patterns)
Collect a set of patterns that fold arithmetic extension on floating point into vector contract for t...
void populateVectorContractCanonicalizeMatmulToMMT(RewritePatternSet &patterns, std::function< LogicalResult(vector::ContractionOp)> constraint=[](vector::ContractionOp) { return success();}, PatternBenefit=1)
Canonicalization of a vector.contract a, b, c with row-major matmul semantics to a contraction with M...
void populateSinkVectorOpsPatterns(RewritePatternSet &patterns, PatternBenefit benefit=1)
Patterns that remove redundant Vector Ops by re-ordering them with e.g.
void populateVectorReductionToContractPatterns(RewritePatternSet &patterns, PatternBenefit benefit=1)
Collect patterns to convert reduction op to vector.contract and fold transpose/broadcast ops into the...
Value createOrFoldDimOp(OpBuilder &b, Location loc, Value source, int64_t dim)
Helper function that creates a memref::DimOp or tensor::DimOp depending on the type of source.
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
Type getType(OpFoldResult ofr)
Returns the int type of the integer in ofr.
void bindDims(MLIRContext *ctx, AffineExprTy &...exprs)
Bind a list of AffineExpr references to DimExpr at positions: [0 .
AffineMap inversePermutation(AffineMap map)
Returns a map of codomain to domain dimensions such that the first codomain dimension for a particula...
Value getValueOrCreateCastToIndexLike(OpBuilder &b, Location loc, Type targetType, Value value)
Create a cast from an index-like value (index or integer) to another index-like value.
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
bool isZeroInteger(OpFoldResult v)
Return "true" if v is an integer value/attribute with constant value 0.
detail::constant_float_predicate_matcher m_AnyZeroFloat()
Matches a constant scalar / vector splat / tensor splat float (both positive and negative) zero.
Value getValueOrCreateConstantIndexOp(OpBuilder &b, Location loc, OpFoldResult ofr)
Converts an OpFoldResult to a Value.
AffineMap compressDims(AffineMap map, const llvm::SmallBitVector &unusedDims)
Drop the dims that are listed in unusedDims.
llvm::SmallBitVector getUnusedDimsBitVector(ArrayRef< AffineMap > maps)
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
LogicalResult matchAndRewrite(vector::ReductionOp op, PatternRewriter &rewriter) const override
BreakDownVectorReduction(MLIRContext *context, unsigned maxNumElementsToExtract, PatternBenefit benefit)
Canonicalization of a vector.contract a, b, c with row-major matmul semantics to a contraction suitab...
LogicalResult matchAndRewrite(vector::ContractionOp op, PatternRewriter &rewriter) const override
std::function< LogicalResult(vector::ContractionOp op)> FilterConstraintType
CanonicalizeContractMatmulToMMT(MLIRContext *context, PatternBenefit benefit, FilterConstraintType constraint)
Pattern to fold chained reduction to a series of vector additions and a final reduction....
LogicalResult matchAndRewrite(vector::ReductionOp op, PatternRewriter &rewriter) const override
For vectors with at least one unit dim, replaces: elementwise(a, b) with: sc_a = shape_cast(a) sc_b =...
OpTraitRewritePattern(MLIRContext *context, PatternBenefit benefit=1)
LogicalResult matchAndRewrite(Operation *op, PatternRewriter &rewriter) const override
Attempt to match against code rooted at the specified operation, which is the same operation code as ...
A pattern to drop unit dims from the iter_args of an scf.for.
LogicalResult matchAndRewrite(scf::ForOp forOp, PatternRewriter &rewriter) const override
A pattern to drop unit dims from vector.transpose.
LogicalResult matchAndRewrite(vector::TransposeOp op, PatternRewriter &rewriter) const override
Pattern to fold arithmetic extensions on floating point data types into vector contraction operations...
LogicalResult matchAndRewrite(vector::ContractionOp contractOp, PatternRewriter &rewriter) const override
Pattern to eliminate redundant zero-constants added to reduction operands. It's enough for there to b...
LogicalResult matchAndRewrite(vector::ReductionOp op, PatternRewriter &rewriter) const override
OpInterfaceRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting a...
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;.
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})
This represents an operation in an abstracted form, suitable for use with the builder APIs.
Attribute propertiesAttr
This Attribute is used to opaquely construct the properties of the operation.
LogicalResult matchAndRewrite(Operation *op, PatternRewriter &rewriter) const final
Wrapper around the RewritePattern method that passes the derived op type.
A pattern for ops that implement MaskableOpInterface and that might be masked (i.e.