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"
44template <
typename IntType>
46 return llvm::to_vector<4>(llvm::map_range(
47 arrayAttr.getAsRange<IntegerAttr>(),
48 [](IntegerAttr attr) { return static_cast<IntType>(attr.getInt()); }));
80struct MultiReduceToContract
84 LogicalResult matchAndRewrite(vector::MultiDimReductionOp reduceOp,
85 PatternRewriter &rewriter)
const override {
86 if (reduceOp.getKind() != vector::CombiningKind::ADD)
88 Operation *mulOp = reduceOp.getSource().getDefiningOp();
89 if (!mulOp || !isa<arith::MulIOp, arith::MulFOp>(mulOp))
91 SmallVector<bool> reductionMask = reduceOp.getReductionMask();
93 SmallVector<AffineExpr> exprs;
94 SmallVector<vector::IteratorType> iteratorTypes;
95 for (
const auto &isReduceDim : llvm::enumerate(reductionMask)) {
96 if (!isReduceDim.value()) {
97 iteratorTypes.push_back(vector::IteratorType::parallel);
100 iteratorTypes.push_back(vector::IteratorType::reduction);
105 0, exprs, reduceOp.getContext());
110 iteratorTypes, [&](IteratorType t) -> mlir::Attribute {
111 return IteratorTypeAttr::get(rewriter.getContext(), t);
140struct CombineContractABTranspose final
144 LogicalResult matchAndRewrite(vector::ContractionOp contractOp,
145 PatternRewriter &rewriter)
const override {
146 SmallVector<AffineMap> maps =
147 llvm::to_vector<4>(contractOp.getIndexingMapsArray());
148 Value
lhs = contractOp.getLhs();
149 Value
rhs = contractOp.getRhs();
151 bool changed =
false;
152 for (Value *operand : {&
lhs, &
rhs}) {
153 AffineMap &map = maps[index++];
154 auto transposeOp = operand->getDefiningOp<vector::TransposeOp>();
158 transposeOp.getPermutation(), contractOp.getContext());
160 *operand = transposeOp.getVector();
166 contractOp,
lhs,
rhs, contractOp.getAcc(),
204struct CombineContractResultTranspose final
208 LogicalResult matchAndRewrite(vector::TransposeOp resTOp,
209 PatternRewriter &rewriter)
const override {
210 auto contractOp = resTOp.getVector().getDefiningOp<vector::ContractionOp>();
211 if (!contractOp || !contractOp->hasOneUse())
214 auto accTOp = contractOp.getAcc().getDefiningOp<vector::TransposeOp>();
218 MLIRContext *context = contractOp.getContext();
219 auto maps = llvm::to_vector<3>(contractOp.getIndexingMapsArray());
220 AffineMap contractMap = maps.back();
231 auto combinedResMap = resTMap.compose(contractMap);
238 maps.back() = combinedResMap;
241 resTOp, contractOp.getLhs(), contractOp.getRhs(), accTOp.getVector(),
274FailureOr<Value> combineContractAndBroadcast(vector::ContractionOp contractOp,
275 MaskingOpInterface maskingOp,
278 llvm::to_vector<4>(contractOp.getIndexingMapsArray());
282 bool changed =
false;
288 auto sc = operand->getDefiningOp<vector::ShapeCastOp>();
289 auto broadcast = operand->getDefiningOp<vector::BroadcastOp>();
293 if (sc && !sc.isBroadcastLike())
295 contractOp,
"Operand defined via vector.shape_cast that has "
296 "non-broadcast semantics");
299 VectorType srcType = sc ? sc.getSourceVectorType()
300 : dyn_cast<VectorType>(
broadcast.getSourceType());
302 sc ? sc.getResultVectorType() :
broadcast.getResultVectorType();
306 if (!srcType || srcType.getRank() >= resType.getRank())
308 int64_t rankDiff = resType.getRank() - srcType.getRank();
309 bool innerDimBroadcast =
false;
311 for (
const auto &dim : llvm::enumerate(srcType.getShape())) {
312 if (dim.value() != resType.getDimSize(rankDiff + dim.index())) {
313 innerDimBroadcast =
true;
320 if (innerDimBroadcast)
325 bool nonUnitDimReductionBroadcast =
false;
326 for (
int64_t i = 0; i < rankDiff; ++i) {
327 if (resType.getDimSize(i) != 1 &&
330 nonUnitDimReductionBroadcast =
true;
334 if (nonUnitDimReductionBroadcast)
338 contractOp.getContext());
339 map = broadcastMap.
compose(map);
355 for (
unsigned i = 0, e = unusedDimsBitVector.size(); i < e; ++i) {
356 if (!unusedDimsBitVector.test(i))
357 iterators.push_back(contractOp.getIteratorTypes().getValue()[i]);
363 VectorType oldMaskType;
364 bool isAnyUnusedDimNonUnit =
false;
366 oldMaskType = cast<VectorType>(maskingOp.getMask().getType());
367 for (
unsigned i = 0, e = unusedDimsBitVector.size(); i < e; ++i) {
368 if (unusedDimsBitVector.test(i) && oldMaskType.getShape()[i] != 1) {
369 isAnyUnusedDimNonUnit =
true;
380 bool hasReductionIteratorApplyingOnBothSides =
false;
381 for (
unsigned i = 0; i < iterators.size(); ++i) {
385 hasReductionIteratorApplyingOnBothSides =
true;
389 if (!hasReductionIteratorApplyingOnBothSides)
397 Operation *newOp = vector::ContractionOp::create(
398 rewriter, contractOp.getLoc(),
lhs,
rhs, contractOp.getAcc(),
403 if (isAnyUnusedDimNonUnit)
405 "Cannont drop non-unit mask dim.");
406 assert(unusedDimsBitVector.size() ==
407 static_cast<size_t>(oldMaskType.getRank()) &&
408 "The mask rank is incorrect!");
412 Value mask = maskingOp.getMask();
413 if (unusedDimsBitVector.count() != 0) {
421 oldMaskType.getShape().drop_front(unusedDimsBitVector.count());
422 auto newShapeScalableDims =
423 oldMaskType.getScalableDims().drop_front(unusedDimsBitVector.count());
424 VectorType maskOpType =
425 VectorType::get(newShape, rewriter.
getI1Type(), newShapeScalableDims);
426 mask = vector::ShapeCastOp::create(rewriter, contractOp.getLoc(),
427 maskOpType, maskingOp.getMask())
436struct CombineContractBroadcastMask
438 using MaskableOpRewritePattern::MaskableOpRewritePattern;
441 matchAndRewriteMaskableOp(vector::ContractionOp contractOp,
442 MaskingOpInterface maskingOp,
443 PatternRewriter &rewriter)
const override {
444 return combineContractAndBroadcast(contractOp, maskingOp, rewriter);
461struct ReorderCastOpsOnBroadcast
463 using OpInterfaceRewritePattern<CastOpInterface>::OpInterfaceRewritePattern;
465 LogicalResult matchAndRewrite(CastOpInterface op,
466 PatternRewriter &rewriter)
const override {
467 if (op->getNumOperands() != 1)
469 if (!isa<VectorType>(op->getResult(0).getType()))
471 auto bcastOp = op->getOperand(0).getDefiningOp<vector::BroadcastOp>();
476 if (
auto vecTy = dyn_cast<VectorType>(bcastOp.getSourceType()))
477 castResTy = vecTy.clone(castResTy);
479 rewriter.
create(op->getLoc(), op->getName().getIdentifier(),
480 bcastOp.getSource(), castResTy, op->getAttrs());
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 =
563 transposeMaps.front());
570 return llvm::map_to_vector<4>(arrayAttr.getAsRange<IntegerAttr>(),
571 [](IntegerAttr attr) { return attr.getInt(); });
583struct BubbleDownVectorBitCastForExtract
587 LogicalResult matchAndRewrite(vector::ExtractOp extractOp,
588 PatternRewriter &rewriter)
const override {
590 if (extractOp.getSourceVectorType().getRank() != 1)
593 auto castOp = extractOp.getSource().getDefiningOp<vector::BitCastOp>();
597 VectorType castSrcType = castOp.getSourceVectorType();
598 VectorType castDstType = castOp.getResultVectorType();
599 assert(castSrcType.getRank() == castDstType.getRank());
604 if (castSrcType.getNumElements() == 1)
609 if (castSrcType.getNumElements() > castDstType.getNumElements())
612 unsigned expandRatio =
613 castDstType.getNumElements() / castSrcType.getNumElements();
616 auto mixedPos = extractOp.getMixedPosition();
617 if (!mixedPos.empty() && !isa<Attribute>(mixedPos[0]))
619 uint64_t index = cast<IntegerAttr>(cast<Attribute>(mixedPos[0])).getInt();
623 Location loc = extractOp.getLoc();
624 Value packedValue = vector::ExtractOp::create(
625 rewriter, loc, castOp.getSource(), index / expandRatio);
626 Type packedVecType = VectorType::get({1}, packedValue.
getType());
627 Value zero = arith::ConstantOp::create(rewriter, loc, packedVecType,
629 packedValue = vector::InsertOp::create(rewriter, loc, packedValue, zero,
634 VectorType packedType =
635 VectorType::get({expandRatio}, castDstType.getElementType());
637 vector::BitCastOp::create(rewriter, loc, packedType, packedValue);
641 index % expandRatio);
658struct BubbleDownBitCastForStridedSliceExtract
662 LogicalResult matchAndRewrite(vector::ExtractStridedSliceOp extractOp,
663 PatternRewriter &rewriter)
const override {
664 auto castOp = extractOp.getSource().getDefiningOp<vector::BitCastOp>();
668 VectorType castSrcType = castOp.getSourceVectorType();
669 VectorType castDstType = castOp.getResultVectorType();
670 assert(castSrcType.getRank() == castDstType.getRank());
672 int64_t castSrcLastDim = castSrcType.getShape().back();
673 int64_t castDstLastDim = castDstType.getShape().back();
675 if (castSrcLastDim > castDstLastDim)
679 if (llvm::any_of(extractOp.getStrides().getAsValueRange<IntegerAttr>(),
680 [](
const APInt &val) { return !val.isOne(); }))
683 unsigned rank = extractOp.getSourceVectorType().getRank();
684 assert(castDstLastDim % castSrcLastDim == 0);
685 int64_t expandRatio = castDstLastDim / castSrcLastDim;
691 ArrayAttr newOffsets = extractOp.getOffsets();
692 if (newOffsets.size() == rank) {
693 SmallVector<int64_t> offsets = getIntValueVector(newOffsets);
694 if (offsets.back() % expandRatio != 0)
696 offsets.back() = offsets.back() / expandRatio;
701 ArrayAttr newSizes = extractOp.getSizes();
702 if (newSizes.size() == rank) {
703 SmallVector<int64_t> sizes = getIntValueVector(newSizes);
704 if (sizes.back() % expandRatio != 0)
706 sizes.back() = sizes.back() / expandRatio;
710 SmallVector<int64_t> dims =
711 llvm::to_vector<4>(cast<VectorType>(extractOp.getType()).getShape());
712 dims.back() = dims.back() / expandRatio;
713 VectorType newExtractType =
714 VectorType::get(dims, castSrcType.getElementType());
716 auto newExtractOp = vector::ExtractStridedSliceOp::create(
717 rewriter, extractOp.getLoc(), newExtractType, castOp.getSource(),
718 newOffsets, newSizes, extractOp.getStrides());
721 extractOp, extractOp.getType(), newExtractOp);
737struct BubbleUpBitCastForInsert :
public OpRewritePattern<vector::BitCastOp> {
740 LogicalResult matchAndRewrite(vector::BitCastOp bitcastOp,
741 PatternRewriter &rewriter)
const override {
742 VectorType castSrcType = bitcastOp.getSourceVectorType();
743 VectorType castDstType = bitcastOp.getResultVectorType();
746 if (castSrcType.getRank() == 0 || castSrcType.isScalable() ||
747 castDstType.isScalable())
750 int64_t castSrcLastDim = castSrcType.getShape().back();
751 int64_t castDstLastDim = castDstType.getShape().back();
752 bool isNumElemsShrink = castSrcLastDim >= castDstLastDim;
754 if (isNumElemsShrink) {
755 assert(castSrcLastDim % castDstLastDim == 0);
756 ratio = castSrcLastDim / castDstLastDim;
758 assert(castDstLastDim % castSrcLastDim == 0);
759 ratio = castDstLastDim / castSrcLastDim;
762 auto insertOp = bitcastOp.getSource().getDefiningOp<vector::InsertOp>();
767 auto insertSrcType = dyn_cast<VectorType>(insertOp.getValueToStoreType());
772 SmallVector<int64_t> srcDims(insertSrcType.getShape());
774 isNumElemsShrink ? srcDims.back() / ratio : srcDims.back() * ratio;
775 VectorType newCastSrcType =
776 VectorType::get(srcDims, castDstType.getElementType());
778 vector::BitCastOp::create(rewriter, bitcastOp.getLoc(), newCastSrcType,
779 insertOp.getValueToStore());
781 SmallVector<int64_t> dstDims(insertOp.getDestVectorType().getShape());
783 isNumElemsShrink ? dstDims.back() / ratio : dstDims.back() * ratio;
784 VectorType newCastDstType =
785 VectorType::get(dstDims, castDstType.getElementType());
788 auto newCastDstOp = vector::BitCastOp::create(
789 rewriter, bitcastOp.getLoc(), newCastDstType, insertOp.getDest());
793 bitcastOp, newCastSrcOp, newCastDstOp, insertOp.getMixedPosition());
809struct BubbleUpBitCastForStridedSliceInsert
813 LogicalResult matchAndRewrite(vector::BitCastOp bitcastOp,
814 PatternRewriter &rewriter)
const override {
815 VectorType castSrcType = bitcastOp.getSourceVectorType();
816 VectorType castDstType = bitcastOp.getResultVectorType();
817 assert(castSrcType.getRank() == castDstType.getRank());
819 if (castSrcType.getRank() == 0)
822 int64_t castSrcLastDim = castSrcType.getShape().back();
823 int64_t castDstLastDim = castDstType.getShape().back();
825 if (castSrcLastDim < castDstLastDim)
828 assert(castSrcLastDim % castDstLastDim == 0);
829 int64_t shrinkRatio = castSrcLastDim / castDstLastDim;
832 bitcastOp.getSource().getDefiningOp<vector::InsertStridedSliceOp>();
837 if (llvm::any_of(insertOp.getStrides().getAsValueRange<IntegerAttr>(),
838 [](
const APInt &val) { return !val.isOne(); }))
841 unsigned rank = insertOp.getSourceVectorType().getRank();
844 if (rank != insertOp.getDestVectorType().getRank())
848 unsigned sourceWidth = castSrcType.getElementType().getIntOrFloatBitWidth();
849 unsigned destinationWidth =
850 castDstType.getElementType().getIntOrFloatBitWidth();
851 unsigned numElements = destinationWidth / sourceWidth;
852 if (insertOp.getSourceVectorType().getNumElements() % numElements != 0)
855 ArrayAttr newOffsets = insertOp.getOffsets();
856 assert(newOffsets.size() == rank);
857 SmallVector<int64_t> offsets = getIntValueVector(newOffsets);
858 if (offsets.back() % shrinkRatio != 0)
860 offsets.back() = offsets.back() / shrinkRatio;
863 SmallVector<int64_t> srcDims =
864 llvm::to_vector<4>(insertOp.getSourceVectorType().getShape());
865 srcDims.back() = srcDims.back() / shrinkRatio;
866 VectorType newCastSrcType =
867 VectorType::get(srcDims, castDstType.getElementType());
870 vector::BitCastOp::create(rewriter, bitcastOp.getLoc(), newCastSrcType,
871 insertOp.getValueToStore());
873 SmallVector<int64_t> dstDims =
874 llvm::to_vector<4>(insertOp.getDestVectorType().getShape());
875 dstDims.back() = dstDims.back() / shrinkRatio;
876 VectorType newCastDstType =
877 VectorType::get(dstDims, castDstType.getElementType());
879 auto newCastDstOp = vector::BitCastOp::create(
880 rewriter, bitcastOp.getLoc(), newCastDstType, insertOp.getDest());
883 bitcastOp, bitcastOp.getType(), newCastSrcOp, newCastDstOp, newOffsets,
884 insertOp.getStrides());
912 BreakDownVectorBitCast(MLIRContext *context,
913 std::function<
bool(vector::BitCastOp)> controlFn,
914 PatternBenefit benefit)
915 : OpRewritePattern(context, benefit), controlFn(std::move(controlFn)) {}
917 LogicalResult matchAndRewrite(vector::BitCastOp bitcastOp,
918 PatternRewriter &rewriter)
const override {
920 if (controlFn && !controlFn(bitcastOp))
923 VectorType castSrcType = bitcastOp.getSourceVectorType();
924 VectorType castDstType = bitcastOp.getResultVectorType();
925 assert(castSrcType.getRank() == castDstType.getRank());
930 if (castSrcType.isScalable())
932 "Scalable vectors are not supported");
935 if (castSrcType.getRank() != 1)
938 int64_t castSrcLastDim = castSrcType.getShape().back();
939 int64_t castDstLastDim = castDstType.getShape().back();
941 if (castSrcLastDim < castDstLastDim)
944 assert(castSrcLastDim % castDstLastDim == 0);
945 int64_t shrinkRatio = castSrcLastDim / castDstLastDim;
947 if (castSrcLastDim == shrinkRatio)
950 Location loc = bitcastOp.getLoc();
951 Type elemType = castDstType.getElementType();
954 Value zero = arith::ConstantOp::create(rewriter, loc, elemType,
956 Value res = BroadcastOp::create(rewriter, loc, castDstType, zero);
958 SmallVector<int64_t> sliceShape = {castDstLastDim};
959 SmallVector<int64_t> strides = {1};
960 VectorType newCastDstType =
961 VectorType::get(SmallVector<int64_t>{castDstLastDim / shrinkRatio},
962 castDstType.getElementType());
964 for (
int i = 0, e = shrinkRatio; i < e; ++i) {
965 Value extracted = ExtractStridedSliceOp::create(
966 rewriter, loc, bitcastOp.getSource(),
967 ArrayRef<int64_t>{i * castDstLastDim}, sliceShape, strides);
969 BitCastOp::create(rewriter, loc, newCastDstType, extracted);
970 res = InsertStridedSliceOp::create(
971 rewriter, loc, bitcast, res,
972 ArrayRef<int64_t>{i * castDstLastDim / shrinkRatio}, strides);
979 std::function<bool(BitCastOp)> controlFn;
982static bool haveSameShapeAndScaling(
Type t,
Type u) {
983 auto tVec = dyn_cast<VectorType>(t);
984 auto uVec = dyn_cast<VectorType>(u);
991 return tVec.getShape() == uVec.getShape() &&
992 tVec.getScalableDims() == uVec.getScalableDims();
997static Type cloneOrReplace(
Type type,
Type newElementType) {
998 if (
auto shapedType = dyn_cast<ShapedType>(type)) {
999 return shapedType.clone(newElementType);
1001 return newElementType;
1006static Value getBroadcastLikeSource(
Value value) {
1012 if (
auto broadcast = dyn_cast<vector::BroadcastOp>(op))
1031struct ReorderElementwiseOpsOnBroadcast final
1034 LogicalResult matchAndRewrite(Operation *op,
1035 PatternRewriter &rewriter)
const override {
1043 op,
"Op doesn't have ElementwiseMappableTraits");
1047 Type resultElemType = resultType.getElementType();
1050 Value broadcastSource;
1052 Operation *definingOp = operand.getDefiningOp();
1055 if (definingOp->
hasTrait<OpTrait::ConstantLike>())
1057 broadcastSource = getBroadcastLikeSource(operand);
1060 if (!broadcastSource)
1062 Type unbroadcastResultType =
1063 cloneOrReplace(broadcastSource.
getType(), resultElemType);
1070 if (isa<vector::FMAOp>(op) && !isa<VectorType>(unbroadcastResultType)) {
1072 op,
"Op only accepts vector types, but the broadcast source is a "
1080 if (!llvm::all_of(op->
getOperands(), [broadcastSource](Value val) {
1081 if (auto source = getBroadcastLikeSource(val))
1082 return haveSameShapeAndScaling(source.getType(),
1083 broadcastSource.getType());
1084 SplatElementsAttr splatConst;
1085 return matchPattern(val, m_Constant(&splatConst));
1089 "not all operands are constants or broadcasts from the same type");
1093 SmallVector<Value> srcValues;
1096 SplatElementsAttr splatConst;
1100 Type newType = cloneOrReplace(unbroadcastResultType, elementType);
1101 if (
auto newTypeShaped = dyn_cast<ShapedType>(newType)) {
1106 Operation *newConstOp =
1107 operand.getDefiningOp()->getDialect()->materializeConstant(
1108 rewriter, newConst, newType, operand.getLoc());
1109 srcValues.push_back(newConstOp->
getResult(0));
1111 srcValues.push_back(operand.getDefiningOp()->getOperand(0));
1116 Operation *elementwiseOp =
1118 unbroadcastResultType, op->
getAttrs());
1122 op, resultType, elementwiseOp->
getResults());
1144class ExtractOpFromElementwise final
1149 LogicalResult matchAndRewrite(vector::ExtractOp op,
1150 PatternRewriter &rewriter)
const override {
1151 Operation *eltwise = op.getSource().getDefiningOp();
1156 isa<vector::FMAOp>(eltwise))
1170 if (!op.getDynamicPosition().empty())
1172 op,
"dynamic position not yet implemented");
1174 Type dstType = op.getType();
1176 OpBuilder::InsertionGuard g(rewriter);
1180 Location loc = eltwise->
getLoc();
1181 SmallVector<OpFoldResult> pos = op.getMixedPosition();
1183 Value newArg = vector::ExtractOp::create(rewriter, loc, arg, pos);
1184 mapping.
map(arg, newArg);
1187 Operation *newEltwise = rewriter.
clone(*eltwise, mapping);
1198static bool isSupportedMemSinkElementType(
Type type) {
1199 if (isa<IndexType>(type))
1220class ExtractOpFromLoad final :
public OpRewritePattern<vector::ExtractOp> {
1225 PatternRewriter &rewriter)
const override {
1226 auto loadOp = op.getSource().getDefiningOp<vector::LoadOp>();
1231 if (!loadOp->hasOneUse())
1234 VectorType loadVecType = loadOp.getVectorType();
1235 if (loadVecType.isScalable())
1237 "scalable vectors are not supported");
1239 MemRefType memType = loadOp.getMemRefType();
1243 if (!isSupportedMemSinkElementType(memType.getElementType()))
1246 int64_t rankOffset = memType.getRank() - loadVecType.getRank();
1250 auto extractVecType = dyn_cast<VectorType>(op.getResult().getType());
1251 int64_t finalRank = 0;
1253 finalRank = extractVecType.getRank();
1255 SmallVector<Value>
indices = loadOp.getIndices();
1256 SmallVector<OpFoldResult> extractPos = op.getMixedPosition();
1261 OpBuilder::InsertionGuard g(rewriter);
1263 Location loc = loadOp.getLoc();
1264 ArithIndexingBuilder idxBuilderf(rewriter, loc);
1265 for (
auto i : llvm::seq<int64_t>(rankOffset,
indices.size() - finalRank)) {
1266 OpFoldResult pos = extractPos[i - rankOffset];
1274 Value base = loadOp.getBase();
1275 if (extractVecType) {
1298class StoreOpFromBroadcast final :
public OpRewritePattern<vector::StoreOp> {
1303 PatternRewriter &rewriter)
const override {
1304 VectorType vecType = op.getVectorType();
1305 if (vecType.isScalable())
1307 "scalable vectors are not supported");
1309 if (isa<VectorType>(op.getMemRefType().getElementType()))
1311 op,
"memrefs of vectors are not supported");
1313 if (vecType.getNumElements() != 1)
1315 op,
"only 1-element vectors are supported");
1317 Value toStore = op.getValueToStore();
1318 Value source = getBroadcastLikeSource(toStore);
1321 op,
"value to store is not from a broadcast");
1328 Value base = op.getBase();
1331 if (isa<VectorType>(source.
getType())) {
1351 bool force32BitVectorIndices,
int64_t dim,
1360 if (dim == 0 && force32BitVectorIndices) {
1363 }
else if (dim == 0) {
1366 }
else if (force32BitVectorIndices) {
1368 llvm::to_vector<4>(llvm::seq<int32_t>(0, dim)));
1371 llvm::to_vector<4>(llvm::seq<int64_t>(0, dim)));
1373 Value indices = arith::ConstantOp::create(rewriter, loc, indicesAttr);
1377 Value ov = vector::BroadcastOp::create(rewriter, loc,
indices.getType(), o);
1386 if (force32BitVectorIndices) {
1389 b = arith::MinSIOp::create(rewriter, loc,
b, maxBound);
1393 vector::BroadcastOp::create(rewriter, loc,
indices.getType(), bound);
1394 return arith::CmpIOp::create(rewriter, loc, arith::CmpIPredicate::slt,
1398template <
typename ConcreteOp>
1401 explicit MaterializeTransferMask(MLIRContext *context,
bool enableIndexOpt,
1402 PatternBenefit benefit = 1)
1403 : mlir::OpRewritePattern<ConcreteOp>(context, benefit),
1404 force32BitVectorIndices(enableIndexOpt) {}
1407 PatternRewriter &rewriter)
const override {
1408 if (!xferOp.hasOutOfBoundsDim())
1411 if (xferOp.getVectorType().getRank() > 1 || xferOp.getIndices().empty())
1414 Location loc = xferOp->getLoc();
1415 VectorType vtp = xferOp.getVectorType();
1422 unsigned lastIndex = llvm::size(xferOp.getIndices()) - 1;
1423 Value off = xferOp.getIndices()[lastIndex];
1426 Value
b = arith::SubIOp::create(rewriter, loc, dim.
getType(), dim, off);
1427 Value mask = vector::CreateMaskOp::create(
1429 VectorType::get(vtp.getShape(), rewriter.
getI1Type(),
1430 vtp.getScalableDims()),
1432 if (xferOp.getMask()) {
1434 mask = arith::AndIOp::create(rewriter, loc, mask, xferOp.getMask());
1438 xferOp.getMaskMutable().assign(mask);
1446 const bool force32BitVectorIndices;
1450class VectorCreateMaskOpConversion
1453 explicit VectorCreateMaskOpConversion(MLIRContext *context,
1454 bool enableIndexOpt,
1455 PatternBenefit benefit = 1)
1456 : mlir::OpRewritePattern<vector::CreateMaskOp>(context, benefit),
1457 force32BitVectorIndices(enableIndexOpt) {}
1459 LogicalResult matchAndRewrite(vector::CreateMaskOp op,
1460 PatternRewriter &rewriter)
const override {
1461 auto dstType = op.getType();
1462 if (cast<VectorType>(dstType).isScalable())
1464 int64_t rank = dstType.getRank();
1468 op, buildVectorComparison(rewriter, op, force32BitVectorIndices,
1469 rank == 0 ? 0 : dstType.getDimSize(0),
1475 const bool force32BitVectorIndices;
1479static bool allI1ConstantValuesSetTo(arith::ConstantOp constantOp,
bool value) {
1480 auto denseAttr = dyn_cast<DenseIntElementsAttr>(constantOp.getValue());
1485 assert(denseAttr.getElementType().isInteger(1) &&
"Unexpected type");
1486 return denseAttr.isSplat() && denseAttr.getSplatValue<
bool>() == value;
1504 PatternRewriter &rewriter)
const override {
1505 auto vecType = dyn_cast<VectorType>(selectOp.getType());
1506 if (!vecType || !vecType.getElementType().isInteger(1))
1510 Value cond = selectOp.getCondition();
1511 if (isa<VectorType>(cond.
getType()))
1515 if (vecType.getRank() != 1 || vecType.isScalable())
1519 if (vecType.getShape()[0] != 1)
1522 auto trueConst = selectOp.getTrueValue().getDefiningOp<arith::ConstantOp>();
1523 if (!trueConst || !allI1ConstantValuesSetTo(trueConst,
true))
1527 selectOp.getFalseValue().getDefiningOp<arith::ConstantOp>();
1528 if (!falseConst || !allI1ConstantValuesSetTo(falseConst,
false))
1532 auto elemType = rewriter.
getIntegerType(vecType.getNumElements());
1533 auto bcastType = VectorType::get({1}, elemType);
1554static FailureOr<size_t>
1558 if (failed(srcType.getStridesAndOffset(srcStrides, srcOffset)))
1561 auto isUnitDim = [](VectorType type,
int dim) {
1562 return type.getDimSize(dim) == 1 && !type.getScalableDims()[dim];
1569 int rankDiff = srcType.getRank() - vectorType.getRank();
1570 for (
int64_t i = 0, e = vectorType.getRank(); i < e; ++i) {
1573 int dim = vectorType.getRank() - i - 1;
1574 if (srcStrides[dim + rankDiff] != 1 ||
1575 srcType.getDimSize(dim + rankDiff) != 1 || !isUnitDim(vectorType, dim))
1587 LogicalResult matchAndRewrite(vector::TransferReadOp readOp,
1590 if (readOp.getTransferRank() == 0)
1593 auto srcType = dyn_cast<MemRefType>(readOp.getBase().getType());
1597 if (!readOp.getPermutationMap().isMinorIdentity())
1600 auto targetType = readOp.getVectorType();
1601 if (targetType.getRank() <= 1)
1604 FailureOr<size_t> maybeDimsToDrop =
1606 if (failed(maybeDimsToDrop))
1609 size_t dimsToDrop = maybeDimsToDrop.value();
1610 if (dimsToDrop == 0)
1613 auto inBounds = readOp.getInBoundsValues();
1614 auto droppedInBounds =
ArrayRef<bool>(inBounds).take_back(dimsToDrop);
1615 if (llvm::is_contained(droppedInBounds,
false))
1618 auto resultTargetVecType =
1619 VectorType::get(targetType.getShape().drop_back(dimsToDrop),
1620 targetType.getElementType(),
1621 targetType.getScalableDims().drop_back(dimsToDrop));
1623 auto loc = readOp.getLoc();
1630 MemRefType resultMemrefType = memref::SubViewOp::inferRankReducedResultType(
1631 srcType.getShape().drop_back(dimsToDrop), srcType, offsets, sizes,
1634 readOp.getInBoundsAttr().getValue().drop_back(dimsToDrop));
1635 Value rankedReducedView =
1636 memref::SubViewOp::create(rewriter, loc, resultMemrefType,
1637 readOp.getBase(), offsets, sizes, strides);
1639 cast<ShapedType>(rankedReducedView.
getType()), resultTargetVecType);
1642 Value mask = readOp.getMask();
1644 auto maskType = cast<VectorType>(mask.
getType());
1645 auto reducedMaskType = VectorType::get(
1646 maskType.getShape().drop_back(dimsToDrop), maskType.getElementType(),
1647 maskType.getScalableDims().drop_back(dimsToDrop));
1648 mask = rewriter.
createOrFold<vector::ShapeCastOp>(loc, reducedMaskType,
1653 rewriter, loc, resultTargetVecType, rankedReducedView,
1654 readOp.getIndices().drop_back(dimsToDrop), AffineMapAttr::get(permMap),
1655 readOp.getPadding(), mask, inBoundsAttr);
1684 LogicalResult matchAndRewrite(vector::TransferWriteOp writeOp,
1687 if (writeOp.getTransferRank() == 0)
1690 auto srcType = dyn_cast<MemRefType>(writeOp.getBase().getType());
1694 if (!writeOp.getPermutationMap().isMinorIdentity())
1697 auto targetType = writeOp.getVectorType();
1698 if (targetType.getRank() <= 1)
1701 FailureOr<size_t> maybeDimsToDrop =
1703 if (failed(maybeDimsToDrop))
1706 size_t dimsToDrop = maybeDimsToDrop.value();
1707 if (dimsToDrop == 0)
1710 auto inBounds = writeOp.getInBoundsValues();
1711 auto droppedInBounds =
ArrayRef<bool>(inBounds).take_back(dimsToDrop);
1712 if (llvm::is_contained(droppedInBounds,
false))
1715 auto resultTargetVecType =
1716 VectorType::get(targetType.getShape().drop_back(dimsToDrop),
1717 targetType.getElementType(),
1718 targetType.getScalableDims().drop_back(dimsToDrop));
1727 MemRefType resultMemrefType = memref::SubViewOp::inferRankReducedResultType(
1728 srcType.getShape().drop_back(dimsToDrop), srcType, offsets, sizes,
1731 writeOp.getInBoundsAttr().getValue().drop_back(dimsToDrop));
1733 Value rankedReducedView =
1734 memref::SubViewOp::create(rewriter, loc, resultMemrefType,
1735 writeOp.getBase(), offsets, sizes, strides);
1737 cast<ShapedType>(rankedReducedView.
getType()), resultTargetVecType);
1739 auto shapeCast = rewriter.
createOrFold<vector::ShapeCastOp>(
1740 loc, resultTargetVecType, writeOp.getVector());
1743 Value mask = writeOp.getMask();
1745 auto maskType = cast<VectorType>(mask.
getType());
1746 auto reducedMaskType = VectorType::get(
1747 maskType.getShape().drop_back(dimsToDrop), maskType.getElementType(),
1748 maskType.getScalableDims().drop_back(dimsToDrop));
1749 mask = rewriter.
createOrFold<vector::ShapeCastOp>(loc, reducedMaskType,
1754 writeOp, shapeCast, rankedReducedView,
1755 writeOp.getIndices().drop_back(dimsToDrop), AffineMapAttr::get(permMap),
1756 mask, inBoundsAttr);
1769 std::function<LogicalResult(vector::ContractionOp op)>;
1774 filter(std::move(constraint)) {}
1778 if (failed(filter(op)))
1784 Value res = op.getAcc();
1788 auto infer = [&](MapList m) {
1795 static constexpr std::array<int64_t, 2> perm = {1, 0};
1796 auto iteratorTypes = op.getIteratorTypes().getValue();
1798 if (iteratorTypes.size() != 3 ||
1805 const auto canonicalForm = infer({{m, k}, {n, k}, {m, n}});
1806 if (maps == canonicalForm)
1811 auto createTranspose = [&rewriter, loc](
Value mat) ->
Value {
1812 if (
auto sext = mat.getDefiningOp<arith::ExtSIOp>()) {
1814 vector::TransposeOp::create(rewriter, loc, sext.getIn(), perm);
1815 VectorType newType =
1816 cast<VectorType>(trans.
getType())
1817 .clone(cast<VectorType>(mat.getType()).getElementType());
1818 return arith::ExtSIOp::create(rewriter, loc, newType, trans);
1820 if (
auto zext = mat.getDefiningOp<arith::ExtUIOp>()) {
1822 vector::TransposeOp::create(rewriter, loc,
zext.getIn(), perm);
1823 VectorType newType =
1824 VectorType::get(cast<VectorType>(trans.
getType()).getShape(),
1825 cast<VectorType>(mat.getType()).getElementType());
1826 return arith::ExtUIOp::create(rewriter, loc, newType, trans);
1828 return vector::TransposeOp::create(rewriter, loc, mat, perm);
1831 if (maps == infer({{m, k}, {k, n}, {m, n}})) {
1832 rhs = createTranspose(
rhs);
1833 }
else if (maps == infer({{k, m}, {n, k}, {m, n}})) {
1834 lhs = createTranspose(
lhs);
1835 }
else if (maps == infer({{k, m}, {k, n}, {m, n}})) {
1836 rhs = createTranspose(
rhs);
1837 lhs = createTranspose(
lhs);
1838 }
else if (maps == infer({{k, m}, {k, n}, {n, m}})) {
1840 rhs = createTranspose(
rhs);
1841 lhs = createTranspose(
lhs);
1842 }
else if (maps == infer({{k, m}, {n, k}, {n, m}})) {
1844 rhs = createTranspose(
rhs);
1845 }
else if (maps == infer({{m, k}, {k, n}, {n, m}})) {
1847 lhs = createTranspose(
lhs);
1848 }
else if (maps == infer({{m, k}, {n, k}, {n, m}})) {
1855 op.getIteratorTypes());
1880template <
typename ExtOp>
1888 auto lhsDefOp = contractOp.getLhs().getDefiningOp<ExtOp>();
1889 auto rhsDefOp = contractOp.getRhs().getDefiningOp<ExtOp>();
1891 if (!lhsDefOp || !rhsDefOp) {
1893 "no defining op on contract operands");
1897 contractOp, lhsDefOp->getOperand(0), rhsDefOp->getOperand(0),
1898 contractOp.getAcc(), contractOp.getIndexingMapsAttr(),
1899 contractOp.getIteratorTypesAttr());
1921 if (op.getKind() != vector::CombiningKind::ADD)
1929 if (!
acc.getType().isIntOrFloat())
1932 auto parentReduction =
acc.getDefiningOp<vector::ReductionOp>();
1933 if (!parentReduction)
1938 if (isa<IntegerType>(
acc.getType())) {
1940 loc, parentReduction.getVector(), op.getVector());
1942 vAdd = arith::AddFOp::create(rewriter, loc, parentReduction.getVector(),
1946 parentReduction.getAcc());
1957 auto inVecShape = inVecTy.getShape();
1960 for (
auto [dim, isScalable] :
1961 llvm::zip_equal(inVecShape, inVecTy.getScalableDims())) {
1962 if (dim == 1 && !isScalable)
1965 newShape.push_back(dim);
1966 newScalableDims.push_back(isScalable);
1969 if (newShape.empty()) {
1970 newShape.push_back(1);
1971 newScalableDims.push_back(
false);
1974 return VectorType::get(newShape, inVecTy.getElementType(), newScalableDims);
2011 if (!resultVectorType)
2018 if (!sourceVectorType)
2020 if (sourceVectorType.getRank() < 2)
2026 auto opVectorType = cast<VectorType>(operand.getType());
2028 if (newVType == opVectorType)
2031 auto opSC = vector::ShapeCastOp::create(rewriter, loc, newVType, operand);
2032 newOperands.push_back(opSC);
2035 VectorType newResultVectorType =
2040 newResultVectorType, op->
getAttrs());
2075 VectorType sourceType = op.getSourceVectorType();
2076 VectorType sourceTypeWithoutUnitDims =
2079 if (sourceType == sourceTypeWithoutUnitDims)
2086 for (
auto [i, dim] : llvm::enumerate(sourceDims)) {
2087 droppedDimsBefore[i] = droppedDims;
2088 if (dim == std::make_tuple(1,
false))
2096 if (sourceDims[idx] == std::make_tuple(1,
false))
2098 newPerm.push_back(idx - droppedDimsBefore[idx]);
2104 if (newPerm.empty()) {
2105 newPerm.push_back(0);
2110 auto dropDimsShapeCast = vector::ShapeCastOp::create(
2111 rewriter, loc, sourceTypeWithoutUnitDims, op.getVector());
2113 auto transposeWithoutUnitDims =
2114 vector::TransposeOp::create(rewriter, loc, dropDimsShapeCast, newPerm);
2117 op, op.getResultVectorType(), transposeWithoutUnitDims);
2154 for (
OpOperand &operand : forOp.getInitArgsMutable()) {
2155 auto vectorType = dyn_cast<VectorType>(operand.get().getType());
2160 if (vectorType == newVectorType)
2165 return vector::ShapeCastOp::create(
b, loc, type, source);
2169 castFn(rewriter, forOp.getLoc(), newVectorType, operand.get());
2171 replaceAndCastForOpIterArg(rewriter, forOp, operand,
2198 if (op.getKind() != vector::CombiningKind::ADD)
2201 Type elemType = op.getSourceVectorType().getElementType();
2204 if (!isa<FloatType>(elemType))
2207 auto vAdd = op.getVector().getDefiningOp<arith::AddFOp>();
2210 auto addLhs = vAdd.getLhs().getDefiningOp<arith::AddFOp>();
2217 auto newAdd = arith::AddFOp::create(rewriter, vAdd.getLoc(),
2218 addLhs.getLhs(), vAdd.getRhs());
2237 unsigned maxNumElementsToExtract,
2240 maxNumElementsToExtract(maxNumElementsToExtract) {}
2244 VectorType type = op.getSourceVectorType();
2245 if (type.isScalable() || op.isMasked())
2247 assert(type.getRank() == 1 &&
"Expected a 1-d vector");
2249 int64_t numElems = type.getNumElements();
2250 if (numElems > maxNumElementsToExtract) {
2252 op, llvm::formatv(
"has too many vector elements ({0}) to break down "
2253 "(max allowed: {1})",
2254 numElems, maxNumElementsToExtract));
2259 for (
auto [idx, extractedElem] : llvm::enumerate(extracted))
2260 extractedElem = vector::ExtractOp::create(rewriter, loc, op.getVector(),
2263 Value res = extracted.front();
2264 for (
auto extractedElem : llvm::drop_begin(extracted))
2266 extractedElem, op.getFastmathAttr());
2269 op.getFastmathAttr());
2276 unsigned maxNumElementsToExtract = 0;
2295template <
typename MulOpType>
2300 bool isValidBroadcastSource(vector::BroadcastOp broadcastOp)
const {
2303 if (!broadcastOp.computeBroadcastedUnitDims().empty())
2306 auto srcType = dyn_cast<VectorType>(broadcastOp.getSourceType());
2307 return srcType && srcType.getRank() != 2;
2312 auto resType = llvm::dyn_cast<VectorType>(mulOp.getResult().getType());
2315 if (resType.getRank() != 2)
2320 auto matchOuterProduct =
2322 Value operandB) -> FailureOr<vector::OuterProductOp> {
2323 auto transposedLhs = operandA.
getDefiningOp<vector::TransposeOp>();
2328 if (permutation.size() != 2 || permutation[0] != 1 || permutation[1] != 0)
2331 auto broadcastedLhs =
2332 transposedLhs.getVector().getDefiningOp<vector::BroadcastOp>();
2333 if (!broadcastedLhs || !isValidBroadcastSource(broadcastedLhs))
2336 auto broadcastedRhs = operandB.getDefiningOp<vector::BroadcastOp>();
2337 if (!broadcastedRhs || !isValidBroadcastSource(broadcastedRhs))
2340 return vector::OuterProductOp::create(
2341 rewriter, mulOp->getLoc(), resType, broadcastedLhs.getSource(),
2342 broadcastedRhs.getSource(),
Value(), vector::CombiningKind::ADD);
2345 Value lhs = mulOp->getOperand(0),
rhs = mulOp->getOperand(1);
2346 auto maybeOuterP = matchOuterProduct(
lhs,
rhs);
2348 if (failed(maybeOuterP))
2349 maybeOuterP = matchOuterProduct(
rhs,
lhs);
2350 if (failed(maybeOuterP))
2352 rewriter.
replaceOp(mulOp, maybeOuterP->getResult());
2366void mlir::vector::populateVectorMaskMaterializationPatterns(
2369 patterns.
add<VectorCreateMaskOpConversion,
2370 MaterializeTransferMask<vector::TransferReadOp>,
2371 MaterializeTransferMask<vector::TransferWriteOp>>(
2372 patterns.
getContext(), force32BitVectorIndices, benefit);
2376void mlir::vector::populateDropUnitDimWithShapeCastPatterns(
2382void mlir::vector::populateBubbleVectorBitCastOpPatterns(
2384 patterns.
add<BubbleDownVectorBitCastForExtract,
2385 BubbleDownBitCastForStridedSliceExtract,
2386 BubbleUpBitCastForInsert, BubbleUpBitCastForStridedSliceInsert>(
2390void mlir::vector::populateBreakDownVectorBitCastOpPatterns(
2392 std::function<
bool(vector::BitCastOp)> controlFn,
PatternBenefit benefit) {
2394 std::move(controlFn), benefit);
2399 std::function<LogicalResult(vector::ContractionOp)> constraint,
2402 std::move(constraint));
2407 patterns.
add<MultiReduceToContract, CombineContractBroadcastMask,
2408 CombineContractABTranspose, CombineContractResultTranspose>(
2421 patterns.
add<ReorderElementwiseOpsOnTranspose, ReorderCastOpsOnBroadcast,
2422 ReorderElementwiseOpsOnBroadcast, ExtractOpFromElementwise>(
2429 patterns.
add<ExtractOpFromLoad, StoreOpFromBroadcast>(patterns.
getContext(),
2433void mlir::vector::populateChainedVectorReductionFoldingPatterns(
2440void mlir::vector::populateBreakDownVectorReductionPatterns(
2444 maxNumElementsToExtract, benefit);
2449 patterns.
add<FoldArithToVectorOuterProduct<arith::MulFOp>,
2450 FoldArithToVectorOuterProduct<arith::MulIOp>>(
2458#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)
StringAttr getIdentifier() const
Return the name of this operation as a StringAttr.
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.
ArrayRef< NamedAttribute > getAttrs()
Return all of the attributes on this operation.
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()
OperationName getName()
The name of an operation is the key identifier for it.
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...
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 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={})
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.