MLIR 24.0.0git
Transforms.cpp
Go to the documentation of this file.
1//===- Transforms.cpp - Linalg transformations as patterns ----------------===//
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//
9// This file implements logic and helpers to expose Linalg transforms as rewrite
10// patterns.
11//
12//===----------------------------------------------------------------------===//
13
28#include "mlir/IR/AffineExpr.h"
31#include "mlir/Support/LLVM.h"
32#include "llvm/ADT/SmallVectorExtras.h"
33#include "llvm/ADT/TypeSwitch.h"
34#include "llvm/Support/Debug.h"
35#include "llvm/Support/DebugLog.h"
36#include "llvm/Support/InterleavedRange.h"
37#include "llvm/Support/raw_ostream.h"
38#include <type_traits>
39#include <utility>
40
41#define DEBUG_TYPE "linalg-transforms"
42
43using namespace mlir;
44using namespace mlir::linalg;
45
46//===----------------------------------------------------------------------===//
47// Transformations exposed as functional-style API calls.
48//===----------------------------------------------------------------------===//
49
50//===----------------------------------------------------------------------===//
51// peelLoop transformation.
52//===----------------------------------------------------------------------===//
53
54/// Try to peel and canonicalize loop `op` and return the new result.
55/// Also applies affine_min/max bounds simplification on the fly where relevant.
56// TODO: Add support for scf.parallel and affine.for loops.
58 Operation *op) {
60 .Case([&](scf::ForOp forOp) {
61 scf::ForOp partialIteration;
62 if (succeeded(scf::peelForLoopAndSimplifyBounds(rewriter, forOp,
63 partialIteration)))
64 return partialIteration->getResults();
65 assert(!partialIteration && "expected that loop was not peeled");
66 return forOp->getResults();
67 })
68 .Default([&](Operation *op) { return op->getResults(); });
69}
70
71/// Peel 'loops' and applies affine_min/max bounds simplification on the fly
72/// where relevant.
75 for (auto loopOp : loops)
76 peelLoop(rewriter, loopOp);
77}
78
79//===----------------------------------------------------------------------===//
80// pack transformation.
81//===----------------------------------------------------------------------===//
82
83#ifndef NDEBUG
84/// Return true if `map` has 0 or 1 result function of AffineDimExpr(dim).
86 bool found = false;
87 for (AffineExpr e : map.getResults()) {
88 if (!e.isFunctionOfDim(dim))
89 continue;
90 if (found)
91 return false;
92 found = true;
93 }
94 return true;
95}
96#endif // NDEBUG
97
99 return llvm::interleaved(ri, ", ", /*Prefix=*/"|", /*Suffix=*/"");
100}
101
102/// Return the index of the first result of `map` that is a function of
103/// AffineDimExpr(dim), std::nullopt otherwise.
104static std::optional<int64_t> getFirstResultIndexFunctionOf(AffineMap map,
105 int64_t dim) {
106 for (int64_t i = 0, e = map.getNumResults(); i < e; ++i) {
107 AffineExpr expr = map.getResult(i);
108 if (!expr.isFunctionOfDim(dim))
109 continue;
110 return i;
111 }
112 return std::nullopt;
113}
114
115/// Perform one step of packing of a LinalgOp's metadata along `dim` into the
116/// `newDim` at `iteratorTypes.size()` by:
117/// 1. Appending `iteratorTypes[newDim]`, equal to `iteratorTypes[dim]`.
118/// 2. Appending a `newDim` to the domain of every indexing map.
119/// 3. For each operand (i.e. for each map in `indexingMaps`), perform packing
120/// by potentially adding a `newDim` result to `map`.
121/// The preserved invariant is that `iteratorTypes.size()` is always equal to
122/// `map.getNumDims()` for every map in `indexingMaps`.
123///
124/// Update `indexingMaps` and `iteratorTypes` inplace as one step of the update.
125/// Return a vector that records the optional packing for each operand.
126/// Return failure if the packed indexing cannot be represented with a LinalgOp.
127///
128/// Further details:
129/// ================
130/// The current implementation of packing (i.e. data tiling) consists of
131/// rewriting a linearized strip-mined form into a higher-dimensional access.
132/// e.g. consider an access `A[I][f(j, k, l)]` and packing by 4; we rewrite
133/// `I` into `4 * i + ii`, where `0 <= ii < 4`.
134/// The access is further rewritten as `A[i][f(j, k, l)][ii]`.
135///
136/// This rewrite into higher dimensional access is not possible for general
137/// AffineExpr in Linalg atm, it is restricted to an AffineDimExpr:
138/// e.g. consider an access `A[I + J][f(j, k, l)]` and packing by 4; we
139/// rewrite `I + J` into `4 * i + ii + J`, where `0 <= ii < 4`.
140/// The rewrite of the access would be a form not representable in Linalg:
141/// `A[i + (ii + J) / 4][f(j, k, l)][(ii + J) % 4]`.
142/// Note however that as `J` and `ii` iterate, the accesses do not have a
143/// particular alignment, so packing does not achieve alignment in this case
144///
145/// In the future, we may want to consider a mixed-form that allows some
146/// alignment in the presence of multiple accesses:
147/// `A[I][f(j, k, l)]` and `B[I + J][f(j, k, l)]`
148/// And would rewrite accesses as:
149/// `A[i][f(j, k, l)][ii]` and `B[4 * i + ii + J][f(j, k, l)]`
150static FailureOr<SmallVector<std::optional<int64_t>>>
153 int64_t dim) {
154 int64_t newDim = iteratorTypes.size();
155 iteratorTypes.push_back(iteratorTypes[dim]);
156
157 SmallVector<std::optional<int64_t>> packedDimPerIndexingMap(
158 indexingMaps.size(), std::nullopt);
160 for (int64_t operandIdx = 0, e = indexingMaps.size(); operandIdx < e;
161 ++operandIdx) {
162 AffineMap map = indexingMaps[operandIdx];
163
164 // Add the `newDim` to map whatever the case.
165 assert(map.getNumDims() == newDim && "num dims invariant violation");
166 map = map.shiftDims(1, newDim);
167
168 // Get the at-most-1 index of the result that is a function of `dim`.
169 // If we can find one, we insert `AffineDimExpr(newDim)` to the map, which
170 // logically chunks dimension `dim` into `K * dim + newDim`, where the
171 // packing factor `K` is specified separately.
172 assert(hasAtMostOneResultFunctionOfDim(map, dim) &&
173 "num results invariant violation");
174 auto maybeOperandDimensionToPack = getFirstResultIndexFunctionOf(map, dim);
175 if (!maybeOperandDimensionToPack.has_value()) {
176 newMaps.push_back(map);
177 continue;
178 }
179
180 // We can only pack AffineDimExpr atm.
181 if (!isa<AffineDimExpr>(map.getResult(maybeOperandDimensionToPack.value())))
182 return failure();
183
184 // Add `newDim` to the results of the map.
185 map = map.insertResult(Builder(map.getContext()).getAffineDimExpr(newDim),
186 map.getNumResults());
187 newMaps.push_back(map);
188
189 // Record the that `operandIdx` is packed.
190 packedDimPerIndexingMap[operandIdx] = maybeOperandDimensionToPack;
191 }
192 indexingMaps = newMaps;
193
194 return packedDimPerIndexingMap;
195}
196
197namespace {
198
199/// Helper struct to encode packing along one dimension of a LinalgOp.
200struct PackedOperandsDim {
201 OpFoldResult packedSize;
202 SmallVector<std::optional<int64_t>> packedDimForEachOperand;
203};
204
205/// Helper struct to encode packing along all dimensions of a LinalgOp.
206struct PackedOperandsDimList {
207 void pushBack(PackedOperandsDim &&packedOperandsDims) {
208 spec.emplace_back(packedOperandsDims);
209 }
210 /// Return all the dims that have been packed for operand @ `operandPos`.
211 SmallVector<int64_t> extractPackedDimsForOperand(int64_t operandPos);
212 /// Return all the pack sizes by which an operand @ `operandPos` is packed.
213 SmallVector<OpFoldResult> extractPackSizesForOperand(int64_t operandPos);
214
215private:
216 SmallVector<PackedOperandsDim> spec;
217};
218
219} // namespace
220
221FailureOr<LowerPackResult> linalg::lowerPack(RewriterBase &rewriter,
222 linalg::PackOp packOp,
223 bool lowerPadLikeWithInsertSlice) {
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 (!packOp.hasPureTensorSemantics())
228 return failure();
229
230 auto packedTensorType =
231 cast<RankedTensorType>(packOp->getResultTypes().front());
232
233 Location loc = packOp->getLoc();
234 OpBuilder::InsertionGuard g(rewriter);
235 rewriter.setInsertionPoint(packOp);
236
237 // 2. Compute the permutation vector to shuffle packed shape into the shape
238 // before any outer or inner permutations have been applied.
239 PackingMetadata packingMetadata;
240 SmallVector<int64_t> packedToStripMinedShapePerm =
241 getPackInverseDestPerm(packOp, packingMetadata);
242
243 // 3. Compute the stripMinedShape: this is the packed shape before any outer
244 // or inner permutations have been applied.
245 SmallVector<int64_t> stripMinedShape(packedTensorType.getShape());
246 applyPermutationToVector(stripMinedShape, packedToStripMinedShapePerm);
247
248 // Also compute the mixed (static+dynamic) strip-mined sizes for the
249 // expand_shape output. This is needed to support dynamic inner tile sizes,
250 // since the shapes cannot be inferred automatically when multiple dynamic
251 // dims appear in a single reassociation group during ExpandShapeOp
252 // construction.
253 SmallVector<OpFoldResult> stripMinedMixedSizes =
254 tensor::getMixedSizes(rewriter, loc, packOp.getDest());
255 applyPermutationToVector(stripMinedMixedSizes, packedToStripMinedShapePerm);
256
257 // 4. Pad the source of packOp to a shape we can expand into stripMinedShape.
258 SmallVector<OpFoldResult> lows(packOp.getSourceRank(),
259 rewriter.getIndexAttr(0));
260 SmallVector<OpFoldResult> highs(packOp.getSourceRank(),
261 rewriter.getIndexAttr(0));
262 for (auto [pos, innerSize] :
263 llvm::zip_equal(packOp.getInnerDimsPos(), packOp.getMixedTiles())) {
264 int outerPos =
265 packedToStripMinedShapePerm[packingMetadata.outerPositions[pos]];
266 OpFoldResult origSize =
267 tensor::getMixedSize(rewriter, loc, packOp.getSource(), pos);
268 OpFoldResult outerSize =
269 tensor::getMixedSize(rewriter, loc, packOp.getDest(), outerPos);
270 AffineExpr s0, d0, d1;
271 bindDims(rewriter.getContext(), d0, d1);
272 bindSymbols(rewriter.getContext(), s0);
273 auto map = AffineMap::get(/*dimCount=*/2, /*symbolCount=*/1, d0 * s0 - d1);
275 rewriter, loc, map, {outerSize, origSize, innerSize});
276 }
277 RankedTensorType collapsed = tensor::CollapseShapeOp::inferCollapsedType(
278 RankedTensorType::Builder(packedTensorType).setShape(stripMinedShape),
279 packingMetadata.reassociations);
280 Value paddingValue = packOp.getPaddingValue();
281 if (!paddingValue) {
282 paddingValue = arith::ConstantOp::create(
283 rewriter, loc, rewriter.getZeroAttr(getElementTypeOrSelf(collapsed)));
284 }
285 auto padOp =
286 tensor::PadOp::create(rewriter, loc, collapsed, packOp.getSource(), lows,
287 highs, paddingValue, /*nofold=*/false);
288
289 LDBG() << "insertPositions: "
290 << llvm::interleaved(packingMetadata.insertPositions);
291 LDBG() << "outerPositions: "
292 << llvm::interleaved(packingMetadata.outerPositions);
293 LDBG() << "packedShape: " << llvm::interleaved(packedTensorType.getShape());
294 LDBG() << "packedToStripMinedShapePerm: "
295 << llvm::interleaved(packedToStripMinedShapePerm);
296 LDBG() << "reassociations: "
297 << llvm::interleaved(llvm::map_range(packingMetadata.reassociations,
299 LDBG() << "stripMinedShape: " << llvm::interleaved(stripMinedShape);
300 LDBG() << "collapsed type: " << collapsed;
301
302 if (lowerPadLikeWithInsertSlice && packOp.isLikePad()) {
303 // Pack ops which operate as simple pads may not produce legal
304 // tensor.insert_slice operations when the packed type does not rank reduce
305 // to the padded type.
306 SliceVerificationResult rankReduces =
307 isRankReducedType(packedTensorType, padOp.getResultType());
308
309 if (rankReduces == SliceVerificationResult::Success) {
310 // This pack is just a plain pad.
311 // Just insert the pad in the higher ranked tensor.
312 // Offsets.
313 SmallVector<OpFoldResult> zeros(packOp.getDestRank(),
314 rewriter.getIndexAttr(0));
315 // Strides.
316 SmallVector<OpFoldResult> ones(packOp.getDestRank(),
317 rewriter.getIndexAttr(1));
319 tensor::getMixedSizes(rewriter, loc, packOp.getDest());
320
321 auto insertSliceOp = tensor::InsertSliceOp::create(
322 rewriter, loc, /*source=*/padOp, /*dest=*/packOp.getDest(),
323 /*offsets=*/zeros, sizes, /*strides=*/ones);
324
325 LDBG() << "insert_slice op: " << insertSliceOp;
326
327 rewriter.replaceOp(packOp, insertSliceOp->getResults());
328
329 return LowerPackResult{padOp, /*reshapeOp=*/nullptr,
330 /*transposeOp=*/nullptr};
331 }
332 }
333
334 // 5. Expand from the padded result to the stripMinedShape.
335 auto expandShapeResultType =
336 RankedTensorType::Builder(packedTensorType).setShape(stripMinedShape);
337 auto reshapeOp = tensor::ExpandShapeOp::create(
338 rewriter, loc, expandShapeResultType, padOp.getResult(),
339 packingMetadata.reassociations, stripMinedMixedSizes);
340
341 // 6. Transpose stripMinedShape to packedShape.
342 SmallVector<int64_t> transpPerm =
343 invertPermutationVector(packedToStripMinedShapePerm);
344 auto transposeOp = linalg::TransposeOp::create(
345 rewriter, loc, reshapeOp.getResult(), packOp.getDest(), transpPerm);
346
347 LDBG() << "reshape op: " << reshapeOp;
348 LDBG() << "transpPerm: " << llvm::interleaved(transpPerm);
349 LDBG() << "transpose op: " << transposeOp;
350
351 // 7. Replace packOp by transposeOp.
352 rewriter.replaceOp(packOp, transposeOp->getResults());
353
354 return LowerPackResult{padOp, reshapeOp, transposeOp};
355}
356
357FailureOr<LowerUnPackOpResult>
358linalg::lowerUnPack(RewriterBase &rewriter, linalg::UnPackOp unPackOp,
359 bool lowerUnpadLikeWithExtractSlice) {
360 // Pack/unpack memref transformations are unsupported. The memref forms
361 // are mainly for bufferization and scalar lowering. Other uses are not
362 // recommended, see #225650 for details.
363 if (!unPackOp.hasPureTensorSemantics())
364 return failure();
365
366 Location loc = unPackOp->getLoc();
367 OpBuilder::InsertionGuard g(rewriter);
368 rewriter.setInsertionPoint(unPackOp);
369
370 auto packedTensorType = cast<RankedTensorType>(unPackOp.getSourceType());
371 int64_t packedRank = packedTensorType.getRank();
372
373 OpFoldResult zero = rewriter.getIndexAttr(0), one = rewriter.getIndexAttr(1);
374 auto destTensorType = cast<RankedTensorType>(unPackOp.getDest().getType());
375 if (lowerUnpadLikeWithExtractSlice && unPackOp.isLikeUnPad()) {
376 // This unpack is just a plain unpad.
377 // Just extract the slice from the higher ranked tensor.
378 ArrayRef<int64_t> destShape = destTensorType.getShape();
379 // The inner dimensions stay the same as the destination tensor, but the
380 // outer ones are additional 1s.
381 SmallVector<OpFoldResult> sizes(packedRank - destShape.size(), one);
382 sizes.append(tensor::getMixedSizes(rewriter, loc, unPackOp.getDest()));
383
384 auto extractSliceOp = tensor::ExtractSliceOp::create(
385 rewriter, loc, destTensorType, unPackOp.getSource(),
386 SmallVector<OpFoldResult>(packedRank, zero), sizes,
387 SmallVector<OpFoldResult>(packedRank, one));
388
389 rewriter.replaceOp(unPackOp, extractSliceOp->getResults());
390
391 return LowerUnPackOpResult{/*emptyOp=*/nullptr, /*transposeOp=*/nullptr,
392 /*reshapeOp=*/nullptr, extractSliceOp,
393 /*copyOp=*/nullptr};
394 }
395
396 // 1. Compute the permutation vector to shuffle packed shape into the shape
397 // before any outer or inner permutations have been applied.
398 PackingMetadata packingMetadata;
399 SmallVector<int64_t> packedToStripMinedShapePerm =
400 getUnPackInverseSrcPerm(unPackOp, packingMetadata);
401
402 // 2. Compute the stripMinedShape: this is the packed shape without outer and
403 // inner permutations.
404 SmallVector<int64_t> stripMinedShape(packedTensorType.getShape());
405 applyPermutationToVector(stripMinedShape, packedToStripMinedShapePerm);
406
407 // 3. Transpose packedShape to stripMinedShape.
408 RankedTensorType stripMinedTensorType =
409 RankedTensorType::Builder(packedTensorType).setShape(stripMinedShape);
410 RankedTensorType collapsedType = tensor::CollapseShapeOp::inferCollapsedType(
411 stripMinedTensorType, packingMetadata.reassociations);
412
413 // Get dynamic dims from input tensor based on packedToStripMinedShapePerm
414 // permutation.
416 tensor::getMixedSizes(rewriter, loc, unPackOp.getSource());
417 applyPermutationToVector(dims, packedToStripMinedShapePerm);
418 auto emptyOp = tensor::EmptyOp::create(rewriter, loc, dims,
419 stripMinedTensorType.getElementType());
420 auto transposeOp =
421 linalg::TransposeOp::create(rewriter, loc, unPackOp.getSource(), emptyOp,
422 packedToStripMinedShapePerm);
423
424 LDBG() << "insertPositions: "
425 << llvm::interleaved(packingMetadata.insertPositions);
426 LDBG() << "packedShape: " << llvm::interleaved(packedTensorType.getShape());
427 LDBG() << "packedToStripMinedShapePerm: "
428 << llvm::interleaved(packedToStripMinedShapePerm);
429 LDBG() << "reassociations: "
430 << llvm::interleaved(llvm::map_range(packingMetadata.reassociations,
432 LDBG() << "stripMinedShape: " << llvm::interleaved(stripMinedShape);
433 LDBG() << "collapsed type: " << collapsedType;
434
435 // 4. Collapse from the stripMinedShape to the padded result.
436 auto reshapeOp = tensor::CollapseShapeOp::create(
437 rewriter, loc, collapsedType, transposeOp->getResult(0),
438 packingMetadata.reassociations);
439
440 // 5. ExtractSlice.
441 int64_t destRank = destTensorType.getRank();
442 auto extractSliceOp = tensor::ExtractSliceOp::create(
443 rewriter, loc, destTensorType, reshapeOp->getResult(0),
444 SmallVector<OpFoldResult>(destRank, zero),
445 tensor::getMixedSizes(rewriter, loc, unPackOp.getDest()),
446 SmallVector<OpFoldResult>(destRank, one));
447
448 // 6. Inject a copy to preserve DPS.
449 auto copyOp = linalg::CopyOp::create(
450 rewriter, loc, extractSliceOp->getResult(0), unPackOp.getDest());
451
452 // 7. Replace unPackOp by copyOp.
453 rewriter.replaceOp(unPackOp, copyOp->getResults());
454
455 return LowerUnPackOpResult{emptyOp, transposeOp, reshapeOp, extractSliceOp,
456 copyOp};
457}
458
460PackedOperandsDimList::extractPackedDimsForOperand(int64_t operandPos) {
462 for (auto &i : spec) {
463 if (!i.packedDimForEachOperand[operandPos].has_value())
464 continue;
465 res.push_back(i.packedDimForEachOperand[operandPos].value());
466 }
467 return res;
468}
469
470SmallVector<OpFoldResult>
471PackedOperandsDimList::extractPackSizesForOperand(int64_t operandPos) {
472 SmallVector<OpFoldResult> res;
473 for (auto &i : spec) {
474 if (!i.packedDimForEachOperand[operandPos].has_value())
475 continue;
476 res.push_back(i.packedSize);
477 }
478 return res;
479}
480
481/// Implement packing of a single LinalgOp by performing packing by
482/// `packedSizes`. There must be one packedSizes entry per `linalgOp` iterator.
483/// Return the packed Linalg op on success, failure otherwise.
484FailureOr<PackResult> linalg::pack(RewriterBase &rewriter,
485 linalg::LinalgOp linalgOp,
486 ArrayRef<OpFoldResult> packedSizes) {
487 if (packedSizes.size() != linalgOp.getNumLoops()) {
488 return rewriter.notifyMatchFailure(linalgOp,
489 "incorrect number of pack sizes");
490 }
491 if (!linalgOp.hasPureTensorSemantics()) {
492 return rewriter.notifyMatchFailure(
493 linalgOp, "expects LinalgOp with pure tensor semantics");
494 }
495
496 Location loc = linalgOp->getLoc();
497 SmallVector<AffineMap> indexingMaps = linalgOp.getIndexingMapsArray();
499 linalgOp.getIteratorTypesArray();
500 LDBG() << "Start packing: " << linalgOp;
501 LDBG() << "maps: " << llvm::interleaved(indexingMaps);
502 LDBG() << "iterators: " << llvm::interleaved(iteratorTypes);
503
506 // Step 1. Pack each dim of the LinalgOp metadata by packedSizes[i].
507 PackedOperandsDimList listOfPackedOperandsDim;
508 for (int64_t i = 0, e = packedSizes.size(); i < e; ++i) {
509 std::optional<int64_t> maybeConstant = getConstantIntValue(packedSizes[i]);
510 // Skip tile sizes explicitly set to 0.
511 if (maybeConstant.has_value() && maybeConstant.value() == 0)
512 continue;
513
514 PackedOperandsDim packedOperandsDims;
515 packedOperandsDims.packedSize = packedSizes[i];
516 FailureOr<SmallVector<std::optional<int64_t>>>
517 maybePackedDimForEachOperand =
518 packLinalgMetadataOnce(indexingMaps, iteratorTypes, i);
519 if (failed(maybePackedDimForEachOperand))
520 return failure();
521 packedOperandsDims.packedDimForEachOperand = *maybePackedDimForEachOperand;
522
523 LDBG() << "++++ After pack size #" << i << ": " << packedSizes[i];
524 LDBG() << "maps: " << llvm::interleaved(indexingMaps);
525 LDBG() << "iterators: " << llvm::interleaved(iteratorTypes);
526 LDBG() << "packedDimForEachOperand: "
527 << llvm::interleaved(packedOperandsDims.packedDimForEachOperand);
528
529 listOfPackedOperandsDim.pushBack(std::move(packedOperandsDims));
530 }
531
532 // Step 2. Propagate packing to all LinalgOp operands.
533 SmallVector<Value> inputsAndInits, results;
534 SmallVector<OpOperand *> initOperands =
535 llvm::to_vector(llvm::make_pointer_range(linalgOp.getDpsInitsMutable()));
536 SmallVector<OpOperand *> inputOperands = linalgOp.getDpsInputOperands();
537 for (const auto &operandsList : {inputOperands, initOperands}) {
538 for (OpOperand *opOperand : operandsList) {
539 int64_t pos = opOperand->getOperandNumber();
540 Value operand = opOperand->get();
541 SmallVector<int64_t> innerPos =
542 listOfPackedOperandsDim.extractPackedDimsForOperand(pos);
543 SmallVector<OpFoldResult> innerPackSizes =
544 listOfPackedOperandsDim.extractPackSizesForOperand(pos);
545 LDBG() << "operand: " << operand;
546 LDBG() << "innerPos: " << llvm::interleaved(innerPos);
547 LDBG() << "innerPackSizes: " << llvm::interleaved(innerPackSizes);
548 if (innerPackSizes.empty()) {
549 inputsAndInits.push_back(operand);
550 continue;
551 }
552 Value dest = linalg::PackOp::createDestinationTensor(
553 rewriter, loc, operand, innerPackSizes, innerPos,
554 /*outerDimsPerm=*/{});
555 ShapedType operandType = cast<ShapedType>(operand.getType());
556 bool areConstantTiles =
557 llvm::all_of(innerPackSizes, [](OpFoldResult tile) {
558 return getConstantIntValue(tile).has_value();
559 });
560 if (areConstantTiles && operandType.hasStaticShape() &&
561 !linalg::PackOp::requirePaddingValue(
562 operandType.getShape(), innerPos,
563 cast<ShapedType>(dest.getType()).getShape(), {},
564 innerPackSizes)) {
565 packOps.push_back(linalg::PackOp::create(rewriter, loc, operand, dest,
566 innerPos, innerPackSizes));
567 } else {
568 // TODO: value of the padding attribute should be determined by
569 // consumers.
570 auto zeroAttr =
571 rewriter.getZeroAttr(getElementTypeOrSelf(dest.getType()));
572 Value zero = arith::ConstantOp::create(rewriter, loc, zeroAttr);
573 packOps.push_back(linalg::PackOp::create(
574 rewriter, loc, operand, dest, innerPos, innerPackSizes, zero));
575 }
576 inputsAndInits.push_back(packOps.back().getResult());
577 }
578 }
579
580 // Step 3. Build the packed op, use the type of `inits` as result types.
581 ValueRange inputs =
582 ValueRange{inputsAndInits}.take_front(linalgOp.getNumDpsInputs());
583 ValueRange inits =
584 ValueRange{inputsAndInits}.take_back(linalgOp.getNumDpsInits());
585 auto packedLinalgOp =
586 linalg::GenericOp::create(rewriter, linalgOp.getLoc(), inits.getTypes(),
587 inputs, inits, indexingMaps, iteratorTypes);
588 packedLinalgOp.getRegion().takeBody(linalgOp->getRegion(0));
589
590 // Step 4. Propagate packing to all the op results.
591 for (OpResult result : packedLinalgOp->getResults()) {
592 int64_t resultNum = result.getResultNumber();
593 linalg::PackOp maybePackedInit =
594 inits[resultNum].getDefiningOp<linalg::PackOp>();
595 if (!maybePackedInit) {
596 results.push_back(result);
597 continue;
598 }
599 // Build the symmetrical UnPackOp to the existing PackOp.
600 unPackOps.push_back(linalg::UnPackOp::create(
601 rewriter, packedLinalgOp->getLoc(), result, maybePackedInit.getSource(),
602 maybePackedInit.getInnerDimsPos(), maybePackedInit.getMixedTiles()));
603 results.push_back(unPackOps.back().getResult());
604 }
605
606 // Step 5. Replace `linalgOp`.
607 rewriter.replaceOp(linalgOp, results);
608
609 // Return packedLinalgOp.
610 return PackResult{packOps,
611 cast<linalg::LinalgOp>(packedLinalgOp.getOperation()),
612 unPackOps};
613}
614
615//===----------------------------------------------------------------------===//
616// packTranspose transformation.
617//===----------------------------------------------------------------------===//
618
619/// Return a copy of `tensorType` after permutation by `permutationVector`.
620// Note: Should be a new method in of MemRef/RankedTensor/VectorType::Builder
621// but this would introduce a dependence on Dialect in IR.
622// TODO: Restructure.
623static RankedTensorType permuteShape(RankedTensorType tensorType,
624 ArrayRef<int64_t> permutationVector) {
625 SmallVector<int64_t> shape(tensorType.getShape());
626 applyPermutationToVector(shape, permutationVector);
627 return RankedTensorType::Builder(tensorType).setShape(shape);
628}
629
630/// Return a new GenericOp obtained by transposing opOperand by the permutation
631/// vector:
632/// - the corresponding indexing map is transposed by `permutation`
633/// - the corresponding operand value is replaced by `transposedValue`
634/// `linalgOp` is replaced by the return op in the process.
635/// Asserts that `transposedValue` is of the proper transposed ShapedType.
637 RewriterBase &rewriter, LinalgOp linalgOp, OpOperand &opOperand,
638 ArrayRef<int64_t> permutation, Value transposedValue) {
639 // Sanity check the operand.
640 assert(linalgOp == opOperand.getOwner() && "linalg op must own the operand");
641
642 // Sanity check of the expected transposed tensor type.
643 auto tensorType = permuteShape(
644 cast<RankedTensorType>(opOperand.get().getType()), permutation);
645 (void)tensorType;
646 assert(tensorType == transposedValue.getType() &&
647 "expected tensor type mismatch");
648
649 // Compute the transposed indexing map.
650 // Sigh unsigned pollution.
651 SmallVector<unsigned> tmpTransposition =
652 llvm::map_to_vector(permutation, [](int64_t i) -> unsigned { return i; });
653 AffineMap permutationMap =
654 AffineMap::getPermutationMap(tmpTransposition, rewriter.getContext());
655 AffineMap transposedMap =
656 permutationMap.compose(linalgOp.getMatchingIndexingMap(&opOperand));
657
658 // Set the transposed indexing map in the proper position.
659 SmallVector<AffineMap> indexingMaps = linalgOp.getIndexingMapsArray();
660 indexingMaps[linalgOp.getIndexingMapIndex(&opOperand)] = transposedMap;
661 // Set the transposedValue in the proper operand position.
662 SmallVector<Value> operands = linalgOp->getOperands();
663 operands[opOperand.getOperandNumber()] = transposedValue;
664
665 ValueRange operandsRef(operands);
666 auto transposedGenericOp = linalg::GenericOp::create(
667 rewriter,
668 /*location=*/linalgOp->getLoc(),
669 /*resultTensorTypes=*/
670 operandsRef.drop_front(linalgOp.getNumDpsInputs()).getTypes(),
671 /*inputs=*/operandsRef.take_front(linalgOp.getNumDpsInputs()),
672 /*outputs=*/operandsRef.drop_front(linalgOp.getNumDpsInputs()),
673 /*indexingMaps=*/indexingMaps,
674 /*iteratorTypes=*/linalgOp.getIteratorTypesArray());
675 transposedGenericOp.getRegion().takeBody(linalgOp->getRegion(0));
676 rewriter.replaceOp(linalgOp, transposedGenericOp->getResults());
677
678 return cast<linalg::LinalgOp>(transposedGenericOp.getOperation());
679}
680
681FailureOr<PackTransposeResult>
682linalg::packTranspose(RewriterBase &rewriter, linalg::PackOp packOp,
683 linalg::LinalgOp linalgOp, linalg::UnPackOp maybeUnPackOp,
684 ArrayRef<int64_t> outerPerm,
685 ArrayRef<int64_t> innerPerm) {
686 Location loc = linalgOp.getLoc();
687
688 // Step 1. Transpose packOp.
689 rewriter.setInsertionPoint(packOp);
690 linalg::PackOp transposedPackOp =
691 packOp.createTransposedClone(rewriter, loc, innerPerm, outerPerm);
692
693 if (packOp.hasPureBufferSemantics() || !packOp.getResult().hasOneUse())
694 return rewriter.notifyMatchFailure(linalgOp, "expect single pack use");
695
696 OpOperand &packUse = *packOp->getUses().begin();
697 if (packUse.getOwner() != linalgOp) {
698 return rewriter.notifyMatchFailure(
699 linalgOp, "not a single use by the LinalgOp target");
700 }
701 if (maybeUnPackOp &&
702 (!linalgOp.isDpsInit(&packUse) ||
703 maybeUnPackOp.getSource() != linalgOp.getTiedOpResult(&packUse))) {
704 return rewriter.notifyMatchFailure(linalgOp,
705 "not produced by the LinalgOp target");
706 }
707
708 // Step 2. Transpose linalgOp.
709 // transposedPackOp.getOuterDimsPerm() may be empty, in which case it is the
710 // identity. Don't rely on it.
711 int64_t numLeadingDims = packOp.getSourceRank();
712 int64_t numTrailingDims = packOp.getInnerDimsPos().size();
713 // Step 2.a. Compute the permutation on the whole operand.
714 // Leading part just reuse the outerPerm.
715 SmallVector<int64_t> permutation(outerPerm);
716 if (permutation.empty())
717 llvm::append_range(permutation, llvm::seq<int64_t>(0, numLeadingDims));
718 // Trailing part needs to reindex positions by `numLeadingDims`.
719 if (innerPerm.empty()) {
720 llvm::append_range(
721 permutation,
722 llvm::seq<int64_t>(numLeadingDims, numLeadingDims + numTrailingDims));
723 } else {
724 llvm::append_range(permutation,
725 llvm::map_range(innerPerm, [&](int64_t pos) {
726 return numLeadingDims + pos;
727 }));
728 }
729 if (!isPermutationVector(permutation))
730 return rewriter.notifyMatchFailure(linalgOp, "invalid permutation");
731
732 // Step 2.b. Save the transposedPackUse operand number in case we need to
733 // get the tied OpResult after `linalgOp` has been replaced.
734 int64_t packUseOperandNumber = packUse.getOperandNumber();
735 // Step 2.c. Actually perform the transposition.
736 rewriter.setInsertionPoint(linalgOp);
737 linalg::LinalgOp transposedLinalgOp = transposeOneLinalgOperandAndReplace(
738 rewriter, linalgOp, packUse, permutation, transposedPackOp.getResult());
739
740 // Step 3. Maybe transpose unPackOp.
741 linalg::UnPackOp transposedUnPackOp;
742 if (maybeUnPackOp) {
743 OpOperand &opOperand =
744 transposedLinalgOp->getOpOperand(packUseOperandNumber);
745 OpResult transposedResult = transposedLinalgOp.getTiedOpResult(&opOperand);
746 rewriter.setInsertionPoint(maybeUnPackOp);
747 transposedUnPackOp = maybeUnPackOp.createTransposedClone(
748 rewriter, loc, transposedResult, innerPerm, outerPerm);
749
750 rewriter.replaceOp(maybeUnPackOp, transposedUnPackOp->getResults());
751 }
752
753 // Step 4. Finally, replace packOp now that we don't need it anymore.
754 if (packOp.hasPureTensorSemantics())
755 rewriter.replaceOp(packOp, transposedPackOp->getResults());
756 else
757 rewriter.eraseOp(packOp);
758
759 return PackTransposeResult{transposedPackOp, transposedLinalgOp,
760 transposedUnPackOp};
761}
762
763//===----------------------------------------------------------------------===//
764// packMatmulGreedily transformation.
765//===----------------------------------------------------------------------===//
766
767/// Pack a LinalgOp by greedily inferring matmul dimensions (m, n, k) where m
768/// and n are proper parallel dimensions and k is a proper reduction
769/// dimension. Packing occurs by rewriting the op as a linalg.generic and
770/// calling linalg::pack by `mnkPackedSizes`. The order of the packed
771/// dimensions is customizable: the `mnkOrder` is a permutation of {0, 1, 2}
772/// to reorder {m, n, k} into one of the 8 possible forms. The outer
773/// dimensions of the operands are not permuted at this time, this is left for
774/// future work.
775FailureOr<PackResult>
776linalg::packMatmulGreedily(RewriterBase &rewriter, LinalgOp linalgOp,
777 ArrayRef<OpFoldResult> mnkPackedSizes,
778 ArrayRef<int64_t> mnkPaddedSizesNextMultipleOf,
779 ArrayRef<int64_t> mnkOrder) {
780 assert(mnkPackedSizes.size() == 3 && "unexpected num of packing sizes");
781 assert((mnkPaddedSizesNextMultipleOf.empty() ||
782 mnkPaddedSizesNextMultipleOf.size() == 3) &&
783 "num of packing sizes next multiple should be empty or of size 3");
784 assert(mnkOrder.size() == 3 && "unexpected mnkOrder size");
785 assert(isPermutationVector(mnkOrder) && "expected a permutation");
786
787 int64_t numLoops = linalgOp.getNumLoops();
788 if (numLoops <= 2) {
789 LDBG() << "need 3+ loops to find a matmul to pack, got " << numLoops
790 << " in: " << linalgOp;
791 return rewriter.notifyMatchFailure(
792 linalgOp, "need 3+ loops to find a matmul to pack");
793 }
794
795 // Locally adjust the desired iterator position of mnk and packing sizes.
796 int64_t numPackedDims = mnkPackedSizes.size();
797 SmallVector<int64_t> mmnnkkPos(numPackedDims);
798 for (int64_t i = 0, e = numPackedDims; i < e; ++i)
799 mmnnkkPos[i] = numLoops - numPackedDims + mnkOrder[i];
800 SmallVector<OpFoldResult> packedSizes(numPackedDims);
801 for (int64_t i = 0, e = numPackedDims; i < e; ++i)
802 packedSizes[mnkOrder[i]] = mnkPackedSizes[i];
803 SmallVector<int64_t> paddedSizesNextMultipleOf(numPackedDims);
804 for (int64_t i = 0, e = numPackedDims; i < e; ++i) {
805 paddedSizesNextMultipleOf[mnkOrder[i]] =
806 mnkPaddedSizesNextMultipleOf.empty() ? 0
807 : mnkPaddedSizesNextMultipleOf[i];
808 }
809
810 // 1. Infer dims that are important for matmul.
811 FailureOr<ContractionDimensions> maybeDimensions =
812 inferContractionDims(linalgOp);
813 if (failed(maybeDimensions)) {
814 LDBG() << "couldn't infer matmul iterators in: " << linalgOp;
815 return rewriter.notifyMatchFailure(linalgOp,
816 "couldn't infer matmul iterators");
817 }
818
819 // 2. Normalize linalgOp to an kmn-matmul-like with [red, par, par] most
820 // minor iterators. In cases with multiple options for m, n, k bias towards
821 // the most minor embedding.
822 // If we wanted a different normalization order, this is where it would have
823 // to plug a heuristic.
824 int64_t mPos = maybeDimensions->m.back(), nPos = maybeDimensions->n.back(),
825 kPos = maybeDimensions->k.back();
826 LDBG() << "Start packing generic op greedily with (m@" << mPos << ", n@"
827 << nPos << ", k@" << kPos << "): " << linalgOp;
828
829 // 2.a. Rewrite as a generic.
830 auto genericOp = dyn_cast<GenericOp>(linalgOp.getOperation());
831 if (!genericOp) {
832 FailureOr<LinalgOp> generalizeResult =
833 generalizeNamedOp(rewriter, linalgOp);
834 assert(succeeded(generalizeResult) &&
835 isa<GenericOp>(generalizeResult->getOperation()) &&
836 "unexpected failure generalizing op");
837 genericOp = cast<GenericOp>(generalizeResult->getOperation());
838 }
839
840 // 2.b. Interchange to move the dimensions (k, m, n) as most-minor
841 // iterators. Note that this only normalized the iteration order and does
842 // not change the indexings of any operand.
843 SmallVector<int64_t> permutation =
844 computePermutationVector(numLoops, {mPos, nPos, kPos}, mmnnkkPos);
845 LDBG() << "perm: " << llvm::interleaved(permutation);
846 // Sign .. unsigned pollution.
847 SmallVector<unsigned> unsignedPerm(permutation.begin(), permutation.end());
848 FailureOr<GenericOp> interchangeResult =
849 interchangeGenericOp(rewriter, genericOp, unsignedPerm);
850 assert(succeeded(interchangeResult) && "unexpected failure interchanging op");
851 genericOp = *interchangeResult;
852 LDBG() << "Generalized Op to pack: " << genericOp;
853
854 // At this point, the op iterators are normalized to {leading, k, m, n}.
855 // The layouts induced by packing will always be:
856 // - LHS{leading_lhs, kk, mm}
857 // - RHS{leading_rhs, kk, nn}
858 // - RES{leading_res, mm, nn}
859 // If we wanted to change the packed order, we would reorder (k, m, n) to
860 // something else above.
861 //
862 // Additional permutations of the outer dims of the operands (i.e.
863 // leading_lhs, leading_rhs and leading_res) could follow by computing the
864 // desired outerPerm for each operand.
865 // This is left for future work.
866
867 // TODO: this creates too much IR, go use reifyResultShapes.
868 SmallVector<Range, 4> loopRanges =
869 cast<LinalgOp>(genericOp.getOperation())
870 .createLoopRanges(rewriter, genericOp.getLoc());
871
872 // Add leading zeros to match numLoops, we only pack the last 3 dimensions
873 // post interchange.
874 LDBG() << "paddedSizesNextMultipleOf: "
875 << llvm::interleaved(paddedSizesNextMultipleOf);
876 LDBG() << "loopRanges: "
877 << llvm::interleaved(
878 llvm::map_range(loopRanges, [](Range r) { return r.size; }));
879 SmallVector<OpFoldResult> adjustedPackedSizes(numLoops - packedSizes.size(),
880 rewriter.getIndexAttr(0));
881 for (int64_t i = 0, e = numPackedDims; i < e; ++i) {
882 if (paddedSizesNextMultipleOf[i] == 0) {
883 adjustedPackedSizes.push_back(packedSizes[i]);
884 continue;
885 }
886 AffineExpr d0, s0;
887 bindDims(rewriter.getContext(), d0);
888 bindSymbols(rewriter.getContext(), s0);
889 adjustedPackedSizes.push_back(affine::makeComposedFoldedAffineApply(
890 rewriter, genericOp->getLoc(), d0.ceilDiv(s0) * s0,
891 {loopRanges[adjustedPackedSizes.size()].size,
892 rewriter.getIndexAttr(paddedSizesNextMultipleOf[i])}));
893 }
894 LDBG() << "adjustedPackedSizes: " << llvm::interleaved(adjustedPackedSizes);
895
896 // TODO: If we wanted to give the genericOp a name after packing, after
897 // calling `pack` would be a good time. One would still need to check that
898 // `containsMostMinorMatmul(packingRes->packedLinalgOp)` is true, since we
899 // also allow degenerate matmul cases (i.e. matvec, dot).
900 return pack(rewriter, genericOp, adjustedPackedSizes);
901}
902
903//===----------------------------------------------------------------------===//
904// Transformations exposed as rewrite patterns.
905//===----------------------------------------------------------------------===//
906
909 assert(!tileSizeComputationFunction && "tile sizes already set");
910 SmallVector<int64_t, 4> tileSizes(ts);
911 tileSizeComputationFunction = [tileSizes](OpBuilder &b, Operation *op) {
913 b.setInsertionPointToStart(
914 &op->getParentOfType<func::FuncOp>().getBody().front());
915 return llvm::map_to_vector<4>(tileSizes, [&](int64_t s) {
916 Value v = arith::ConstantIndexOp::create(b, op->getLoc(), s);
917 return v;
918 });
919 };
920 return *this;
921}
922
924 memref::CopyOp copyOp, PatternRewriter &rewriter) const {
925 return vectorizeCopy(rewriter, copyOp);
926}
927
928/// Filling `dest` using FillOp constant padding value if possible.
929/// Otherwise, generate a tensor::GenerateOp.
931 RewriterBase &rewriter, tensor::PadOp padOp, Value dest,
932 const SmallVector<Value> &dynSizes) const {
933 auto padValue = padOp.getConstantPaddingValue();
934 if (padValue) {
935 // Move the padding value defined inside the PadOp block to outside.
936 if (padValue.getParentBlock() == &padOp.getRegion().front())
937 rewriter.moveOpBefore(padValue.getDefiningOp(), padOp);
938 return FillOp::create(rewriter, padOp.getLoc(), padValue, dest).result();
939 }
940
941 // Fill could not be optimized: Lower to tensor::GenerateOp with region.
942 auto generateOp = tensor::GenerateOp::create(rewriter, padOp.getLoc(),
943 padOp.getResultType(), dynSizes);
944 // Copy region to new op.
945 IRMapping bvm;
946 padOp.getRegion().cloneInto(&generateOp.getRegion(), bvm);
947 return generateOp;
948}
949
950LogicalResult
952 PatternRewriter &rewriter) const {
953 // Given an OpFoldResult, return an index-typed value.
954 auto getIdxValue = [&](OpFoldResult ofr) {
955 if (auto val = llvm::dyn_cast_if_present<Value>(ofr))
956 return val;
958 rewriter, padOp.getLoc(),
959 cast<IntegerAttr>(cast<Attribute>(ofr)).getInt())
960 .getResult();
961 };
962
963 auto resultType = padOp.getResultType();
964 // Compute size of EmptyOp. Any combination of static/dynamic is supported.
965 SmallVector<Value> dynSizes;
966 SmallVector<int64_t> staticSizes;
967 for (unsigned dim = 0; dim < resultType.getRank(); ++dim) {
968 if (resultType.isDynamicDim(dim)) {
969 auto srcSize = getIdxValue(tensor::getMixedSize(rewriter, padOp.getLoc(),
970 padOp.getSource(), dim));
971 // Add low and high padding value.
972 auto plusLow = rewriter.createOrFold<arith::AddIOp>(
973 padOp.getLoc(), srcSize, getIdxValue(padOp.getMixedLowPad()[dim]));
974 auto plusHigh = rewriter.createOrFold<arith::AddIOp>(
975 padOp.getLoc(), plusLow, getIdxValue(padOp.getMixedHighPad()[dim]));
976 dynSizes.push_back(plusHigh);
977 }
978 staticSizes.push_back(resultType.getDimSize(dim));
979 }
980
981 // Init tensor and fill it with padding.
982 Value emptyTensor =
983 tensor::EmptyOp::create(rewriter, padOp.getLoc(), staticSizes,
984 resultType.getElementType(), dynSizes);
985 Value fill = createFillOrGenerateOp(rewriter, padOp, emptyTensor, dynSizes);
986
987 // Generate a InsertSliceOp for copying the PadOp source.
988 auto sourceType = padOp.getSourceType();
989 // Compute size of source of tensor::PadOp.
991 tensor::getMixedSizes(rewriter, padOp.getLoc(), padOp.getSource());
992 // Strides of InsertSliceOp are all 1.
993 SmallVector<OpFoldResult> strides(sourceType.getRank(),
994 rewriter.getIndexAttr(1));
995 rewriter.replaceOpWithNewOp<tensor::InsertSliceOp>(
996 padOp, padOp.getSource(), fill, padOp.getMixedLowPad(), srcSizes,
997 strides);
998
999 return success();
1000}
1001
1003 tensor::ExtractSliceOp sliceOp, PatternRewriter &rewriter) const {
1004 if (!sliceOp.hasUnitStride())
1005 return failure();
1006
1007 auto padOp = sliceOp.getSource().getDefiningOp<tensor::PadOp>();
1008 if (!padOp)
1009 return failure();
1010
1011 bool zeroSliceGuard = true;
1012 if (controlFn) {
1013 if (std::optional<bool> control = controlFn(sliceOp))
1014 zeroSliceGuard = *control;
1015 else
1016 return failure();
1017 }
1018
1019 FailureOr<TilingResult> tilingResult =
1020 tensor::bubbleUpPadSlice(rewriter, padOp, sliceOp.getMixedOffsets(),
1021 sliceOp.getMixedSizes(), zeroSliceGuard);
1022 if (failed(tilingResult))
1023 return failure();
1024
1025 RankedTensorType sourceType = sliceOp.getSourceType();
1026 RankedTensorType resultType = sliceOp.getResultType();
1027
1028 // If the extract_slice is not rank-reduced, all shapes are static and the
1029 // data source is actually used. Rewrite into pad(extract_slice(x)).
1030 if (sourceType.getRank() == resultType.getRank()) {
1031 rewriter.replaceOp(sliceOp, tilingResult->tiledValues);
1032 return success();
1033 }
1034
1035 // Handle rank-reduced slice by creating another extract_slice op.
1037 rewriter, sliceOp.getLoc(), tilingResult->tiledValues[0], resultType);
1038
1039 rewriter.replaceOp(sliceOp, rankReduced);
1040 return success();
1041}
1042
1043/// If padding value is set, returns a tensor.pad Op for the source tensor,
1044/// with the output shape matching the output of `packOp`. Otherwise, returns
1045/// the source directly.
1046///
1047/// This method assumes that all outer dims for this pack Op are 1.
1049 linalg::PackOp packOp) {
1050 Value input = packOp.getSource();
1051 // Pack/unpack memref transformations are unsupported. The memref forms
1052 // are mainly for bufferization and scalar lowering. Other uses are not
1053 // recommended, see #225650 for details.
1054 if (!packOp.hasPureTensorSemantics())
1055 return input;
1056
1057 if (!packOp.getPaddingValue()) {
1058 return input;
1059 }
1060
1061 assert(llvm::all_of(packOp.getAllOuterDims(),
1062 [](int64_t val) { return val == 1; }) &&
1063 "some outer dims are != 1");
1064
1065 Location loc = packOp.getLoc();
1066 ShapedType inputType = packOp.getSourceType();
1067 int64_t inputRank = inputType.getRank();
1068
1069 DenseMap<int64_t, OpFoldResult> tileAndPosMapping =
1070 packOp.getDimAndTileMapping();
1071
1072 // The sizes of dynamic tiles
1073 SmallVector<Value> dynamicTileSizes;
1074
1075 // Collect dims for the padded shape.
1076 SmallVector<int64_t> paddedShape;
1077 for (int64_t dimIdx = 0; dimIdx < inputRank; ++dimIdx) {
1078 // 1. Non-tiled outer dims.
1079 // These dims should be 1 and we simply preserve them.
1080 if (!tileAndPosMapping.count(dimIdx)) {
1081 int64_t inputDimSize = inputType.getDimSize(dimIdx);
1082 assert(inputDimSize == 1 &&
1083 "with all outer dims == 1, this non-tiled input dim should be 1!");
1084 paddedShape.push_back(inputDimSize);
1085 continue;
1086 }
1087
1088 // 2. Tiled outer dims
1089 // As all outer dims == 1, it is safe to use the tile size for the padded
1090 // shape.
1091 OpFoldResult tileSizeForDim = tileAndPosMapping.lookup(dimIdx);
1092
1093 // 2.1 Static tile sizes
1094 std::optional<int64_t> cstTileSize = getConstantIntValue(tileSizeForDim);
1095 if (cstTileSize.has_value()) {
1096 paddedShape.push_back(cstTileSize.value());
1097 continue;
1098 }
1099
1100 // 2.2 Dynamic tile sizes
1101 paddedShape.push_back(ShapedType::kDynamic);
1102
1103 // Get the value that holds the dynamic size.
1104 dynamicTileSizes.push_back(llvm::dyn_cast<Value>(tileSizeForDim));
1105 }
1106 auto resultType =
1107 RankedTensorType::get(paddedShape, inputType.getElementType());
1108 return tensor::createPadHighOp(resultType, input, packOp.getPaddingValue(),
1109 /*nofold=*/false, loc, builder,
1110 dynamicTileSizes);
1111}
1112
1113// Normalizes a permutation on a higher rank space to its actual size, e.g.
1114// perm = [1, 4, 2]
1115// becomes
1116// norm = [0, 2, 1]
1117static SmallVector<int64_t>
1119 constexpr int64_t kNonTiledMarker = -1;
1120 SmallVector<int64_t> vec(rank, kNonTiledMarker);
1121 for (auto [index, value] : llvm::enumerate(perm))
1122 vec[value] = index;
1123 SmallVector<int64_t> normalizedPerm = llvm::filter_to_vector(
1124 vec, [&](int64_t v) { return v != kNonTiledMarker; });
1125 // This inverts the permutation in addition to normalizing so invert back.
1126 return invertPermutationVector(normalizedPerm);
1127}
1128
1129// Gets the normalized permutation implied by innerDimsPos and outerDimsPerm
1130// assuming rank reduction of unit outer dims.
1131static SmallVector<int64_t>
1133 ArrayRef<int64_t> innerDimsPos,
1134 ArrayRef<int64_t> outerDimsPerm) {
1135 SmallVector<int64_t> rankReducedOuterDimsPerm;
1136 SmallVector<int64_t> outerDims;
1137 SmallVector<int64_t> innerDims;
1138 int64_t dim = 0;
1139 int64_t unpackedRank = shape.size();
1140 for (auto i : llvm::seq<unsigned>(0, unpackedRank)) {
1141 if (llvm::is_contained(innerDimsPos, i)) {
1142 innerDims.push_back(dim++);
1143 continue;
1144 }
1145 if (shape[i] == 1)
1146 continue;
1147 outerDims.push_back(dim++);
1148 if (!outerDimsPerm.empty())
1149 rankReducedOuterDimsPerm.push_back(outerDimsPerm[i]);
1150 }
1151
1152 // Get the position of the inner dims after permutation.
1153 SmallVector<int64_t> innerPerm =
1154 getPackUnpackNormalizedPerm(unpackedRank, innerDimsPos);
1155 applyPermutationToVector<int64_t>(innerDims, innerPerm);
1156
1157 // Ditto for the outer dims.
1158 SmallVector<int64_t> perm = outerDims;
1159
1160 rankReducedOuterDimsPerm =
1161 getPackUnpackNormalizedPerm(unpackedRank, rankReducedOuterDimsPerm);
1162 if (!rankReducedOuterDimsPerm.empty())
1163 applyPermutationToVector<int64_t>(perm, rankReducedOuterDimsPerm);
1164
1165 // The tile always ends up as the inner most dims after packing.
1166 perm.append(innerDims);
1167
1168 return perm;
1169}
1170
1172 linalg::PackOp packOp, PatternRewriter &rewriter) const {
1173 // Pack/unpack memref transformations are unsupported. The memref forms
1174 // are mainly for bufferization and scalar lowering. Other uses are not
1175 // recommended, see #225650 for details.
1176 if (!packOp.hasPureTensorSemantics())
1177 return failure();
1178
1179 if (llvm::any_of(packOp.getTiledOuterDims(),
1180 [](int64_t dim) { return dim != 1; })) {
1181 return rewriter.notifyMatchFailure(
1182 packOp, "not all outer dimensions of the result are 1s");
1183 }
1184
1185 // When a padding value is set, getPackOpSourceOrPaddedSource only supports
1186 // the case where every outer dim (including un-tiled ones) is 1. Bail out
1187 // instead of hitting an assertion on a non-unit un-tiled outer dim.
1188 // FIXME: Handle this case by decomposing the padded pack instead of bailing
1189 // out; a non-unit un-tiled outer dim should be supported here.
1190 if (packOp.getPaddingValue() &&
1191 llvm::any_of(packOp.getAllOuterDims(),
1192 [](int64_t dim) { return dim != 1; })) {
1193 return rewriter.notifyMatchFailure(
1194 packOp, "cannot decompose padded pack with a non-unit un-tiled outer "
1195 "dimension");
1196 }
1197
1198 ArrayRef<int64_t> innerDimsPos = packOp.getInnerDimsPos();
1199 auto outerDimsPerm = packOp.getOuterDimsPerm();
1200
1201 // Verify that there are no:
1202 // * non-unit + un-tiled-outer-dims,
1203 // that are permuted. Supporting such cases would require refining the logic
1204 // that generates the Transpose Op.
1205 if (!llvm::all_of(outerDimsPerm, [&innerDimsPos, &packOp](int64_t dim) {
1206 static int prev = 0;
1207 // Skip tiled dims - these can be permuted.
1208 if (llvm::is_contained(innerDimsPos, dim))
1209 return true;
1210
1211 // Check whether this dim has been permuted. Permuting unit dims is fine
1212 // as that's effectively a no-op.
1213 if (dim < prev && (packOp.getResult().getType().getShape()[prev] != 1 ||
1214 packOp.getResult().getType().getShape()[dim] != 1))
1215 return false;
1216
1217 prev = dim;
1218 return true;
1219 })) {
1220 return rewriter.notifyMatchFailure(
1221 packOp, "At least one non-unit and un-tiled outer dim is permuted, "
1222 "this is not supported ATM!");
1223 }
1224
1225 Location loc = packOp.getLoc();
1226
1227 int64_t srcRank = packOp.getSourceRank();
1228
1229 // 1. Get the input that is going to be packed. If the input requires padding,
1230 // add a padding operation and return that as the input.
1231 Value input = getPackOpSourceOrPaddedSource(rewriter, packOp);
1232
1233 // 2. Transpose the input to match the inner tile order:
1234 // %init = tensor.empty()
1235 // %transposed_tile = linalg.transpose ins(%source_or_padded_source),
1236 // outs(%init)
1237 // Assumptions made:
1238 // - All tiled outer dims are 1 - the corresponding transposition order
1239 // doesn't matter, but requires all dim indices to be present.
1240 // - Un-tiled outer dims remain un-permuted.
1241
1242 // 2.1 Get the permutation for linalg.transpose:
1243 // [ untiled-dims, inner-dims-pos ]
1244 // Note, this logic assumes that the untiled dims are not permuted.
1245 SmallVector<int64_t> srcPermForTranspose;
1246 for (int64_t i = 0; i < srcRank; i++) {
1247 // We assume the `k` dimensions of the inner dim position, where `k` is the
1248 // rank of the inner tiling, correspond to the last `k` indices of the
1249 // transpose permutation. This is done by adding the indices not contained
1250 // in the inner dimension position in order from 0 to `n`. Where n is the
1251 // rank of the source tensor. For example if we have a source tensor with
1252 // indices [0, 1, 2, 3] and inner dim position of [3, 0], the remaining
1253 // indices are [1, 2]. and the transpose will be [1, 2, 3, 0].
1254 if (llvm::is_contained(innerDimsPos, i))
1255 continue;
1256 srcPermForTranspose.push_back(i);
1257 }
1258 srcPermForTranspose.append(innerDimsPos.begin(), innerDimsPos.end());
1259
1260 // 2.2 Create the init tensor for linalg.transpose with the correct shape:
1261 // [ untiled-dims, tiled-dims ]
1262 ShapedType inputTy = cast<ShapedType>(input.getType());
1263 SmallVector<OpFoldResult> shapeForEmptyOp;
1264 for (int64_t i = 0; i < srcRank; i++) {
1265 if (llvm::is_contained(innerDimsPos, i)) {
1266 // The tiled dims are appended after this loop.
1267 continue;
1268 }
1269 if (inputTy.isStaticDim(i))
1270 shapeForEmptyOp.push_back(rewriter.getIndexAttr(inputTy.getShape()[i]));
1271 else
1272 shapeForEmptyOp.emplace_back(
1273 tensor::DimOp::create(rewriter, loc, input, i).getResult());
1274 }
1275 shapeForEmptyOp.append(packOp.getMixedTiles());
1276
1277 // getMixedTiles() may contain Values pointing to constant ops (as opposed to
1278 // constant attributes with the corresponding value). Replace those with
1279 // attributes. This is to match the behaviour in
1280 // `getPackOpSourceOrPaddedSource`, which replaces constant SSA values with
1281 // attributes.
1282 llvm::transform(shapeForEmptyOp, shapeForEmptyOp.begin(),
1283 [&](OpFoldResult ofr) {
1284 if (auto val = llvm::dyn_cast<Value>(ofr))
1285 return getAsOpFoldResult(val);
1286 return ofr;
1287 });
1288
1289 LDBG() << "Pack permutation: " << packOp;
1290 LDBG() << "perm: " << llvm::interleaved(srcPermForTranspose);
1291 LDBG() << "Shape of empty tensor: " << llvm::interleaved(shapeForEmptyOp);
1292
1293 Value empty = tensor::EmptyOp::create(
1294 rewriter, loc, shapeForEmptyOp, packOp.getSourceType().getElementType());
1295
1296 // 2.3 Create linalg.transpose
1297 auto transposedOp = linalg::TransposeOp::create(rewriter, loc, input, empty,
1298 srcPermForTranspose);
1299
1300 // 3. Insert the inner tile into the destination tensor:
1301 // %inserted_tile = tensor.insert_slice(%transposed_tile)
1302
1303 // Compute the sizes attribute:
1304 // [ outer-dims, tile-sizes ]
1305 // Note that the output from the transpose Op excludes the tiled outer dims.
1306 // However, given the assumption that:
1307 // * all tiled outer dims == 1,
1308 // we can just use a rank-expanding tensor.insert_slice.
1309 SmallVector<OpFoldResult> writeSizes;
1310 for (auto size : packOp.getAllOuterDims()) {
1311 writeSizes.push_back(rewriter.getIndexAttr(size));
1312 }
1313
1314 for (auto tileSize : packOp.getMixedTiles()) {
1315 auto [_, tileSizeOfr] =
1316 getSimplifiedOfrAndStaticSizePair(tileSize, rewriter);
1317 writeSizes.push_back(tileSizeOfr);
1318 }
1319
1320 auto insert = tensor::InsertSliceOp::create(
1321 rewriter, loc, transposedOp.getResult()[0], packOp.getDest(), writeSizes);
1322
1323 // 4. Replace tensor.packOp with tensor.insert_slice created above
1324 rewriter.replaceOp(packOp, insert.getResult());
1325
1326 return success();
1327}
1328
1330 linalg::UnPackOp unpackOp, PatternRewriter &rewriter) const {
1331 if (!unpackOp.hasPureTensorSemantics())
1332 return failure();
1333
1334 int64_t destRank = unpackOp.getDestRank();
1335 ArrayRef<int64_t> srcShape = unpackOp.getSourceType().getShape();
1336 ArrayRef<int64_t> innerDimsPos = unpackOp.getInnerDimsPos();
1337 if (llvm::any_of(unpackOp.getTiledOuterDims(),
1338 [](int64_t dim) { return dim != 1; })) {
1339 return rewriter.notifyMatchFailure(
1340 unpackOp,
1341 "require the tiled outer dimensions of the result are all 1s");
1342 }
1343
1344 // 1. Use rank-reduced tensor.extract_slice op to extract the tile:
1345 // %extracted_tile = tensor.extract_slice(%unpack_op_input)
1346 Location loc = unpackOp.getLoc();
1347 Value source = unpackOp.getSource();
1348 DenseMap<int64_t, OpFoldResult> dimAndTileMapping =
1349 unpackOp.getDimAndTileMapping();
1350 Attribute oneIdxAttr = rewriter.getIndexAttr(1);
1351
1352 // The shape for ExtractSliceOp. Note that this will consist of 3 blocks of
1353 // dims:
1354 // [ outer-untiled-dims, outer-tiled-dims, tile-sizes ]
1355 SmallVector<int64_t> readShapeForExtractSlice;
1356 // The sizes attribute for ExtractSliceOp. Due to rank-reducing (and
1357 // outer-tiled-dims being all 1), this will be
1358 // [ outer-untiled-dims, tile-sizes ]
1359 SmallVector<OpFoldResult> extractSliceSizes;
1360
1361 // Shape for EmptyOp that's used as the init value for TransposeOp below.
1362 // This should be:
1363 // [ outer-untiled-dims, tile-sizes ]
1364 // However, skip unit dims - TransposeOp (below) applies rank-reduced
1365 // permutation.
1366 SmallVector<OpFoldResult> shapeForEmptyOp;
1367
1368 for (auto i : llvm::seq<unsigned>(0, destRank)) {
1369 // Compute sizes attribute for ExtractSliceOp - outer-tiled-dims.
1370 //
1371 // As all outer tiled dims are 1, so the corresponding
1372 // slice size to read will also 1. As this will be rank-reducing "extract
1373 // slice" (i.e. the unit dims will be "collapsed"), there's no need to
1374 // update:
1375 // * the output shape for ExtractSliceOp, nor
1376 // * the shape for EmptyOp.
1377 if (dimAndTileMapping.count(i)) {
1378 extractSliceSizes.push_back(oneIdxAttr);
1379 continue;
1380 }
1381
1382 // Compute sizes attribute for ExtractSliceOp + EmptyOp -
1383 // outer-untiled-dims
1384 if (ShapedType::isDynamic(srcShape[i])) {
1385 OpFoldResult dynamicDim =
1386 tensor::DimOp::create(rewriter, loc, source, i).getResult();
1387 extractSliceSizes.push_back(dynamicDim);
1388 shapeForEmptyOp.push_back(dynamicDim);
1389 } else {
1390 extractSliceSizes.push_back(rewriter.getIndexAttr(srcShape[i]));
1391 if (srcShape[i] != 1)
1392 shapeForEmptyOp.push_back(rewriter.getIndexAttr(srcShape[i]));
1393 }
1394 // Compute the output shape for ExtractSliceOp - outer-untiled-dims (take
1395 // into account rank-reducing)
1396 if (srcShape[i] != 1) {
1397 readShapeForExtractSlice.push_back(srcShape[i]);
1398 }
1399 }
1400 // Append the tile sizes to "sizes attribute" for ExtractSliceOp and the
1401 // shape for EmptyOp.
1402 auto mixedTiles = unpackOp.getMixedTiles();
1403 extractSliceSizes.append(mixedTiles.begin(), mixedTiles.end());
1404 shapeForEmptyOp.append(mixedTiles.begin(), mixedTiles.end());
1405
1406 // Explicitly create the type for extract_slice op because the inner tile
1407 // size could be 1. We want to represent the whole inner tile in this case.
1408 auto tileShape = srcShape.drop_front(destRank);
1409 // Append the inner tile shape to the permuted and rank-reduced outer shape.
1410 readShapeForExtractSlice.append(tileShape.begin(), tileShape.end());
1411 Type elemType = unpackOp.getSourceType().getElementType();
1412 auto readType = RankedTensorType::get(readShapeForExtractSlice, elemType);
1413 Value innerTile = tensor::ExtractSliceOp::create(
1414 rewriter, loc, readType, unpackOp.getSource(), extractSliceSizes);
1415
1416 // 2. Transpose the tile to match the outer corresponding tile order.
1418 srcShape.take_front(destRank), innerDimsPos, unpackOp.getOuterDimsPerm());
1419 // Unpack is a transition out of packed space so we invert the permutation.
1420 perm = invertPermutationVector(perm);
1421 applyPermutationToVector<OpFoldResult>(shapeForEmptyOp, perm);
1422
1423 Value empty =
1424 tensor::EmptyOp::create(rewriter, loc, shapeForEmptyOp, elemType);
1425 auto transposedOp =
1426 linalg::TransposeOp::create(rewriter, loc, innerTile, empty, perm);
1427
1428 // 3. Handle in-complete tiles if needed. It truncates trailing data from the
1429 // transposed tile.
1430 SmallVector<OpFoldResult> tileSizes;
1431 ArrayRef<int64_t> destShape = unpackOp.getDestType().getShape();
1432 for (auto i : llvm::seq<unsigned>(0, destRank)) {
1433 if (dimAndTileMapping.count(i) || destShape[i] != 1)
1434 tileSizes.push_back(
1435 tensor::getMixedSize(rewriter, loc, unpackOp.getDest(), i));
1436 }
1437
1438 auto partialTile =
1439 tensor::ExtractSliceOp::create(rewriter, loc, RankedTensorType(),
1440 transposedOp.getResult()[0], tileSizes);
1441
1442 // 4. Insert the result to the destination tensor.
1443 SmallVector<OpFoldResult> writeSizes;
1444 for (int i = 0, idx = 0; i < destRank; ++i) {
1445 if (dimAndTileMapping.count(i) || destShape[i] != 1)
1446 writeSizes.push_back(tileSizes[idx++]);
1447 else
1448 writeSizes.push_back(oneIdxAttr);
1449 }
1450 auto insert = tensor::InsertSliceOp::create(rewriter, loc, partialTile,
1451 unpackOp.getDest(), writeSizes);
1452 rewriter.replaceOp(unpackOp, insert.getResult());
1453
1454 return success();
1455}
1456
1457//===----------------------------------------------------------------------===//
1458// Generic DownscaleSizeOneWindowedConvolution
1459//===----------------------------------------------------------------------===//
1460//
1461/// Returns the indices of affine map results that reference any of the given
1462/// dimensions.
1465 SmallVector<unsigned> resultIndices;
1466 for (unsigned dim : dims) {
1467 for (unsigned i = 0, e = map.getNumResults(); i < e; ++i) {
1468 AffineExpr expr = map.getResult(i);
1469 if (expr.isFunctionOfDim(dim)) {
1470 resultIndices.push_back(i);
1471 break;
1472 }
1473 }
1474 }
1475 return resultIndices;
1476}
1477
1478/// Helper to create a rank-reducing extract_slice that removes specific
1479/// dimensions from a tensor.
1481 Location loc, Value tensor,
1482 ArrayRef<unsigned> dimsToRemove) {
1483 auto tensorType = cast<RankedTensorType>(tensor.getType());
1484 int64_t rank = tensorType.getRank();
1485
1486 // Compute new shape by removing the specified dimensions.
1487 SmallVector<int64_t> newShape;
1488 for (int64_t i = 0; i < rank; ++i) {
1489 if (!llvm::is_contained(dimsToRemove, i))
1490 newShape.push_back(tensorType.getDimSize(i));
1491 }
1492
1493 auto newType = RankedTensorType::get(newShape, tensorType.getElementType());
1495 tensor, newType);
1496}
1497
1498/// Drops specified dimensions from an AffineExpr and compresses remaining
1499/// dimension indices. Returns std::nullopt if the expression only references
1500/// the dropped dimensions.
1501static std::optional<AffineExpr>
1503 unsigned newNumDims, MLIRContext *ctx) {
1504 // Check if expr only references dimensions to be dropped.
1505 bool onlyReferencesDroppedDims = true;
1506 for (unsigned d = 0; d < newNumDims + dimsToDrop.size(); ++d) {
1507 if (expr.isFunctionOfDim(d) && !llvm::is_contained(dimsToDrop, d)) {
1508 onlyReferencesDroppedDims = false;
1509 break;
1510 }
1511 }
1512 if (onlyReferencesDroppedDims && llvm::any_of(dimsToDrop, [&](unsigned d) {
1513 return expr.isFunctionOfDim(d);
1514 }))
1515 return std::nullopt;
1516
1517 // Replace dimensions: compute new index for each old dimension.
1518 // Dropped dimensions get mapped to constant 0, others get compressed.
1519 SmallVector<AffineExpr> dimReplacements;
1520 unsigned newDimIdx = 0;
1521 for (unsigned d = 0; d < newNumDims + dimsToDrop.size(); ++d) {
1522 if (llvm::is_contained(dimsToDrop, d)) {
1523 dimReplacements.push_back(getAffineConstantExpr(0, ctx));
1524 } else {
1525 dimReplacements.push_back(getAffineDimExpr(newDimIdx++, ctx));
1526 }
1527 }
1528
1529 return expr.replaceDims(dimReplacements);
1530}
1531
1532FailureOr<LinalgOp>
1534 LinalgOp op) {
1535 auto maybeDims = inferConvolutionDims(op);
1536 if (failed(maybeDims))
1537 return failure();
1538
1539 // Currently supports only 2D convolutions.
1540 if (maybeDims->outputImage.size() != 2 || maybeDims->filterLoop.size() != 2)
1541 return failure();
1542
1543 if (op.hasPureBufferSemantics())
1544 return failure();
1545
1546 // Get loop domain indices for spatial dimensions.
1547 unsigned outSpatial0 = maybeDims->outputImage[0];
1548 unsigned outSpatial1 = maybeDims->outputImage[1];
1549 unsigned filterSpatial0 = maybeDims->filterLoop[0];
1550 unsigned filterSpatial1 = maybeDims->filterLoop[1];
1551
1552 // Get sizes from loop bounds.
1553 SmallVector<int64_t, 4> loopRanges = op.getStaticLoopRanges();
1554 int64_t outSize0 = loopRanges[outSpatial0];
1555 int64_t outSize1 = loopRanges[outSpatial1];
1556 int64_t filterSize0 = loopRanges[filterSpatial0];
1557 int64_t filterSize1 = loopRanges[filterSpatial1];
1558
1559 // Check if we can downscale by removing a spatial dimension.
1560 bool canRemoveSpatial0 = (filterSize0 == 1 && outSize0 == 1);
1561 bool canRemoveSpatial1 = (filterSize1 == 1 && outSize1 == 1);
1562 if (!canRemoveSpatial0 && !canRemoveSpatial1)
1563 return failure();
1564
1565 // Determine which loop dims to remove (output spatial + corresponding filter)
1566 // and sort for correct index compression when removing dimensions from affine
1567 // maps.
1568 SmallVector<unsigned> loopDimsToRemove;
1569 if (canRemoveSpatial0) {
1570 loopDimsToRemove.push_back(outSpatial0);
1571 loopDimsToRemove.push_back(filterSpatial0);
1572 } else {
1573 loopDimsToRemove.push_back(outSpatial1);
1574 loopDimsToRemove.push_back(filterSpatial1);
1575 }
1576 llvm::sort(loopDimsToRemove);
1577
1578 // Create new indexing maps with dimensions removed.
1579 SmallVector<AffineMap> newMaps;
1580 MLIRContext *ctx = op.getContext();
1581 unsigned numDims = op.getNumLoops();
1582 unsigned newNumDims = numDims - loopDimsToRemove.size();
1583 for (AffineMap map : op.getIndexingMapsArray()) {
1584 SmallVector<AffineExpr> newResults;
1585 for (AffineExpr expr : map.getResults()) {
1586 auto newExpr =
1587 dropDimsAndCompress(expr, loopDimsToRemove, newNumDims, ctx);
1588 if (newExpr)
1589 newResults.push_back(*newExpr);
1590 }
1591 newMaps.push_back(AffineMap::get(newNumDims, 0, newResults, ctx));
1592 }
1593
1594 // Create new iterator types.
1596 auto iterTypes = op.getIteratorTypesArray();
1597 for (unsigned idx = 0; idx < iterTypes.size(); ++idx) {
1598 if (!llvm::is_contained(loopDimsToRemove, idx))
1599 newIterTypes.push_back(iterTypes[idx]);
1600 }
1601
1602 // Rank-reduce operands using extract_slice.
1603 Location loc = op.getLoc();
1604 SmallVector<Value> newInputs;
1605 for (OpOperand *input : op.getDpsInputOperands()) {
1606 AffineMap map = op.getMatchingIndexingMap(input);
1607 SmallVector<unsigned> tensorDimsToRemove =
1608 getResultIndicesReferencingDims(map, loopDimsToRemove);
1609 Value reduced = createRankReducingExtractSlice(rewriter, loc, input->get(),
1610 tensorDimsToRemove);
1611 newInputs.push_back(reduced);
1612 }
1613
1614 OpOperand &output = *op.getDpsInitsMutable().begin();
1615 AffineMap outputMap = op.getMatchingIndexingMap(&output);
1616 SmallVector<unsigned> outputDimsToRemove =
1617 getResultIndicesReferencingDims(outputMap, loopDimsToRemove);
1618 Value newOutput = createRankReducingExtractSlice(rewriter, loc, output.get(),
1619 outputDimsToRemove);
1620
1621 // Create new linalg.generic with reduced dimensions.
1622 auto newOp =
1623 linalg::GenericOp::create(rewriter, loc, TypeRange{newOutput.getType()},
1624 newInputs, newOutput, newMaps, newIterTypes);
1625 rewriter.inlineRegionBefore(op->getRegion(0), newOp.getRegion(),
1626 newOp.getRegion().begin());
1627
1628 // Try to specialize the generic back to a named op only if the input was
1629 // already a specialized (named) op.
1630 LinalgOp resultOp = newOp;
1631 if (!isa<GenericOp>(op)) {
1632 FailureOr<LinalgOp> specializedOp = specializeGenericOp(rewriter, newOp);
1633 if (succeeded(specializedOp))
1634 resultOp = *specializedOp;
1635 }
1636
1637 // Insert result back into original shape.
1639 rewriter, loc, resultOp->getResult(0), output.get());
1640
1641 rewriter.replaceOp(op, result);
1642 return resultOp;
1643}
1644
1645namespace {
1646/// Pattern wrapper around `downscaleSizeOneWindowedConvolution`.
1647struct DownscaleSizeOneWindowedConvolution final
1648 : public OpInterfaceRewritePattern<LinalgOp> {
1649 DownscaleSizeOneWindowedConvolution(MLIRContext *context,
1650 PatternBenefit benefit = 1)
1651 : OpInterfaceRewritePattern<LinalgOp>(context, benefit) {}
1652
1653 LogicalResult matchAndRewrite(LinalgOp op,
1654 PatternRewriter &rewriter) const override {
1656 }
1657};
1658} // namespace
1659
1661 PatternBenefit benefit) {
1662 patterns.add<DownscaleSizeOneWindowedConvolution>(patterns.getContext(),
1663 benefit);
1664}
1665
1670
return success()
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
static RankedTensorType permuteShape(RankedTensorType tensorType, ArrayRef< int64_t > permutationVector)
Return a copy of tensorType after permutation by permutationVector.
static std::optional< AffineExpr > dropDimsAndCompress(AffineExpr expr, ArrayRef< unsigned > dimsToDrop, unsigned newNumDims, MLIRContext *ctx)
Drops specified dimensions from an AffineExpr and compresses remaining dimension indices.
static SmallVector< int64_t > getPackUnpackRankReducedPerm(ArrayRef< int64_t > shape, ArrayRef< int64_t > innerDimsPos, ArrayRef< int64_t > outerDimsPerm)
static Value createRankReducingExtractSlice(RewriterBase &rewriter, Location loc, Value tensor, ArrayRef< unsigned > dimsToRemove)
Helper to create a rank-reducing extract_slice that removes specific dimensions from a tensor.
static FailureOr< SmallVector< std::optional< int64_t > > > packLinalgMetadataOnce(SmallVectorImpl< AffineMap > &indexingMaps, SmallVectorImpl< utils::IteratorType > &iteratorTypes, int64_t dim)
Perform one step of packing of a LinalgOp's metadata along dim into the newDim at iteratorTypes....
static LinalgOp transposeOneLinalgOperandAndReplace(RewriterBase &rewriter, LinalgOp linalgOp, OpOperand &opOperand, ArrayRef< int64_t > permutation, Value transposedValue)
Return a new GenericOp obtained by transposing opOperand by the permutation vector:
static SmallVector< int64_t > getPackUnpackNormalizedPerm(int rank, ArrayRef< int64_t > perm)
static bool hasAtMostOneResultFunctionOfDim(AffineMap map, int64_t dim)
Return true if map has 0 or 1 result function of AffineDimExpr(dim).
static Value getPackOpSourceOrPaddedSource(OpBuilder &builder, linalg::PackOp packOp)
If padding value is set, returns a tensor.pad Op for the source tensor, with the output shape matchin...
static SmallVector< unsigned > getResultIndicesReferencingDims(AffineMap map, ArrayRef< unsigned > dims)
Returns the indices of affine map results that reference any of the given dimensions.
static std::string stringifyReassocIndices(ReassociationIndicesRef ri)
static std::optional< int64_t > getFirstResultIndexFunctionOf(AffineMap map, int64_t dim)
Return the index of the first result of map that is a function of AffineDimExpr(dim),...
Base type for affine expression.
Definition AffineExpr.h:68
bool isFunctionOfDim(unsigned position) const
Return true if the affine expression involves AffineDimExpr position.
AffineExpr replaceDims(ArrayRef< AffineExpr > dimReplacements) const
Dim-only version of replaceDimsAndSymbols.
AffineExpr ceilDiv(uint64_t v) const
A multi-dimensional affine map Affine map's are immutable like Type's, and they are uniqued.
Definition AffineMap.h:46
MLIRContext * getContext() const
static AffineMap get(MLIRContext *context)
Returns a zero result affine map with no dimensions or symbols: () -> ().
AffineMap shiftDims(unsigned shift, unsigned offset=0) const
Replace dims[offset ... numDims) by dims[offset + shift ... shift + numDims).
Definition AffineMap.h:267
AffineMap insertResult(AffineExpr expr, unsigned pos) const
Returns a new AffineMap with the same number of dims and symbols and an extra result inserted at pos.
Definition AffineMap.h:315
unsigned getNumDims() const
ArrayRef< AffineExpr > getResults() const
unsigned getNumResults() const
AffineExpr getResult(unsigned idx) const
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.
Attributes are known-constant values of operations.
Definition Attributes.h:25
This class is a general helper class for creating context-global objects like types,...
Definition Builders.h:51
IntegerAttr getIndexAttr(int64_t value)
Definition Builders.cpp:116
TypedAttr getZeroAttr(Type type)
Definition Builders.cpp:333
AffineExpr getAffineDimExpr(unsigned position)
Definition Builders.cpp:373
MLIRContext * getContext() const
Definition Builders.h:56
This is a utility class for mapping one set of IR entities to another.
Definition IRMapping.h:26
IRValueT get() const
Return the current value being used by this operand.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
MLIRContext is the top-level object for a collection of MLIR operations.
Definition MLIRContext.h:63
RAII guard to reset the insertion point of the builder when destroyed.
Definition Builders.h:351
This class helps build Operations.
Definition Builders.h:210
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
Definition Builders.h:401
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...
Definition Builders.h:528
This class represents a single result from folding an operation.
This class represents an operand of an operation.
Definition Value.h:254
unsigned getOperandNumber() const
Return which operand this is in the OpOperand list of the Operation.
Definition Value.cpp:226
This is a value defined by a result of an operation.
Definition Value.h:454
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
result_range getResults()
Definition Operation.h:440
This class represents the benefit of a pattern match in a unitless scheme that ranges from 0 (very li...
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
This is a builder type that keeps local references to arguments.
Builder & setShape(ArrayRef< int64_t > newShape)
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.
This class coordinates the application of a rewrite on a set of IR, providing a way for clients to tr...
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.
void moveOpBefore(Operation *op, Operation *existingOp)
Unlink this operation from its current block and insert it right before existingOp which may be in th...
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 inlineRegionBefore(Region &region, Region &parent, Region::iterator before)
Move the blocks that belong to "region" before the given position in another region "parent".
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
This class provides an abstraction over the various different ranges of value types.
Definition TypeRange.h:40
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
type_range getTypes() const
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Type getType() const
Return the type of this value.
Definition Value.h:105
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
Definition ArithOps.cpp:398
Operation * getOwner() const
Return the owner of this operand.
Definition UseDefLists.h:38
OpFoldResult makeComposedFoldedAffineApply(OpBuilder &b, Location loc, AffineMap map, ArrayRef< OpFoldResult > operands, bool composeAffineMin=false)
Constructs an AffineApplyOp that applies map to operands after composing the map with the maps of any...
SmallVector< int64_t > getUnPackInverseSrcPerm(linalg::UnPackOp, PackingMetadata &metadata)
Compute inverse permutation for the source tensor (i.e.
FailureOr< PackTransposeResult > packTranspose(RewriterBase &rewriter, linalg::PackOp packOp, linalg::LinalgOp linalgOp, linalg::UnPackOp maybeUnPackOp, ArrayRef< int64_t > outerPerm, ArrayRef< int64_t > innerPerm)
Transpose a single PackOp -> LinalgOp -> UnPackOp chain and return the transposed PackOp -> LinalgOp ...
FailureOr< LowerUnPackOpResult > lowerUnPack(RewriterBase &rewriter, linalg::UnPackOp unPackOp, bool lowerUnpadLikeWithExtractSlice=true)
Rewrite pack as empty + transpose + reshape + extract_slice + copy.
void peelLoops(RewriterBase &rewriter, ArrayRef< scf::ForOp > loops)
Peel 'loops' and applies affine_min/max bounds simplification on the fly where relevant.
FailureOr< ConvolutionDimensions > inferConvolutionDims(LinalgOp linalgOp)
Find at least 1 parallel (output_image) and reduction (filter_loop) dimension candidates that form a ...
FailureOr< LinalgOp > generalizeNamedOp(RewriterBase &rewriter, LinalgOp linalgOp, bool emitCategoryOps=false)
Create a GenericOp or CategoryOp from the given named operation linalgOp and replace the given linalg...
FailureOr< LinalgOp > specializeGenericOp(RewriterBase &rewriter, GenericOp genericOp, bool emitCategoryOps=false)
Replace the given GenericOp with a namedOp or categoryOp.
void populateDecomposeConvolutionPatterns(RewritePatternSet &patterns, PatternBenefit benefit=1)
Linalg decompose convolutions patterns.
LogicalResult vectorizeCopy(RewriterBase &builder, memref::CopyOp copyOp)
Emit a suitable vector form for a Copy op with fully static shape.
FailureOr< GenericOp > interchangeGenericOp(RewriterBase &rewriter, GenericOp genericOp, ArrayRef< unsigned > interchangeVector)
Interchange the iterator_types and iterator_maps dimensions and adapts the index accesses of op.
SmallVector< int64_t > getPackInverseDestPerm(linalg::PackOp packOp, PackingMetadata &metadata)
Compute inverse permutation for the destination tensor (i.e.
void populateDecomposePackUnpackPatterns(RewritePatternSet &patterns)
Populates patterns to decompose linalg.pack and linalg.unpack Ops into e.g.
FailureOr< ContractionDimensions > inferContractionDims(LinalgOp linalgOp)
Find at least 2 parallel (m and n) and 1 reduction (k) dimension candidates that form a matmul subcom...
FailureOr< PackResult > packMatmulGreedily(RewriterBase &rewriter, LinalgOp linalgOp, ArrayRef< OpFoldResult > mnkPackedSizes, ArrayRef< int64_t > mnkPaddedSizesNextMultipleOf, ArrayRef< int64_t > mnkOrder)
Pack a LinalgOp by greedily inferring matmul dimensions (m, n, k) where m and n are proper parallel d...
FailureOr< PackResult > pack(RewriterBase &rewriter, linalg::LinalgOp linalgOp, ArrayRef< OpFoldResult > packedSizes)
Implement packing of a single LinalgOp by packedSizes.
SmallVector< Value > peelLoop(RewriterBase &rewriter, Operation *op)
Try to peel and canonicalize loop op and return the new result.
void populateDecomposePadPatterns(RewritePatternSet &patterns)
Populates patterns to decompose tensor.pad into e.g.
FailureOr< LinalgOp > downscaleSizeOneWindowedConvolution(RewriterBase &rewriter, LinalgOp op)
Rewrite convolution/pooling/depthwise ops with size-1 window dimensions into lower-dimensional ops.
FailureOr< LowerPackResult > lowerPack(RewriterBase &rewriter, linalg::PackOp packOp, bool lowerPadLikeWithInsertSlice=true)
Rewrite pack as pad + reshape + transpose.
LogicalResult peelForLoopAndSimplifyBounds(RewriterBase &rewriter, ForOp forOp, scf::ForOp &partialIteration)
Rewrite a for loop with bounds/step that potentially do not divide evenly into a for loop where the s...
FailureOr< TilingResult > bubbleUpPadSlice(OpBuilder &b, tensor::PadOp padOp, ArrayRef< OpFoldResult > offsets, ArrayRef< OpFoldResult > sizes, bool generateZeroSliceGuard=true)
Bubbles up a slice of this pad by taking the slice first and then performing the padding.
PadOp createPadHighOp(RankedTensorType resType, Value source, Value pad, bool nofold, Location loc, OpBuilder &builder, ValueRange dynOutDims={})
Definition Utils.cpp:23
Value createCanonicalRankReducingInsertSliceOp(OpBuilder &b, Location loc, Value tensor, Value dest)
Create a rank-reducing InsertSliceOp @[0 .
Value createCanonicalRankReducingExtractSliceOp(OpBuilder &b, Location loc, Value tensor, RankedTensorType targetType)
Create a rank-reducing ExtractSliceOp @[0 .
OpFoldResult getMixedSize(OpBuilder &builder, Location loc, Value value, int64_t dim)
Return the dimension of the given tensor value.
Definition TensorOps.cpp:82
SmallVector< OpFoldResult > getMixedSizes(OpBuilder &builder, Location loc, Value value)
Return the dimensions of the given tensor value.
Definition TensorOps.cpp:91
Include the generated interface declarations.
SliceVerificationResult
Enum that captures information related to verifier error conditions on slice insert/extract type of o...
ArrayRef< int64_t > ReassociationIndicesRef
std::optional< int64_t > getConstantIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
void bindDims(MLIRContext *ctx, AffineExprTy &...exprs)
Bind a list of AffineExpr references to DimExpr at positions: [0 .
Definition AffineExpr.h:311
SmallVector< int64_t > computePermutationVector(int64_t permSize, ArrayRef< int64_t > positions, ArrayRef< int64_t > desiredPositions)
Return a permutation vector of size permSize that would result in moving positions into desiredPositi...
Type getElementTypeOrSelf(Type type)
Return the element type or return the type itself.
void bindSymbols(MLIRContext *ctx, AffineExprTy &...exprs)
Bind a list of AffineExpr references to SymbolExpr at positions: [0 .
Definition AffineExpr.h:325
SmallVector< Loops, 8 > tile(ArrayRef< scf::ForOp > forOps, ArrayRef< Value > sizes, ArrayRef< scf::ForOp > targets)
Performs tiling fo imperfectly nested loops (with interchange) by strip-mining the forOps by sizes an...
Definition Utils.cpp:1380
AffineExpr getAffineConstantExpr(int64_t constant, MLIRContext *context)
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
Definition LLVM.h:120
std::pair< int64_t, OpFoldResult > getSimplifiedOfrAndStaticSizePair(OpFoldResult ofr, Builder &b)
Given OpFoldResult representing dim size value (*), generates a pair of sizes:
void applyPermutationToVector(SmallVector< T, N > &inVec, ArrayRef< int64_t > permutation)
Apply the permutation defined by permutation to inVec.
AffineExpr getAffineDimExpr(unsigned position, MLIRContext *context)
These free functions allow clients of the API to not use classes in detail.
SliceVerificationResult isRankReducedType(ShapedType originalType, ShapedType candidateReducedType)
Check if originalType can be rank reduced to candidateReducedType type by dropping some dimensions wi...
bool isPermutationVector(ArrayRef< int64_t > interchange)
Method to check if an interchange vector is a permutation.
SmallVector< int64_t > invertPermutationVector(ArrayRef< int64_t > permutation)
Helper method to apply to inverse a permutation.
OpInterfaceRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting a...
Represents a range (offset, size, and stride) where each element of the triple may be dynamic or stat...
OpFoldResult size
LogicalResult matchAndRewrite(memref::CopyOp copyOp, PatternRewriter &rewriter) const override
Rewrites a linalg::PackOp into a sequence of:
LogicalResult matchAndRewrite(linalg::PackOp packOp, PatternRewriter &rewriter) const override
Rewrites a linalg::UnPackOp into a sequence of:
LogicalResult matchAndRewrite(linalg::UnPackOp unpackOp, PatternRewriter &rewriter) const override
Rewrite a tensor::PadOp into a sequence of EmptyOp, FillOp and InsertSliceOp.
LogicalResult matchAndRewrite(tensor::PadOp padOp, PatternRewriter &rewriter) const override
Value createFillOrGenerateOp(RewriterBase &rewriter, tensor::PadOp padOp, Value dest, const SmallVector< Value > &dynSizes) const
Filling dest using FillOp constant padding value if possible.
LogicalResult matchAndRewrite(tensor::ExtractSliceOp sliceOp, PatternRewriter &rewriter) const override
LinalgTilingOptions & setTileSizes(const SmallVector< Value, 4 > &ts)
Set the tileSizeComputationFunction to return the values ts.
Definition Transforms.h:204
TileSizeComputationFunction tileSizeComputationFunction
Computation function that returns the tile sizes for each operation.
Definition Transforms.h:194
Struct to hold the result of a pack call.
Struct to hold the result of a packTranspose call.