MLIR 24.0.0git
PackAndUnpackPatterns.cpp
Go to the documentation of this file.
1//===- FoldIntoPackAndUnpackPatterns.cpp ----------------------------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8
16
17namespace mlir {
18namespace linalg {
19namespace {
20
21/// Returns the number of shape sizes that is either dynamic or greater than 1.
22static int64_t getNumGtOneDims(ArrayRef<int64_t> shape) {
23 return llvm::count_if(
24 shape, [](int64_t v) { return ShapedType::isDynamic(v) || v > 1; });
25}
26
27/// Returns the index of the first non-unit size in `sizes`. Returns -1 if
28/// there are no non-unit sizes.
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;
32}
33
34/// Check whether `op` is effectively a 1D pack/unpack. Example:
35///
36/// %pack = linalg.pack %src
37/// inner_dims_pos = [0, 1]
38/// inner_tiles = [1, 2] into %dest
39/// : tensor<1x32xf32> -> tensor<1x16x1x2xf32>
40///
41/// Returns success() if there is:
42/// * only 1 non-unit dim in the un-packed domain,
43/// * only 1 non-unit inner tile size, and
44/// * the unique non-unit tile size is applied to the unique non-unit
45/// un-packed dim.
46template <typename PackOrUnpackOp>
47static LogicalResult isPackOnEffectively1D(RewriterBase &rewriter,
48 PackOrUnpackOp *op) {
49 // Obtain the unpacked shape.
50 auto pack = dyn_cast<linalg::PackOp>(op);
51 auto unpack = dyn_cast<linalg::UnPackOp>(op);
52
53 ArrayRef<int64_t> unpackedShape = pack ? pack->getSourceType().getShape()
54 : unpack->getDestType().getShape();
55
56 // Obtain the inner tile sizes.
57 ArrayRef<int64_t> innerTileSizes = op->getStaticInnerTiles();
58
59 // Make sure that there is exactly single non-unit unpacked dim.
60 if (getNumGtOneDims(unpackedShape) != 1) {
61 return rewriter.notifyMatchFailure(
62 *op, "expects non-packed domain to have at most one non-unit dims");
63 }
64
65 // Make sure that there is at most one non-unit inner tile size.
66 auto numNonUnitInnerTiles = getNumGtOneDims(innerTileSizes);
67 if (numNonUnitInnerTiles > 1) {
68 return rewriter.notifyMatchFailure(
69 *op, "expects at most one non-unit inner tiles");
70 }
71
72 // If there are no non-unit tiles, there is nothing else to check.
73 if (numNonUnitInnerTiles == 0)
74 return success();
75
76 // Get the index of the unique non-unit unpacked dim.
77 int64_t nonUnitDimIdx = getFirstNonUnitSizeIdx(unpackedShape);
78
79 // Get the index of the dim that the unique non-unit tile is applied to.
80 int64_t nonUnitTileDestDimIdx = getFirstNonUnitSizeIdx(innerTileSizes);
81
82 // Make sure that the unique non-unit tile is applied to the unique unit dim.
83 if (nonUnitTileDestDimIdx != nonUnitDimIdx) {
84 return rewriter.notifyMatchFailure(
85 *op, "expects at most one non-unit inner tiles");
86 }
87
88 return success();
89}
90
91// If the `linalgOp` represents a transpose, return the permutation vector for
92// the transpose. Otherwise, return failure.
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())
98 return failure();
99
100 if (linalgOp.getNumDpsInputs() != 1 || linalgOp.getNumDpsInits() != 1)
101 return failure();
102 auto mapRange = linalgOp.getIndexingMapsArray();
103 if (!mapRange.front().isPermutation() || !mapRange.back().isPermutation() ||
104 mapRange.front() == mapRange.back()) {
105 return failure();
106 }
107 if (!llvm::hasSingleElement(linalgOp.getBlock()->getOperations()))
108 return failure();
109 AffineMap outMap = mapRange.back();
110 AffineMap inMap = mapRange.front();
111 // To get the permutation, look at each output index and find which
112 // dimension in the input we're reading from for that index.
113 return llvm::map_to_vector(outMap.getResults(),
114 [&](AffineExpr expr) -> int64_t {
115 return *inMap.getResultPosition(expr);
116 });
117}
118
119/// Packing one-dimensional tensor can be expressed as an expand shape op.
120struct SimplifyPackToExpandShape : public OpRewritePattern<PackOp> {
121 using OpRewritePattern<PackOp>::OpRewritePattern;
122
123 FailureOr<Value>
124 insertExpand(RewriterBase &rewriter, Location loc, Value operand,
125 Type newOperandType,
126 ArrayRef<ReassociationIndices> reassociation) const {
127 if (operand.getType() == newOperandType)
128 return operand;
129 return tensor::ExpandShapeOp::create(rewriter, loc, newOperandType, operand,
130 reassociation)
131 .getResult();
132 }
133
134 /// Returns success() if it is only packing on the innermost dimension.
135 LogicalResult isPackOnInnerMostDim(RewriterBase &rewriter,
136 PackOp packOp) const {
137 auto outerDimsPerm = packOp.getOuterDimsPerm();
138 if (!outerDimsPerm.empty() && !isIdentityPermutation(outerDimsPerm)) {
139 return rewriter.notifyMatchFailure(
140 packOp,
141 "expects outer_dims_perm is empty or an identity permutation");
142 }
143
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");
149 }
150 return success();
151 }
152
153 LogicalResult matchAndRewrite(PackOp packOp,
154 PatternRewriter &rewriter) const override {
155 if (packOp.getPaddingValue())
156 return rewriter.notifyMatchFailure(packOp, "expects no padding value");
157 // Pack/unpack memref transformations are unsupported. The memref forms
158 // are mainly for bufferization and scalar lowering. Other uses are not
159 // recommended, see #225650 for details.
160 if (!packOp.hasPureTensorSemantics())
161 return failure();
162
163 ShapedType sourceType = packOp.getSourceType();
164 if (failed(isPackOnInnerMostDim(rewriter, packOp)) &&
165 failed(isPackOnEffectively1D(rewriter, &packOp)) &&
166 !packOp.isLikePad()) {
167 return failure();
168 }
169
170 ShapedType destType = packOp.getDestType();
171 auto reassociation =
172 getReassociationIndicesForReshape(sourceType, destType);
173 if (!reassociation)
174 return failure();
175 FailureOr<Value> expanded =
176 insertExpand(rewriter, packOp.getLoc(), packOp.getSource(), destType,
177 *reassociation);
178 if (failed(expanded)) {
179 return rewriter.notifyMatchFailure(
180 packOp, "unable to expand source of tensor.pack");
181 }
182 rewriter.replaceOp(packOp, *expanded);
183 return success();
184 }
185};
186
187struct SimplifyUnPackToCollapseShape : public OpRewritePattern<UnPackOp> {
188 using OpRewritePattern<UnPackOp>::OpRewritePattern;
189
190 Value insertCollapse(RewriterBase &rewriter, Location loc, Value operand,
191 Type newOperandType, ArrayAttr reassociation) const {
192 if (operand.getType() == newOperandType)
193 return operand;
194 return tensor::CollapseShapeOp::create(rewriter, loc, newOperandType,
195 operand, reassociation);
196 }
197
198 /// Returns success() if it is unpacking on the innermost dimension.
199 LogicalResult isUnpackOnInnerMostDim(RewriterBase &rewriter,
200 UnPackOp unpackOp) const {
201 auto outerDimsPerm = unpackOp.getOuterDimsPerm();
202 if (!outerDimsPerm.empty() && !isIdentityPermutation(outerDimsPerm)) {
203 return rewriter.notifyMatchFailure(
204 unpackOp,
205 "expects outer_dims_perm is empty or an identity permutation");
206 }
207
208 ShapedType sourceType = unpackOp.getSourceType();
209 ShapedType destType = unpackOp.getDestType();
210 if (!sourceType.hasStaticShape() || !destType.hasStaticShape())
211 return rewriter.notifyMatchFailure(unpackOp, "expects static shapes");
212
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");
217 }
218
219 return success();
220 }
221
222 LogicalResult matchAndRewrite(UnPackOp unpackOp,
223 PatternRewriter &rewriter) const override {
224 // Pack/unpack memref transformations are unsupported. The memref forms
225 // are mainly for bufferization and scalar lowering. Other uses are not
226 // recommended, see #225650 for details.
227 if (!unpackOp.hasPureTensorSemantics())
228 return failure();
229
230 ShapedType destType = unpackOp.getDestType();
231 if (failed(isUnpackOnInnerMostDim(rewriter, unpackOp)) &&
232 failed(isPackOnEffectively1D(rewriter, &unpackOp)) &&
233 !unpackOp.isLikeUnPad()) {
234 return failure();
235 }
236
237 ShapedType sourceType = unpackOp.getSourceType();
238 auto reassociation =
239 getReassociationIndicesForReshape(sourceType, destType);
240 if (!reassociation)
241 return failure();
242 Value collapsed = insertCollapse(
243 rewriter, unpackOp.getLoc(), unpackOp.getSource(), destType,
244 getReassociationIndicesAttribute(rewriter, *reassociation));
245 rewriter.replaceOp(unpackOp, collapsed);
246 return success();
247 }
248};
249
250/// Fold a `pad` -> `pack` into `pack` if they have the same padding values and
251/// the pad op has zero low paddings, or if `pack` has no padding values.
252struct FoldPadWithPackOp : public OpRewritePattern<PackOp> {
253public:
254 FoldPadWithPackOp(MLIRContext *context, ControlFoldIntoPackUnpackFn controlFn)
255 : OpRewritePattern<PackOp>(context), controlFn(std::move(controlFn)) {}
256
257 LogicalResult matchAndRewrite(PackOp packOp,
258 PatternRewriter &rewriter) const override {
259 auto padOp = packOp.getSource().getDefiningOp<tensor::PadOp>();
260
261 if (!padOp || padOp.getNofold() || !padOp.hasZeroLowPad())
262 return failure();
263
264 // User controlled folding function.
265 if (controlFn && !controlFn(&packOp.getSourceMutable()))
266 return failure();
267
268 Value constantPaddingValue = padOp.getConstantPaddingValue();
269 if (!constantPaddingValue)
270 return failure();
271
272 if (auto paddingValue = packOp.getPaddingValue())
273 if (!isEqualConstantIntOrValue(paddingValue, constantPaddingValue))
274 return failure();
275
276 // Folding is not allowed if it were to introduce artificial padding.
277 // Folding is also disabled in the case of dynamic dimensions and/or tile
278 // sizes - that is because it would be impossible to compute the padding
279 // size and hence to establish whether "artificial" padding would be
280 // created.
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))
288 return failure();
289 if (ShapedType::isDynamic(outerShapeWithoutTranspose[pos]))
290 return failure();
291 if (ShapedType::isDynamic(tileSize))
292 return failure();
293 std::optional<int64_t> cstHigh = getConstantIntValue(high);
294 if (!cstHigh)
295 return failure();
296 int64_t paddingSize = outerShapeWithoutTranspose[pos] * tileSize -
297 unpackedType.getDimSize(pos);
298 // Do not fold the op if it requires artificial padding.
299 if (paddingSize + cstHigh.value() >= tileSize)
300 return failure();
301 }
302
303 rewriter.replaceOpWithNewOp<PackOp>(
304 packOp, padOp.getSource(), packOp.getDest(), packOp.getInnerDimsPos(),
305 packOp.getMixedTiles(), constantPaddingValue,
306 packOp.getOuterDimsPerm());
307 return success();
308 }
309
310private:
312};
313
314/// Fold a `unpack` -> `extract_slice` into the `unpack` since it already
315/// has extract_slice semantics.
316struct FoldUnpackWithExtractSliceOp
317 : public OpRewritePattern<tensor::ExtractSliceOp> {
318public:
319 FoldUnpackWithExtractSliceOp(MLIRContext *context,
321 : OpRewritePattern<tensor::ExtractSliceOp>(context),
322 controlFn(std::move(controlFn)) {}
323
324 LogicalResult matchAndRewrite(tensor::ExtractSliceOp sliceOp,
325 PatternRewriter &rewriter) const override {
326 auto unpackOp = sliceOp.getSource().getDefiningOp<UnPackOp>();
327 if (!unpackOp)
328 return failure();
329
330 // Pack/unpack memref transformations are unsupported. The memref forms
331 // are mainly for bufferization and scalar lowering. Other uses are not
332 // recommended, see #225650 for details.
333 if (!unpackOp.hasPureTensorSemantics())
334 return failure();
335
336 // User controlled folding function.
337 if (controlFn && !controlFn(&sliceOp.getSourceMutable()))
338 return failure();
339
340 if (!unpackOp.canFoldSliceOp(sliceOp))
341 return failure();
342
343 // Create a new empty output tensor.
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());
350 return success();
351 }
352
353private:
355};
356
357// Applies 'permutation' on 'inVec' and stores the result in resVec.
358// 'inVec' may be empty, in that case it's one-to-one mapping with permutation.
359// `rank` sets the boundary for permutation i.e., the permutation dim can't be
360// greater than the rank specified. If it's so then return false.
361// For e.g., permutation {1, 0, 3, 2} with rank 2 is allowed since the values in
362// permutation[:rank] doesn't exceed rank, whereas, permutation {1, 3, 0, 2} is
363// not allowed since `3` exceeds the value of the rank in the given range.
364static bool checkAndPermute(ArrayRef<int64_t> permutation,
365 ArrayRef<int64_t> inVec,
366 SmallVectorImpl<int64_t> &resVec, int64_t rank) {
367
368 for (unsigned int i = 0; i < rank; ++i) {
369 int64_t remappedPosition = permutation[i];
370 if (remappedPosition >= rank)
371 return false;
372 if (!inVec.empty())
373 remappedPosition = inVec[remappedPosition];
374 resVec.push_back(remappedPosition);
375 }
376
377 return true;
378}
379
380/// Fold 'pack' -> 'transpose' into 'pack' since 'pack' already has transpose
381/// semantics.
382struct FoldProducerPackWithConsumerLinalgTransposeOp
383 : public OpInterfaceRewritePattern<linalg::LinalgOp> {
384
385public:
386 FoldProducerPackWithConsumerLinalgTransposeOp(
387 MLIRContext *context, ControlFoldIntoPackUnpackFn controlFn)
388 : OpInterfaceRewritePattern<linalg::LinalgOp>(context),
389 controlFn(std::move(controlFn)) {}
390
391 LogicalResult matchAndRewrite(linalg::LinalgOp linalgOp,
392 PatternRewriter &rewriter) const override {
393 auto packOp = linalgOp->getOperand(0).getDefiningOp<PackOp>();
394
395 if (!packOp)
396 return failure();
397
398 // Pack/unpack memref transformations are unsupported. The memref forms
399 // are mainly for bufferization and scalar lowering. Other uses are not
400 // recommended, see #225650 for details.
401 if (!packOp.hasPureTensorSemantics())
402 return failure();
403
404 // User controlled folding function.
405 if (controlFn && !controlFn(&linalgOp->getOpOperand(0)))
406 return failure();
407
408 FailureOr<SmallVector<int64_t>> maybePerm =
409 getTransposeOpPermutation(linalgOp);
410 if (failed(maybePerm))
411 return failure();
412
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();
421
422 if (!checkAndPermute(transposePerm, outerDimsPerm, newOuterDimsPermVec,
423 srcRank))
424 return rewriter.notifyMatchFailure(
425 linalgOp,
426 "Cannot fold in tensor.pack if a tile dimension was transposed "
427 "with a non-tile dimension in linalg.transpose.");
428
429 // Process transpose operation for tiled inner dimensions
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]);
434 }
435
436 Value output = packOp.createDestinationTensor(
437 rewriter, linalgOp.getLoc(), packOp.getSource(), newMixedInnerTilesVec,
438 newInnerDimsPosVec, newOuterDimsPermVec);
439
440 rewriter.replaceOpWithNewOp<PackOp>(
441 linalgOp, packOp.getSource(), output, newInnerDimsPosVec,
442 newMixedInnerTilesVec, packOp.getPaddingValue(), newOuterDimsPermVec);
443
444 return success();
445 }
446
447private:
449};
450
451/// Fold 'transpose' -> 'pack' into 'pack' since 'pack' already has transpose
452/// semantics.
453struct FoldConsumerPackWithProducerLinalgTransposeOp
454 : public OpRewritePattern<PackOp> {
455
456public:
457 FoldConsumerPackWithProducerLinalgTransposeOp(
458 MLIRContext *context, ControlFoldIntoPackUnpackFn controlFn)
459 : OpRewritePattern<PackOp>(context), controlFn(std::move(controlFn)) {}
460
461 LogicalResult matchAndRewrite(PackOp packOp,
462 PatternRewriter &rewriter) const override {
463 // Pack/unpack memref transformations are unsupported. The memref forms
464 // are mainly for bufferization and scalar lowering. Other uses are not
465 // recommended, see #225650 for details.
466 if (!packOp.hasPureTensorSemantics())
467 return failure();
468
469 auto linalgOp = packOp.getSource().getDefiningOp<linalg::LinalgOp>();
470 if (!linalgOp)
471 return failure();
472
473 // User controlled folding function.
474 if (controlFn && !controlFn(&packOp.getSourceMutable()))
475 return failure();
476
477 FailureOr<SmallVector<int64_t>> maybePerm =
478 getTransposeOpPermutation(linalgOp);
479 if (failed(maybePerm))
480 return failure();
481
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);
488
489 if (!outerDimsPerm.empty())
490 applyPermutationToVector(newOuterDimsPermVec, outerDimsPerm);
491
492 // Can't use applyPermutationToVector for newInnerDimsPosVec since input and
493 // permutation rank won't necessarily be equal in all cases.
494 for (auto dim : innerDimsPos)
495 newInnerDimsPosVec.push_back(transposePermutation[dim]);
496
497 Value output = packOp.createDestinationTensor(
498 rewriter, packOp.getLoc(), linalgOp->getOperand(0),
499 packOp.getMixedTiles(), newInnerDimsPosVec, newOuterDimsPermVec);
500
501 rewriter.replaceOpWithNewOp<PackOp>(
502 packOp, linalgOp->getOperand(0), output, newInnerDimsPosVec,
503 packOp.getMixedTiles(), packOp.getPaddingValue(), newOuterDimsPermVec);
504
505 return success();
506 }
507
508private:
510};
511
512/// Fold 'unpack' -> 'transpose' into 'unpack' since 'unpack' already has
513/// transpose semantics.
514struct FoldProducerUnPackWithConsumerLinalgTransposeOp
515 : public OpInterfaceRewritePattern<linalg::LinalgOp> {
516
517public:
518 FoldProducerUnPackWithConsumerLinalgTransposeOp(
519 MLIRContext *context, ControlFoldIntoPackUnpackFn controlFn)
520 : OpInterfaceRewritePattern<linalg::LinalgOp>(context),
521 controlFn(std::move(controlFn)) {}
522
523 LogicalResult matchAndRewrite(linalg::LinalgOp linalgOp,
524 PatternRewriter &rewriter) const override {
525 auto unPackOp = linalgOp->getOperand(0).getDefiningOp<UnPackOp>();
526
527 if (!unPackOp)
528 return failure();
529
530 // Pack/unpack memref transformations are unsupported. The memref forms
531 // are mainly for bufferization and scalar lowering. Other uses are not
532 // recommended, see #225650 for details.
533 if (!unPackOp.hasPureTensorSemantics())
534 return failure();
535
536 // User controlled folding function.
537 if (controlFn && !controlFn(&linalgOp->getOpOperand(0)))
538 return failure();
539
540 FailureOr<SmallVector<int64_t>> maybePerm =
541 getTransposeOpPermutation(linalgOp);
542 if (failed(maybePerm))
543 return failure();
544
545 auto outerDimsPerm = unPackOp.getOuterDimsPerm();
546 auto innerDimsPos = unPackOp.getInnerDimsPos();
547 SmallVector<int64_t> newInnerDimsPosVec;
548 SmallVector<int64_t> newOuterDimsPermVec =
549 invertPermutationVector(maybePerm.value());
550
551 // Can't use applyPermutationToVector for newInnerDimsPosVec since input and
552 // permutation rank won't necessarily be equal in all cases.
553 for (auto dim : innerDimsPos)
554 newInnerDimsPosVec.push_back(newOuterDimsPermVec[dim]);
555
556 if (!outerDimsPerm.empty())
557 applyPermutationToVector(newOuterDimsPermVec, outerDimsPerm);
558
559 // Reuse the destination of the transpose op.
560 rewriter.replaceOpWithNewOp<UnPackOp>(
561 linalgOp, unPackOp.getSource(), linalgOp.getDpsInits()[0],
562 newInnerDimsPosVec, unPackOp.getMixedTiles(), newOuterDimsPermVec);
563
564 return success();
565 }
566
567private:
569};
570
571/// Fold 'transpose' -> 'unpack' into 'unpack' since 'unpack' already has
572/// transpose semantics.
573struct FoldConsumerUnPackWithProducerLinalgTransposeOp
574 : public OpRewritePattern<UnPackOp> {
575 using OpRewritePattern<UnPackOp>::OpRewritePattern;
576
577public:
578 FoldConsumerUnPackWithProducerLinalgTransposeOp(
579 MLIRContext *context, ControlFoldIntoPackUnpackFn controlFn)
580 : OpRewritePattern<UnPackOp>(context), controlFn(std::move(controlFn)) {}
581
582 LogicalResult matchAndRewrite(UnPackOp unPackOp,
583 PatternRewriter &rewriter) const override {
584 // Pack/unpack memref transformations are unsupported. The memref forms
585 // are mainly for bufferization and scalar lowering. Other uses are not
586 // recommended, see #225650 for details.
587 if (!unPackOp.hasPureTensorSemantics())
588 return failure();
589
590 auto linalgOp = unPackOp.getSource().getDefiningOp<linalg::LinalgOp>();
591 if (!linalgOp)
592 return failure();
593
594 // User controlled folding function.
595 if (controlFn && !controlFn(&unPackOp.getSourceMutable()))
596 return failure();
597
598 FailureOr<SmallVector<int64_t>> maybePerm =
599 getTransposeOpPermutation(linalgOp);
600 if (failed(maybePerm))
601 return failure();
602
603 SmallVector<SmallVector<OpFoldResult>> unpackOpResultDims;
604 if (failed(reifyResultShapes(rewriter, unPackOp, unpackOpResultDims))) {
605 return failure();
606 }
607
608 SmallVector<int64_t> inverseTransposePerm =
609 invertPermutationVector(maybePerm.value());
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(
620 unPackOp,
621 "Cannot fold in tensor.unpack if a tile dimension was transposed "
622 "with a non-tile dimension in linalg.transpose.");
623
624 // Process transpose operation for tiled inner dimensions
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]);
629 }
630
631 auto elemType =
632 cast<ShapedType>(unPackOp->getResultTypes()[0]).getElementType();
633 Value output = tensor::EmptyOp::create(rewriter, unPackOp->getLoc(),
634 unpackOpResultDims[0], elemType);
635
636 rewriter.replaceOpWithNewOp<UnPackOp>(
637 unPackOp, linalgOp->getOperand(0), output, newInnerDimsPosVec,
638 newMixedInnerTilesVec, newOuterDimsPermVec);
639
640 return success();
641 }
642
643private:
645};
646
647/// tensor.empty does not define any tensor contents, so an unpadded pack
648/// can be folded away.
649struct FoldEmptyTensorWithPackOp : public OpRewritePattern<PackOp> {
650 using OpRewritePattern<PackOp>::OpRewritePattern;
651
652 LogicalResult matchAndRewrite(PackOp packOp,
653 PatternRewriter &rewriter) const override {
654 // Pack/unpack memref transformations are unsupported. The memref forms
655 // are mainly for bufferization and scalar lowering. Other uses are not
656 // recommended, see #225650 for details.
657 if (!packOp.hasPureTensorSemantics())
658 return failure();
659
660 // Check for tensor.empty source.
661 auto emptyOp = packOp.getSource().getDefiningOp<tensor::EmptyOp>();
662 if (!emptyOp)
663 return failure();
664
665 // Check for padding.
666 // Packing with padding cannot be simply removed.
667 if (packOp.getPaddingValue())
668 return rewriter.notifyMatchFailure(packOp, "expects no padding value");
669
670 // Replace the pack directly with its destination.
671 rewriter.replaceOp(packOp, packOp.getDest());
672
673 return success();
674 }
675};
676
677/// tensor.empty does not define any tensor contents, so an unpack
678/// can be folded away.
679struct FoldEmptyTensorWithUnPackOp : public OpRewritePattern<UnPackOp> {
680 using OpRewritePattern<UnPackOp>::OpRewritePattern;
681
682 LogicalResult matchAndRewrite(UnPackOp unPackOp,
683 PatternRewriter &rewriter) const override {
684 // Pack/unpack memref transformations are unsupported. The memref forms
685 // are mainly for bufferization and scalar lowering. Other uses are not
686 // recommended, see #225650 for details.
687 if (!unPackOp.hasPureTensorSemantics())
688 return failure();
689
690 // Check for tensor.empty source.
691 auto emptyOp = unPackOp.getSource().getDefiningOp<tensor::EmptyOp>();
692 if (!emptyOp)
693 return failure();
694
695 // Replace the unpack directly with its destination.
696 rewriter.replaceOp(unPackOp, unPackOp.getDest());
697
698 return success();
699 }
700};
701
702} // namespace
703
705 RewritePatternSet &patterns, const ControlFoldIntoPackUnpackFn &controlFn) {
706 patterns.insert<FoldUnpackWithExtractSliceOp, FoldPadWithPackOp,
707 FoldProducerPackWithConsumerLinalgTransposeOp,
708 FoldConsumerPackWithProducerLinalgTransposeOp,
709 FoldConsumerUnPackWithProducerLinalgTransposeOp,
710 FoldProducerUnPackWithConsumerLinalgTransposeOp>(
711 patterns.getContext(), controlFn);
712}
713
715 patterns.add<SimplifyPackToExpandShape, SimplifyUnPackToCollapseShape>(
716 patterns.getContext());
717}
718
720 RewritePatternSet &patterns) {
721 patterns.add<FoldEmptyTensorWithPackOp, FoldEmptyTensorWithUnPackOp>(
722 patterns.getContext());
723}
724
725} // namespace linalg
726} // namespace mlir
return success()
ArrayAttr()
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.
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
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.