22static int64_t getNumGtOneDims(ArrayRef<int64_t> shape) {
23 return llvm::count_if(
24 shape, [](int64_t v) {
return ShapedType::isDynamic(v) || v > 1; });
29static int64_t getFirstNonUnitSizeIdx(ArrayRef<int64_t> sizes) {
30 const auto *it = llvm::find_if(sizes, [](int64_t dim) {
return dim != 1; });
31 return (it != sizes.end()) ? std::distance(sizes.begin(), it) : -1;
46template <
typename PackOrUnpackOp>
47static LogicalResult isPackOnEffectively1D(RewriterBase &rewriter,
50 auto pack = dyn_cast<linalg::PackOp>(op);
51 auto unpack = dyn_cast<linalg::UnPackOp>(op);
53 ArrayRef<int64_t> unpackedShape =
pack ?
pack->getSourceType().getShape()
54 : unpack->getDestType().getShape();
57 ArrayRef<int64_t> innerTileSizes = op->getStaticInnerTiles();
60 if (getNumGtOneDims(unpackedShape) != 1) {
61 return rewriter.notifyMatchFailure(
62 *op,
"expects non-packed domain to have at most one non-unit dims");
66 auto numNonUnitInnerTiles = getNumGtOneDims(innerTileSizes);
67 if (numNonUnitInnerTiles > 1) {
68 return rewriter.notifyMatchFailure(
69 *op,
"expects at most one non-unit inner tiles");
73 if (numNonUnitInnerTiles == 0)
77 int64_t nonUnitDimIdx = getFirstNonUnitSizeIdx(unpackedShape);
80 int64_t nonUnitTileDestDimIdx = getFirstNonUnitSizeIdx(innerTileSizes);
83 if (nonUnitTileDestDimIdx != nonUnitDimIdx) {
84 return rewriter.notifyMatchFailure(
85 *op,
"expects at most one non-unit inner tiles");
93static FailureOr<SmallVector<int64_t>>
94getTransposeOpPermutation(linalg::LinalgOp linalgOp) {
95 if (
auto transposeOp = dyn_cast<linalg::TransposeOp>(linalgOp.getOperation()))
96 return SmallVector<int64_t>(transposeOp.getPermutation());
97 if (linalgOp.getNumParallelLoops() != linalgOp.getNumLoops())
100 if (linalgOp.getNumDpsInputs() != 1 || linalgOp.getNumDpsInits() != 1)
102 auto mapRange = linalgOp.getIndexingMapsArray();
103 if (!mapRange.front().isPermutation() || !mapRange.back().isPermutation() ||
104 mapRange.front() == mapRange.back()) {
107 if (!llvm::hasSingleElement(linalgOp.getBlock()->getOperations()))
109 AffineMap outMap = mapRange.back();
110 AffineMap inMap = mapRange.front();
113 return llvm::map_to_vector(outMap.getResults(),
114 [&](AffineExpr expr) -> int64_t {
115 return *inMap.getResultPosition(expr);
120struct SimplifyPackToExpandShape :
public OpRewritePattern<PackOp> {
121 using OpRewritePattern<PackOp>::OpRewritePattern;
124 insertExpand(RewriterBase &rewriter, Location loc, Value operand,
126 ArrayRef<ReassociationIndices> reassociation)
const {
127 if (operand.getType() == newOperandType)
129 return tensor::ExpandShapeOp::create(rewriter, loc, newOperandType, operand,
135 LogicalResult isPackOnInnerMostDim(RewriterBase &rewriter,
136 PackOp packOp)
const {
137 auto outerDimsPerm = packOp.getOuterDimsPerm();
139 return rewriter.notifyMatchFailure(
141 "expects outer_dims_perm is empty or an identity permutation");
144 int64_t srcRank = packOp.getSourceRank();
145 ArrayRef<int64_t> dimsPos = packOp.getInnerDimsPos();
146 if (dimsPos.size() != 1 || (dimsPos[0] + 1 != srcRank)) {
147 return rewriter.notifyMatchFailure(
148 packOp,
"expects packing at the innermost dimension");
153 LogicalResult matchAndRewrite(PackOp packOp,
154 PatternRewriter &rewriter)
const override {
155 if (packOp.getPaddingValue())
156 return rewriter.notifyMatchFailure(packOp,
"expects no padding value");
160 if (!packOp.hasPureTensorSemantics())
163 ShapedType sourceType = packOp.getSourceType();
164 if (
failed(isPackOnInnerMostDim(rewriter, packOp)) &&
165 failed(isPackOnEffectively1D(rewriter, &packOp)) &&
166 !packOp.isLikePad()) {
170 ShapedType destType = packOp.getDestType();
175 FailureOr<Value> expanded =
176 insertExpand(rewriter, packOp.getLoc(), packOp.getSource(), destType,
179 return rewriter.notifyMatchFailure(
180 packOp,
"unable to expand source of tensor.pack");
182 rewriter.replaceOp(packOp, *expanded);
187struct SimplifyUnPackToCollapseShape :
public OpRewritePattern<UnPackOp> {
188 using OpRewritePattern<UnPackOp>::OpRewritePattern;
190 Value insertCollapse(RewriterBase &rewriter, Location loc, Value operand,
191 Type newOperandType,
ArrayAttr reassociation)
const {
192 if (operand.getType() == newOperandType)
194 return tensor::CollapseShapeOp::create(rewriter, loc, newOperandType,
195 operand, reassociation);
199 LogicalResult isUnpackOnInnerMostDim(RewriterBase &rewriter,
200 UnPackOp unpackOp)
const {
201 auto outerDimsPerm = unpackOp.getOuterDimsPerm();
203 return rewriter.notifyMatchFailure(
205 "expects outer_dims_perm is empty or an identity permutation");
208 ShapedType sourceType = unpackOp.getSourceType();
209 ShapedType destType = unpackOp.getDestType();
210 if (!sourceType.hasStaticShape() || !destType.hasStaticShape())
211 return rewriter.notifyMatchFailure(unpackOp,
"expects static shapes");
213 ArrayRef<int64_t> dimsPos = unpackOp.getInnerDimsPos();
214 if (dimsPos.size() != 1 || (dimsPos[0] + 1 != destType.getRank())) {
215 return rewriter.notifyMatchFailure(
216 unpackOp,
"expects unpacking on the innermost dimension");
222 LogicalResult matchAndRewrite(UnPackOp unpackOp,
223 PatternRewriter &rewriter)
const override {
227 if (!unpackOp.hasPureTensorSemantics())
230 ShapedType destType = unpackOp.getDestType();
231 if (
failed(isUnpackOnInnerMostDim(rewriter, unpackOp)) &&
232 failed(isPackOnEffectively1D(rewriter, &unpackOp)) &&
233 !unpackOp.isLikeUnPad()) {
237 ShapedType sourceType = unpackOp.getSourceType();
242 Value collapsed = insertCollapse(
243 rewriter, unpackOp.getLoc(), unpackOp.getSource(), destType,
245 rewriter.replaceOp(unpackOp, collapsed);
252struct FoldPadWithPackOp :
public OpRewritePattern<PackOp> {
255 : OpRewritePattern<PackOp>(context), controlFn(std::move(controlFn)) {}
257 LogicalResult matchAndRewrite(PackOp packOp,
258 PatternRewriter &rewriter)
const override {
259 auto padOp = packOp.getSource().getDefiningOp<tensor::PadOp>();
261 if (!padOp || padOp.getNofold() || !padOp.hasZeroLowPad())
265 if (controlFn && !controlFn(&packOp.getSourceMutable()))
268 Value constantPaddingValue = padOp.getConstantPaddingValue();
269 if (!constantPaddingValue)
272 if (
auto paddingValue = packOp.getPaddingValue())
281 ShapedType unpackedType = packOp.getSourceType();
282 SmallVector<int64_t> outerShapeWithoutTranspose =
284 for (
auto [pos, tileSize, high] :
285 llvm::zip_equal(packOp.getInnerDimsPos(), packOp.getStaticInnerTiles(),
286 padOp.getMixedHighPad())) {
287 if (unpackedType.isDynamicDim(pos))
289 if (ShapedType::isDynamic(outerShapeWithoutTranspose[pos]))
291 if (ShapedType::isDynamic(tileSize))
296 int64_t paddingSize = outerShapeWithoutTranspose[pos] * tileSize -
297 unpackedType.getDimSize(pos);
299 if (paddingSize + cstHigh.value() >= tileSize)
303 rewriter.replaceOpWithNewOp<PackOp>(
304 packOp, padOp.getSource(), packOp.getDest(), packOp.getInnerDimsPos(),
305 packOp.getMixedTiles(), constantPaddingValue,
306 packOp.getOuterDimsPerm());
316struct FoldUnpackWithExtractSliceOp
317 :
public OpRewritePattern<tensor::ExtractSliceOp> {
319 FoldUnpackWithExtractSliceOp(MLIRContext *context,
321 : OpRewritePattern<tensor::ExtractSliceOp>(context),
322 controlFn(std::move(controlFn)) {}
324 LogicalResult matchAndRewrite(tensor::ExtractSliceOp sliceOp,
325 PatternRewriter &rewriter)
const override {
326 auto unpackOp = sliceOp.getSource().getDefiningOp<UnPackOp>();
333 if (!unpackOp.hasPureTensorSemantics())
337 if (controlFn && !controlFn(&sliceOp.getSourceMutable()))
340 if (!unpackOp.canFoldSliceOp(sliceOp))
344 Type elementType = unpackOp.getDestType().getElementType();
345 Value output = tensor::EmptyOp::create(
346 rewriter, sliceOp.getLoc(), sliceOp.getMixedSizes(), elementType);
347 rewriter.replaceOpWithNewOp<UnPackOp>(
348 sliceOp, unpackOp.getSource(), output, unpackOp.getInnerDimsPos(),
349 unpackOp.getMixedTiles(), unpackOp.getOuterDimsPerm());
364static bool checkAndPermute(ArrayRef<int64_t> permutation,
365 ArrayRef<int64_t> inVec,
366 SmallVectorImpl<int64_t> &resVec, int64_t rank) {
368 for (
unsigned int i = 0; i < rank; ++i) {
369 int64_t remappedPosition = permutation[i];
370 if (remappedPosition >= rank)
373 remappedPosition = inVec[remappedPosition];
374 resVec.push_back(remappedPosition);
382struct FoldProducerPackWithConsumerLinalgTransposeOp
383 :
public OpInterfaceRewritePattern<linalg::LinalgOp> {
386 FoldProducerPackWithConsumerLinalgTransposeOp(
388 : OpInterfaceRewritePattern<linalg::LinalgOp>(context),
389 controlFn(std::move(controlFn)) {}
391 LogicalResult matchAndRewrite(linalg::LinalgOp linalgOp,
392 PatternRewriter &rewriter)
const override {
393 auto packOp = linalgOp->getOperand(0).getDefiningOp<PackOp>();
401 if (!packOp.hasPureTensorSemantics())
405 if (controlFn && !controlFn(&linalgOp->getOpOperand(0)))
408 FailureOr<SmallVector<int64_t>> maybePerm =
409 getTransposeOpPermutation(linalgOp);
413 auto innerDimsPos = packOp.getInnerDimsPos();
414 auto mixedInnerTiles = packOp.getMixedTiles();
415 auto outerDimsPerm = packOp.getOuterDimsPerm();
416 const auto &transposePerm = maybePerm.value();
417 SmallVector<int64_t> newOuterDimsPermVec;
418 SmallVector<int64_t> newInnerDimsPosVec;
419 SmallVector<OpFoldResult> newMixedInnerTilesVec;
420 int64_t srcRank = packOp.getSourceRank();
422 if (!checkAndPermute(transposePerm, outerDimsPerm, newOuterDimsPermVec,
424 return rewriter.notifyMatchFailure(
426 "Cannot fold in tensor.pack if a tile dimension was transposed "
427 "with a non-tile dimension in linalg.transpose.");
430 for (
unsigned int i = srcRank; i < transposePerm.size(); ++i) {
431 int64_t remappedPosition = transposePerm[i] - srcRank;
432 newMixedInnerTilesVec.push_back(mixedInnerTiles[remappedPosition]);
433 newInnerDimsPosVec.push_back(innerDimsPos[remappedPosition]);
436 Value output = packOp.createDestinationTensor(
437 rewriter, linalgOp.getLoc(), packOp.getSource(), newMixedInnerTilesVec,
438 newInnerDimsPosVec, newOuterDimsPermVec);
440 rewriter.replaceOpWithNewOp<PackOp>(
441 linalgOp, packOp.getSource(), output, newInnerDimsPosVec,
442 newMixedInnerTilesVec, packOp.getPaddingValue(), newOuterDimsPermVec);
453struct FoldConsumerPackWithProducerLinalgTransposeOp
454 :
public OpRewritePattern<PackOp> {
457 FoldConsumerPackWithProducerLinalgTransposeOp(
459 : OpRewritePattern<PackOp>(context), controlFn(std::move(controlFn)) {}
461 LogicalResult matchAndRewrite(PackOp packOp,
462 PatternRewriter &rewriter)
const override {
466 if (!packOp.hasPureTensorSemantics())
469 auto linalgOp = packOp.getSource().getDefiningOp<linalg::LinalgOp>();
474 if (controlFn && !controlFn(&packOp.getSourceMutable()))
477 FailureOr<SmallVector<int64_t>> maybePerm =
478 getTransposeOpPermutation(linalgOp);
482 auto transposePermutation = maybePerm.value();
483 auto outerDimsPerm = packOp.getOuterDimsPerm();
484 auto innerDimsPos = packOp.getInnerDimsPos();
485 SmallVector<int64_t> newInnerDimsPosVec;
486 SmallVector<int64_t> newOuterDimsPermVec =
487 llvm::to_vector(transposePermutation);
489 if (!outerDimsPerm.empty())
494 for (
auto dim : innerDimsPos)
495 newInnerDimsPosVec.push_back(transposePermutation[dim]);
497 Value output = packOp.createDestinationTensor(
498 rewriter, packOp.getLoc(), linalgOp->getOperand(0),
499 packOp.getMixedTiles(), newInnerDimsPosVec, newOuterDimsPermVec);
501 rewriter.replaceOpWithNewOp<PackOp>(
502 packOp, linalgOp->getOperand(0), output, newInnerDimsPosVec,
503 packOp.getMixedTiles(), packOp.getPaddingValue(), newOuterDimsPermVec);
514struct FoldProducerUnPackWithConsumerLinalgTransposeOp
515 :
public OpInterfaceRewritePattern<linalg::LinalgOp> {
518 FoldProducerUnPackWithConsumerLinalgTransposeOp(
520 : OpInterfaceRewritePattern<linalg::LinalgOp>(context),
521 controlFn(std::move(controlFn)) {}
523 LogicalResult matchAndRewrite(linalg::LinalgOp linalgOp,
524 PatternRewriter &rewriter)
const override {
525 auto unPackOp = linalgOp->getOperand(0).getDefiningOp<UnPackOp>();
533 if (!unPackOp.hasPureTensorSemantics())
537 if (controlFn && !controlFn(&linalgOp->getOpOperand(0)))
540 FailureOr<SmallVector<int64_t>> maybePerm =
541 getTransposeOpPermutation(linalgOp);
545 auto outerDimsPerm = unPackOp.getOuterDimsPerm();
546 auto innerDimsPos = unPackOp.getInnerDimsPos();
547 SmallVector<int64_t> newInnerDimsPosVec;
548 SmallVector<int64_t> newOuterDimsPermVec =
553 for (
auto dim : innerDimsPos)
554 newInnerDimsPosVec.push_back(newOuterDimsPermVec[dim]);
556 if (!outerDimsPerm.empty())
560 rewriter.replaceOpWithNewOp<UnPackOp>(
561 linalgOp, unPackOp.getSource(), linalgOp.getDpsInits()[0],
562 newInnerDimsPosVec, unPackOp.getMixedTiles(), newOuterDimsPermVec);
573struct FoldConsumerUnPackWithProducerLinalgTransposeOp
574 :
public OpRewritePattern<UnPackOp> {
575 using OpRewritePattern<UnPackOp>::OpRewritePattern;
578 FoldConsumerUnPackWithProducerLinalgTransposeOp(
580 : OpRewritePattern<UnPackOp>(context), controlFn(std::move(controlFn)) {}
582 LogicalResult matchAndRewrite(UnPackOp unPackOp,
583 PatternRewriter &rewriter)
const override {
587 if (!unPackOp.hasPureTensorSemantics())
590 auto linalgOp = unPackOp.getSource().getDefiningOp<linalg::LinalgOp>();
595 if (controlFn && !controlFn(&unPackOp.getSourceMutable()))
598 FailureOr<SmallVector<int64_t>> maybePerm =
599 getTransposeOpPermutation(linalgOp);
603 SmallVector<SmallVector<OpFoldResult>> unpackOpResultDims;
608 SmallVector<int64_t> inverseTransposePerm =
610 auto outerDimsPerm = unPackOp.getOuterDimsPerm();
611 auto innerDimsPos = unPackOp.getInnerDimsPos();
612 int64_t destRank = unPackOp.getSourceRank() - innerDimsPos.size();
613 auto mixedInnerTilesVec = unPackOp.getMixedTiles();
614 SmallVector<int64_t> newOuterDimsPermVec;
615 SmallVector<int64_t> newInnerDimsPosVec;
616 SmallVector<OpFoldResult> newMixedInnerTilesVec;
617 if (!checkAndPermute(inverseTransposePerm, outerDimsPerm,
618 newOuterDimsPermVec, destRank))
619 return rewriter.notifyMatchFailure(
621 "Cannot fold in tensor.unpack if a tile dimension was transposed "
622 "with a non-tile dimension in linalg.transpose.");
625 for (
unsigned int i = destRank; i < inverseTransposePerm.size(); ++i) {
626 int64_t remappedPosition = inverseTransposePerm[i] - destRank;
627 newMixedInnerTilesVec.push_back(mixedInnerTilesVec[remappedPosition]);
628 newInnerDimsPosVec.push_back(innerDimsPos[remappedPosition]);
632 cast<ShapedType>(unPackOp->getResultTypes()[0]).getElementType();
633 Value output = tensor::EmptyOp::create(rewriter, unPackOp->getLoc(),
634 unpackOpResultDims[0], elemType);
636 rewriter.replaceOpWithNewOp<UnPackOp>(
637 unPackOp, linalgOp->getOperand(0), output, newInnerDimsPosVec,
638 newMixedInnerTilesVec, newOuterDimsPermVec);
649struct FoldEmptyTensorWithPackOp :
public OpRewritePattern<PackOp> {
650 using OpRewritePattern<PackOp>::OpRewritePattern;
652 LogicalResult matchAndRewrite(PackOp packOp,
653 PatternRewriter &rewriter)
const override {
657 if (!packOp.hasPureTensorSemantics())
661 auto emptyOp = packOp.getSource().getDefiningOp<tensor::EmptyOp>();
667 if (packOp.getPaddingValue())
668 return rewriter.notifyMatchFailure(packOp,
"expects no padding value");
671 rewriter.replaceOp(packOp, packOp.getDest());
679struct FoldEmptyTensorWithUnPackOp :
public OpRewritePattern<UnPackOp> {
680 using OpRewritePattern<UnPackOp>::OpRewritePattern;
682 LogicalResult matchAndRewrite(UnPackOp unPackOp,
683 PatternRewriter &rewriter)
const override {
687 if (!unPackOp.hasPureTensorSemantics())
691 auto emptyOp = unPackOp.getSource().getDefiningOp<tensor::EmptyOp>();
696 rewriter.replaceOp(unPackOp, unPackOp.getDest());
706 patterns.
insert<FoldUnpackWithExtractSliceOp, FoldPadWithPackOp,
707 FoldProducerPackWithConsumerLinalgTransposeOp,
708 FoldConsumerPackWithProducerLinalgTransposeOp,
709 FoldConsumerUnPackWithProducerLinalgTransposeOp,
710 FoldProducerUnPackWithConsumerLinalgTransposeOp>(
715 patterns.
add<SimplifyPackToExpandShape, SimplifyUnPackToCollapseShape>(
721 patterns.
add<FoldEmptyTensorWithPackOp, FoldEmptyTensorWithUnPackOp>(
RewritePatternSet & insert(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
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.
void populateSimplifyPackAndUnpackPatterns(RewritePatternSet &patterns)
Populates patterns with patterns that simplify tensor.pack and tensor.unpack operations.
void populateFoldPackUnpackIntoTensorEmptyPatterns(RewritePatternSet &patterns)
Populates patterns with patterns that fold operations like linalg.pack and linalg....
void populateFoldIntoPackAndUnpackPatterns(RewritePatternSet &patterns, const ControlFoldIntoPackUnpackFn &controlFn=nullptr)
Populates patterns with patterns that fold operations like tensor.pad and tensor.extract_slice into t...
FailureOr< PackResult > pack(RewriterBase &rewriter, linalg::LinalgOp linalgOp, ArrayRef< OpFoldResult > packedSizes)
Implement packing of a single LinalgOp by packedSizes.
std::function< bool(OpOperand *opOperand)> ControlFoldIntoPackUnpackFn
Function type which is used to control folding operations like tensor.pad and tensor....
SmallVector< int64_t > getPackedOuterShapeWithoutTransposition(OpTy packOrUnPack)
Returns the outer shape in the packed domain before applying the transposition.
Include the generated interface declarations.
std::optional< int64_t > getConstantIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
LogicalResult reifyResultShapes(OpBuilder &b, Operation *op, ReifiedRankedShapedTypeDims &reifiedReturnShapes)
Reify the shape of the result of an operation (typically in terms of the shape of its operands).
bool isEqualConstantIntOrValue(OpFoldResult ofr1, OpFoldResult ofr2)
Return true if ofr1 and ofr2 are the same integer constant attribute values or the same SSA value.
std::optional< SmallVector< ReassociationIndices > > getReassociationIndicesForReshape(ShapedType sourceType, ShapedType targetType)
Return the reassociations maps to use to reshape given the source type and the target type when possi...
bool isIdentityPermutation(ArrayRef< int64_t > permutation)
Returns true if permutation is an identity permutation.
void applyPermutationToVector(SmallVector< T, N > &inVec, ArrayRef< int64_t > permutation)
Apply the permutation defined by permutation to inVec.
ArrayAttr getReassociationIndicesAttribute(Builder &b, ArrayRef< ReassociationIndices > reassociation)
Wraps a list of reassociations in an ArrayAttr.
SmallVector< int64_t > invertPermutationVector(ArrayRef< int64_t > permutation)
Helper method to apply to inverse a permutation.