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");
158 if (!packOp.hasPureTensorSemantics())
161 ShapedType sourceType = packOp.getSourceType();
162 if (
failed(isPackOnInnerMostDim(rewriter, packOp)) &&
163 failed(isPackOnEffectively1D(rewriter, &packOp)) &&
164 !packOp.isLikePad()) {
168 ShapedType destType = packOp.getDestType();
173 FailureOr<Value> expanded =
174 insertExpand(rewriter, packOp.getLoc(), packOp.getSource(), destType,
177 return rewriter.notifyMatchFailure(
178 packOp,
"unable to expand source of tensor.pack");
180 rewriter.replaceOp(packOp, *expanded);
185struct SimplifyUnPackToCollapseShape :
public OpRewritePattern<UnPackOp> {
186 using OpRewritePattern<UnPackOp>::OpRewritePattern;
188 Value insertCollapse(RewriterBase &rewriter, Location loc, Value operand,
189 Type newOperandType,
ArrayAttr reassociation)
const {
190 if (operand.getType() == newOperandType)
192 return tensor::CollapseShapeOp::create(rewriter, loc, newOperandType,
193 operand, reassociation);
197 LogicalResult isUnpackOnInnerMostDim(RewriterBase &rewriter,
198 UnPackOp unpackOp)
const {
199 auto outerDimsPerm = unpackOp.getOuterDimsPerm();
201 return rewriter.notifyMatchFailure(
203 "expects outer_dims_perm is empty or an identity permutation");
206 ShapedType sourceType = unpackOp.getSourceType();
207 ShapedType destType = unpackOp.getDestType();
208 if (!sourceType.hasStaticShape() || !destType.hasStaticShape())
209 return rewriter.notifyMatchFailure(unpackOp,
"expects static shapes");
211 ArrayRef<int64_t> dimsPos = unpackOp.getInnerDimsPos();
212 if (dimsPos.size() != 1 || (dimsPos[0] + 1 != destType.getRank())) {
213 return rewriter.notifyMatchFailure(
214 unpackOp,
"expects unpacking on the innermost dimension");
220 LogicalResult matchAndRewrite(UnPackOp unpackOp,
221 PatternRewriter &rewriter)
const override {
223 if (!unpackOp.hasPureTensorSemantics())
226 ShapedType destType = unpackOp.getDestType();
227 if (
failed(isUnpackOnInnerMostDim(rewriter, unpackOp)) &&
228 failed(isPackOnEffectively1D(rewriter, &unpackOp)) &&
229 !unpackOp.isLikeUnPad()) {
233 ShapedType sourceType = unpackOp.getSourceType();
238 Value collapsed = insertCollapse(
239 rewriter, unpackOp.getLoc(), unpackOp.getSource(), destType,
241 rewriter.replaceOp(unpackOp, collapsed);
248struct FoldPadWithPackOp :
public OpRewritePattern<PackOp> {
251 : OpRewritePattern<PackOp>(context), controlFn(std::move(controlFn)) {}
253 LogicalResult matchAndRewrite(PackOp packOp,
254 PatternRewriter &rewriter)
const override {
255 auto padOp = packOp.getSource().getDefiningOp<tensor::PadOp>();
257 if (!padOp || padOp.getNofold() || !padOp.hasZeroLowPad())
261 if (controlFn && !controlFn(&packOp.getSourceMutable()))
264 Value constantPaddingValue = padOp.getConstantPaddingValue();
265 if (!constantPaddingValue)
268 if (
auto paddingValue = packOp.getPaddingValue())
277 ShapedType unpackedType = packOp.getSourceType();
278 SmallVector<int64_t> outerShapeWithoutTranspose =
280 for (
auto [pos, tileSize, high] :
281 llvm::zip_equal(packOp.getInnerDimsPos(), packOp.getStaticInnerTiles(),
282 padOp.getMixedHighPad())) {
283 if (unpackedType.isDynamicDim(pos))
285 if (ShapedType::isDynamic(outerShapeWithoutTranspose[pos]))
287 if (ShapedType::isDynamic(tileSize))
292 int64_t paddingSize = outerShapeWithoutTranspose[pos] * tileSize -
293 unpackedType.getDimSize(pos);
295 if (paddingSize + cstHigh.value() >= tileSize)
299 rewriter.replaceOpWithNewOp<PackOp>(
300 packOp, padOp.getSource(), packOp.getDest(), packOp.getInnerDimsPos(),
301 packOp.getMixedTiles(), constantPaddingValue,
302 packOp.getOuterDimsPerm());
312struct FoldUnpackWithExtractSliceOp
313 :
public OpRewritePattern<tensor::ExtractSliceOp> {
315 FoldUnpackWithExtractSliceOp(MLIRContext *context,
317 : OpRewritePattern<tensor::ExtractSliceOp>(context),
318 controlFn(std::move(controlFn)) {}
320 LogicalResult matchAndRewrite(tensor::ExtractSliceOp sliceOp,
321 PatternRewriter &rewriter)
const override {
322 auto unpackOp = sliceOp.getSource().getDefiningOp<UnPackOp>();
327 if (!unpackOp.hasPureTensorSemantics())
331 if (controlFn && !controlFn(&sliceOp.getSourceMutable()))
334 if (!unpackOp.canFoldSliceOp(sliceOp))
338 Type elementType = unpackOp.getDestType().getElementType();
339 Value output = tensor::EmptyOp::create(
340 rewriter, sliceOp.getLoc(), sliceOp.getMixedSizes(), elementType);
341 rewriter.replaceOpWithNewOp<UnPackOp>(
342 sliceOp, unpackOp.getSource(), output, unpackOp.getInnerDimsPos(),
343 unpackOp.getMixedTiles(), unpackOp.getOuterDimsPerm());
358static bool checkAndPermute(ArrayRef<int64_t> permutation,
359 ArrayRef<int64_t> inVec,
360 SmallVectorImpl<int64_t> &resVec, int64_t rank) {
362 for (
unsigned int i = 0; i < rank; ++i) {
363 int64_t remappedPosition = permutation[i];
364 if (remappedPosition >= rank)
367 remappedPosition = inVec[remappedPosition];
368 resVec.push_back(remappedPosition);
376struct FoldProducerPackWithConsumerLinalgTransposeOp
377 :
public OpInterfaceRewritePattern<linalg::LinalgOp> {
380 FoldProducerPackWithConsumerLinalgTransposeOp(
382 : OpInterfaceRewritePattern<linalg::LinalgOp>(context),
383 controlFn(std::move(controlFn)) {}
385 LogicalResult matchAndRewrite(linalg::LinalgOp linalgOp,
386 PatternRewriter &rewriter)
const override {
387 auto packOp = linalgOp->getOperand(0).getDefiningOp<PackOp>();
393 if (!packOp.hasPureTensorSemantics())
397 if (controlFn && !controlFn(&linalgOp->getOpOperand(0)))
400 FailureOr<SmallVector<int64_t>> maybePerm =
401 getTransposeOpPermutation(linalgOp);
405 auto innerDimsPos = packOp.getInnerDimsPos();
406 auto mixedInnerTiles = packOp.getMixedTiles();
407 auto outerDimsPerm = packOp.getOuterDimsPerm();
408 const auto &transposePerm = maybePerm.value();
409 SmallVector<int64_t> newOuterDimsPermVec;
410 SmallVector<int64_t> newInnerDimsPosVec;
411 SmallVector<OpFoldResult> newMixedInnerTilesVec;
412 int64_t srcRank = packOp.getSourceRank();
414 if (!checkAndPermute(transposePerm, outerDimsPerm, newOuterDimsPermVec,
416 return rewriter.notifyMatchFailure(
418 "Cannot fold in tensor.pack if a tile dimension was transposed "
419 "with a non-tile dimension in linalg.transpose.");
422 for (
unsigned int i = srcRank; i < transposePerm.size(); ++i) {
423 int64_t remappedPosition = transposePerm[i] - srcRank;
424 newMixedInnerTilesVec.push_back(mixedInnerTiles[remappedPosition]);
425 newInnerDimsPosVec.push_back(innerDimsPos[remappedPosition]);
428 Value output = packOp.createDestinationTensor(
429 rewriter, linalgOp.getLoc(), packOp.getSource(), newMixedInnerTilesVec,
430 newInnerDimsPosVec, newOuterDimsPermVec);
432 rewriter.replaceOpWithNewOp<PackOp>(
433 linalgOp, packOp.getSource(), output, newInnerDimsPosVec,
434 newMixedInnerTilesVec, packOp.getPaddingValue(), newOuterDimsPermVec);
445struct FoldConsumerPackWithProducerLinalgTransposeOp
446 :
public OpRewritePattern<PackOp> {
449 FoldConsumerPackWithProducerLinalgTransposeOp(
451 : OpRewritePattern<PackOp>(context), controlFn(std::move(controlFn)) {}
453 LogicalResult matchAndRewrite(PackOp packOp,
454 PatternRewriter &rewriter)
const override {
456 if (!packOp.hasPureTensorSemantics())
459 auto linalgOp = packOp.getSource().getDefiningOp<linalg::LinalgOp>();
464 if (controlFn && !controlFn(&packOp.getSourceMutable()))
467 FailureOr<SmallVector<int64_t>> maybePerm =
468 getTransposeOpPermutation(linalgOp);
472 auto transposePermutation = maybePerm.value();
473 auto outerDimsPerm = packOp.getOuterDimsPerm();
474 auto innerDimsPos = packOp.getInnerDimsPos();
475 SmallVector<int64_t> newInnerDimsPosVec;
476 SmallVector<int64_t> newOuterDimsPermVec =
477 llvm::to_vector(transposePermutation);
479 if (!outerDimsPerm.empty())
484 for (
auto dim : innerDimsPos)
485 newInnerDimsPosVec.push_back(transposePermutation[dim]);
487 Value output = packOp.createDestinationTensor(
488 rewriter, packOp.getLoc(), linalgOp->getOperand(0),
489 packOp.getMixedTiles(), newInnerDimsPosVec, newOuterDimsPermVec);
491 rewriter.replaceOpWithNewOp<PackOp>(
492 packOp, linalgOp->getOperand(0), output, newInnerDimsPosVec,
493 packOp.getMixedTiles(), packOp.getPaddingValue(), newOuterDimsPermVec);
504struct FoldProducerUnPackWithConsumerLinalgTransposeOp
505 :
public OpInterfaceRewritePattern<linalg::LinalgOp> {
508 FoldProducerUnPackWithConsumerLinalgTransposeOp(
510 : OpInterfaceRewritePattern<linalg::LinalgOp>(context),
511 controlFn(std::move(controlFn)) {}
513 LogicalResult matchAndRewrite(linalg::LinalgOp linalgOp,
514 PatternRewriter &rewriter)
const override {
515 auto unPackOp = linalgOp->getOperand(0).getDefiningOp<UnPackOp>();
521 if (!unPackOp.hasPureTensorSemantics())
525 if (controlFn && !controlFn(&linalgOp->getOpOperand(0)))
528 FailureOr<SmallVector<int64_t>> maybePerm =
529 getTransposeOpPermutation(linalgOp);
533 auto outerDimsPerm = unPackOp.getOuterDimsPerm();
534 auto innerDimsPos = unPackOp.getInnerDimsPos();
535 SmallVector<int64_t> newInnerDimsPosVec;
536 SmallVector<int64_t> newOuterDimsPermVec =
541 for (
auto dim : innerDimsPos)
542 newInnerDimsPosVec.push_back(newOuterDimsPermVec[dim]);
544 if (!outerDimsPerm.empty())
548 rewriter.replaceOpWithNewOp<UnPackOp>(
549 linalgOp, unPackOp.getSource(), linalgOp.getDpsInits()[0],
550 newInnerDimsPosVec, unPackOp.getMixedTiles(), newOuterDimsPermVec);
561struct FoldConsumerUnPackWithProducerLinalgTransposeOp
562 :
public OpRewritePattern<UnPackOp> {
563 using OpRewritePattern<UnPackOp>::OpRewritePattern;
566 FoldConsumerUnPackWithProducerLinalgTransposeOp(
568 : OpRewritePattern<UnPackOp>(context), controlFn(std::move(controlFn)) {}
570 LogicalResult matchAndRewrite(UnPackOp unPackOp,
571 PatternRewriter &rewriter)
const override {
573 if (!unPackOp.hasPureTensorSemantics())
576 auto linalgOp = unPackOp.getSource().getDefiningOp<linalg::LinalgOp>();
581 if (controlFn && !controlFn(&unPackOp.getSourceMutable()))
584 FailureOr<SmallVector<int64_t>> maybePerm =
585 getTransposeOpPermutation(linalgOp);
589 SmallVector<SmallVector<OpFoldResult>> unpackOpResultDims;
594 SmallVector<int64_t> inverseTransposePerm =
596 auto outerDimsPerm = unPackOp.getOuterDimsPerm();
597 auto innerDimsPos = unPackOp.getInnerDimsPos();
598 int64_t destRank = unPackOp.getSourceRank() - innerDimsPos.size();
599 auto mixedInnerTilesVec = unPackOp.getMixedTiles();
600 SmallVector<int64_t> newOuterDimsPermVec;
601 SmallVector<int64_t> newInnerDimsPosVec;
602 SmallVector<OpFoldResult> newMixedInnerTilesVec;
603 if (!checkAndPermute(inverseTransposePerm, outerDimsPerm,
604 newOuterDimsPermVec, destRank))
605 return rewriter.notifyMatchFailure(
607 "Cannot fold in tensor.unpack if a tile dimension was transposed "
608 "with a non-tile dimension in linalg.transpose.");
611 for (
unsigned int i = destRank; i < inverseTransposePerm.size(); ++i) {
612 int64_t remappedPosition = inverseTransposePerm[i] - destRank;
613 newMixedInnerTilesVec.push_back(mixedInnerTilesVec[remappedPosition]);
614 newInnerDimsPosVec.push_back(innerDimsPos[remappedPosition]);
618 cast<ShapedType>(unPackOp->getResultTypes()[0]).getElementType();
619 Value output = tensor::EmptyOp::create(rewriter, unPackOp->getLoc(),
620 unpackOpResultDims[0], elemType);
622 rewriter.replaceOpWithNewOp<UnPackOp>(
623 unPackOp, linalgOp->getOperand(0), output, newInnerDimsPosVec,
624 newMixedInnerTilesVec, newOuterDimsPermVec);
635struct FoldEmptyTensorWithPackOp :
public OpRewritePattern<PackOp> {
636 using OpRewritePattern<PackOp>::OpRewritePattern;
638 LogicalResult matchAndRewrite(PackOp packOp,
639 PatternRewriter &rewriter)
const override {
641 if (!packOp.hasPureTensorSemantics())
645 auto emptyOp = packOp.getSource().getDefiningOp<tensor::EmptyOp>();
651 if (packOp.getPaddingValue())
652 return rewriter.notifyMatchFailure(packOp,
"expects no padding value");
655 rewriter.replaceOp(packOp, packOp.getDest());
663struct FoldEmptyTensorWithUnPackOp :
public OpRewritePattern<UnPackOp> {
664 using OpRewritePattern<UnPackOp>::OpRewritePattern;
666 LogicalResult matchAndRewrite(UnPackOp unPackOp,
667 PatternRewriter &rewriter)
const override {
669 if (!unPackOp.hasPureTensorSemantics())
673 auto emptyOp = unPackOp.getSource().getDefiningOp<tensor::EmptyOp>();
678 rewriter.replaceOp(unPackOp, unPackOp.getDest());
688 patterns.
insert<FoldUnpackWithExtractSliceOp, FoldPadWithPackOp,
689 FoldProducerPackWithConsumerLinalgTransposeOp,
690 FoldConsumerPackWithProducerLinalgTransposeOp,
691 FoldConsumerUnPackWithProducerLinalgTransposeOp,
692 FoldProducerUnPackWithConsumerLinalgTransposeOp>(
697 patterns.
add<SimplifyPackToExpandShape, SimplifyUnPackToCollapseShape>(
703 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.