MLIR 24.0.0git
ReshapeOpsUtils.h
Go to the documentation of this file.
1//===- ReshapeOpsUtils.h - Utilities used by reshape ops --*- C++ -*------===//
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 header file defines utilities and common canonicalization patterns for
10// reshape operations.
11//
12//===----------------------------------------------------------------------===//
13
14#ifndef MLIR_DIALECT_UTILS_RESHAPEOPSUTILS_H
15#define MLIR_DIALECT_UTILS_RESHAPEOPSUTILS_H
16
21#include "mlir/Support/LLVM.h"
22#include "llvm/ADT/STLExtras.h"
23#include "llvm/ADT/StringRef.h"
24#include <optional>
25
26namespace mlir {
27
31
32/// Attribute name for the ArrayAttr which encodes reassociation indices.
33constexpr StringRef getReassociationAttrName() { return "reassociation"; }
34
35/// Compose reassociation maps that are used in pair of reshape ops where one
36/// is a producer and other is the consumer. Only valid to use this method when
37/// both the producer and consumer are collapsing dimensions or both are
38/// expanding dimensions.
39///
40/// For example,
41/// producerReassociation = [[0, 1], [2], [3, 4]]
42/// consumerReassociation = [[0, 1], [2]]
43///
44/// is folded into
45///
46/// result = [[0, 1, 2], [3, 4]].
47std::optional<SmallVector<ReassociationIndices>> composeReassociationIndices(
48 ArrayRef<ReassociationIndices> producerReassociations,
49 ArrayRef<ReassociationIndices> consumerReassociations,
50 MLIRContext *context);
51
52/// Convert reassociation indices to affine expressions.
53SmallVector<SmallVector<AffineExpr, 2>, 2> convertReassociationIndicesToExprs(
54 MLIRContext *context, ArrayRef<ReassociationIndices> reassociationIndices);
55
56/// Constructs affine maps out of Array<Array<AffineExpr>>.
57SmallVector<AffineMap, 4>
58getSymbolLessAffineMaps(ArrayRef<ReassociationExprs> reassociation);
59
60/// Wraps a list of reassociations in an ArrayAttr.
63 ArrayRef<ReassociationIndices> reassociation);
64
65/// Convert Array<Array<AffineExpr>> to Array<Array<int64_t>>.
66SmallVector<ReassociationIndices, 2> convertReassociationMapsToIndices(
67 ArrayRef<ReassociationExprs> reassociationExprs);
68
69/// Return the reassociations maps to use to reshape given the source type and
70/// the target type when possible. Return std::nullopt when this computation
71/// failed.
72std::optional<SmallVector<ReassociationIndices>>
73getReassociationIndicesForReshape(ShapedType sourceType, ShapedType targetType);
74
75/// Returns the reassociation maps to collapse `sourceShape` to `targetShape` if
76/// possible.
77std::optional<SmallVector<ReassociationIndices>>
78getReassociationIndicesForCollapse(ArrayRef<int64_t> sourceShape,
79 ArrayRef<int64_t> targetShape);
80
81/// Return true if the reassociation specification is valid, false otherwise.
82/// When false, the `invalidIndex` integer pointer is optionally filled with the
83/// index of the offending reassociation map.
84bool isReassociationValid(ArrayRef<AffineMap> reassociation,
85 int *invalidIndex = nullptr);
86
87/// Verify that none of the reassociation groups is empty.
88template <typename ReshapeOpTy>
89LogicalResult verifyReassociationIndicesNotEmpty(ReshapeOpTy op) {
90 if (llvm::any_of(
91 op.getReassociationIndices(),
92 [](const ReassociationIndices &group) { return group.empty(); })) {
93 return op.emitOpError("reassociation indices must not be empty");
94 }
95 return success();
96}
97
98template <typename ReshapeOpTy, typename InverseReshapeOpTy>
99OpFoldResult foldReshapeOp(ReshapeOpTy reshapeOp,
100 ArrayRef<Attribute> operands) {
101 // Fold identity reshape.
102 if (reshapeOp.getSrcType() == reshapeOp.getType())
103 return reshapeOp.getSrc();
104
105 // Reshape of a constant can be replaced with a new constant, but only when
106 // the result type has a static shape. DenseElementsAttr::reshape requires
107 // a static shape to preserve the element count invariant.
108 if (auto elements = dyn_cast_or_null<DenseElementsAttr>(operands.front())) {
109 auto resultType = cast<ShapedType>(reshapeOp.getResult().getType());
110 if (resultType.hasStaticShape())
111 return elements.reshape(resultType);
112 }
113
114 // Fold if the producer reshape source has the same shape with at most 1
115 // dynamic dimension.
116 auto reshapeSrcOp =
117 reshapeOp.getSrc().template getDefiningOp<InverseReshapeOpTy>();
118 if (!reshapeSrcOp)
119 return nullptr;
120 auto srcType = reshapeSrcOp.getSrcType();
121 auto resultType = reshapeOp.getResultType();
122 if (srcType != resultType)
123 return nullptr;
124
125 if (llvm::count_if(srcType.getShape(), ShapedType::isDynamic) < 2) {
126 return reshapeSrcOp.getSrc();
127 }
128
129 // Fold producer-consumer reshape ops when they are perfect inverses of each
130 // other:
131 // 1) Reassociation indices are equivalent.
132 // 2) Boundary types are equivalent.
133 // 3) No reassociations have more than 1 dynamic dimension, and reassociated
134 // shapes are equal for each reassociation.
135 auto reassociations = reshapeOp.getReassociationIndices();
136 if (reassociations != reshapeSrcOp.getReassociationIndices())
137 return nullptr;
138 // If the reshapes are expanding and then collapsing, the ops can be folded
139 // despite multiple dynamic dimensions.
140 if (srcType.getRank() < reshapeSrcOp.getResultType().getRank())
141 return reshapeSrcOp.getSrc();
142 if (llvm::all_of(reassociations, [&](auto reInd) {
143 ArrayRef<int64_t> srcSlice =
144 srcType.getShape().slice(reInd.front(), reInd.size());
145 return llvm::count_if(srcSlice, ShapedType::isDynamic) < 2;
146 })) {
147 return reshapeSrcOp.getSrc();
148 }
149 return nullptr;
150}
151
152/// Common verifier for reshape-like types. Fills `expandedType` and
153///`collapsedType` with the proper `src` or `result` type.
154template <typename Op, typename T>
155LogicalResult verifyReshapeLikeTypes(Op op, T expandedType, T collapsedType,
156 bool isExpansion) {
157
158 unsigned expandedRank = expandedType.getRank();
159 unsigned collapsedRank = collapsedType.getRank();
160 if (expandedRank < collapsedRank)
161 return op.emitOpError("expected the expanded type, ")
162 << expandedType << " to have a higher (or same) rank "
163 << "than the collapsed type, " << collapsedType << '.';
164
165 if (collapsedRank != op.getReassociation().size())
166 return op.emitOpError("expected collapsed rank (")
167 << collapsedRank << ") to equal the number of reassociation maps ("
168 << op.getReassociation().size() << ").";
169
170 auto maps = op.getReassociationMaps();
171 for (auto it : llvm::enumerate(maps))
172 if (it.value().getNumDims() != expandedRank)
173 return op.emitOpError("expected reassociation map #")
174 << it.index() << " to have size equal to the expanded rank ("
175 << expandedRank << "), but it is " << it.value().getNumDims()
176 << '.';
177
178 int invalidIdx = 0;
179 if (!isReassociationValid(maps, &invalidIdx))
180 return op.emitOpError("expected reassociation map #")
181 << invalidIdx << " to be valid and contiguous.";
182
183 return reshapeLikeShapesAreCompatible(
184 [&](const Twine &msg) { return op->emitOpError(msg); },
185 collapsedType.getShape(), expandedType.getShape(),
186 op.getReassociationIndices(), isExpansion);
187}
188
189/// Verify that shapes of the reshaped types using following rule:
190/// if a dimension in the collapsed type is static, then the corresponding
191/// dimensions in the expanded shape should be
192/// a) static
193/// b) the product should be same as the collaped shape.
194LogicalResult reshapeLikeShapesAreCompatible(
195 function_ref<LogicalResult(const Twine &)> emitError,
196 ArrayRef<int64_t> collapsedShape, ArrayRef<int64_t> expandedShape,
197 ArrayRef<ReassociationIndices> reassociationMaps, bool isExpandingReshape);
198
199/// Returns true iff the type is a MemRefType and has a non-identity layout.
200bool hasNonIdentityLayout(Type type);
201
202enum class ReshapeOpKind { kExpand, kCollapse };
203
204/// Pattern to collapse producer/consumer reshape ops that are both collapsing
205/// dimensions or are both expanding dimensions.
206template <typename ReshapeOpTy, ReshapeOpKind opKind>
207struct ComposeReassociativeReshapeOps : public OpRewritePattern<ReshapeOpTy> {
208 using OpRewritePattern<ReshapeOpTy>::OpRewritePattern;
209 LogicalResult matchAndRewrite(ReshapeOpTy reshapeOp,
210 PatternRewriter &rewriter) const override {
211 auto srcReshapeOp =
212 reshapeOp.getSrc().template getDefiningOp<ReshapeOpTy>();
213 if (!srcReshapeOp)
214 return failure();
215
216 ShapedType resultType = reshapeOp.getResultType();
217
218 if (hasNonIdentityLayout(srcReshapeOp.getSrc().getType()) ||
219 hasNonIdentityLayout(reshapeOp.getSrc().getType()) ||
220 hasNonIdentityLayout(reshapeOp.getResult().getType()))
221 return failure();
222
223 std::optional<SmallVector<ReassociationIndices>> reassociationIndices =
224 composeReassociationIndices(srcReshapeOp.getReassociationIndices(),
225 reshapeOp.getReassociationIndices(),
226 rewriter.getContext());
227 if (!reassociationIndices)
228 return failure();
229
230 if constexpr (opKind == ReshapeOpKind::kExpand) {
231 SmallVector<OpFoldResult> outputShape(
232 getMixedValues(reshapeOp.getStaticOutputShape(),
233 reshapeOp.getOutputShape(), rewriter));
234 rewriter.replaceOpWithNewOp<ReshapeOpTy>(
235 reshapeOp, resultType, srcReshapeOp.getSrc(), *reassociationIndices,
236 outputShape);
237 } else {
238 rewriter.replaceOpWithNewOp<ReshapeOpTy>(
239 reshapeOp, resultType, srcReshapeOp.getSrc(), *reassociationIndices);
240 }
241 return success();
242 }
243};
244
245/// Pattern to compose
246/// `collapse_shape(expand_shape(%src, reassociation_1), reassociation_2)`.
247/// In that case both `srcType` and `resultType` can be expressed as a function
248/// of `intermediateType`.
249/// In order to demonstrate the approach, let's assume that `rank(srcType) >
250/// `rank(resultType)`, i.e. the resulting operation should be `collapse_shape`.
251/// In that case, we can iterate over every set of indices in `reassociation_2`
252/// and try to find ids of sets of indices in `reassociation_1` that cover it
253/// completely.
254///
255/// Example:
256///
257/// %0 = tensor.expand_shape %arg [[0], [1], [2, 3]]
258/// : tensor<?x?x?xi64> into tensor<?x?x?x1xi64>
259/// %1 = tensor.collapse_shape %0 [[0, 1], [2, 3]]
260/// : tensor<?x?x?x1xi64> into tensor<?x?xi64>
261///
262/// can be canonicalized into
263///
264/// %0 = tensor.collapse_shape %arg [[0, 1], [2]]
265/// : tensor<?x?x?xi64> into tensor<?x?xi64>
266///
267/// because [0] and [1] from `expand_shape` reassociation cover completely
268/// `[0, 1]` from `collapse_shape`. If it is impossible to find such union of
269/// indices, then we fail.
270//
271/// When `rank(srcType) < rank(resultType)`, then we just swap `reassociation_1`
272/// `reassociation_2` and produce `expand_shape`.
273template <typename CollapseOpTy, typename ExpandOpTy, typename CastOpTy,
274 typename DimOpTy, typename TensorTy>
275struct ComposeCollapseOfExpandOp : public OpRewritePattern<CollapseOpTy> {
276 using OpRewritePattern<CollapseOpTy>::OpRewritePattern;
277 LogicalResult matchAndRewrite(CollapseOpTy collapseOp,
278 PatternRewriter &rewriter) const override {
279 auto expandOp = collapseOp.getSrc().template getDefiningOp<ExpandOpTy>();
280 if (!expandOp)
281 return failure();
282
283 ShapedType srcType = expandOp.getSrcType();
284 ShapedType resultType = collapseOp.getResultType();
285
286 if (hasNonIdentityLayout(collapseOp.getSrc().getType()) ||
287 hasNonIdentityLayout(expandOp.getSrc().getType()) ||
288 hasNonIdentityLayout(expandOp.getResult().getType()))
289 return failure();
290
291 int64_t srcRank = srcType.getRank();
292 int64_t resultRank = resultType.getRank();
293 if (srcType == resultType)
294 return failure();
295
296 SmallVector<ReassociationIndices, 4> higherRankReassociation,
297 lowerRankReassociation;
298
299 if (srcRank > resultRank) {
300 higherRankReassociation = expandOp.getReassociationIndices();
301 lowerRankReassociation = collapseOp.getReassociationIndices();
302 } else {
303 higherRankReassociation = collapseOp.getReassociationIndices();
304 lowerRankReassociation = expandOp.getReassociationIndices();
305 }
306
307 size_t higherRankIndicesID = 0;
308 SmallVector<ReassociationIndices, 4> composedReassociation;
309 for (const auto &lowerRankIndices : lowerRankReassociation) {
310 ReassociationIndices composedIndices;
311 while (higherRankIndicesID < higherRankReassociation.size()) {
312 auto rightmostIndex =
313 higherRankReassociation[higherRankIndicesID].back();
314 if (rightmostIndex > lowerRankIndices.back())
315 return failure();
316 composedIndices.push_back(higherRankIndicesID++);
317 if (rightmostIndex == lowerRankIndices.back())
318 break;
319 }
320 composedReassociation.push_back(composedIndices);
321 }
322 if (srcRank > resultRank) {
323 rewriter.replaceOpWithNewOp<CollapseOpTy>(
324 collapseOp, resultType, expandOp.getSrc(), composedReassociation);
325 } else if (srcRank < resultRank) {
326 // Compute the dynamic output shape for the new expand_shape op.
327 Location loc = collapseOp.getLoc();
328 SmallVector<OpFoldResult> origOutputShape =
329 expandOp.getMixedOutputShape();
330 SmallVector<OpFoldResult> newOutputShape;
331 for (const ReassociationIndices &indices :
332 collapseOp.getReassociationIndices()) {
333 int64_t numStaticElems = 1;
334 SmallVector<Value> dynamicSizes;
335 for (int64_t idx : indices) {
336 OpFoldResult size = origOutputShape[idx];
337 if (std::optional<int64_t> maybeCst = getConstantIntValue(size)) {
338 numStaticElems *= maybeCst.value();
339 continue;
340 }
341 dynamicSizes.push_back(cast<Value>(size));
342 }
343 if (dynamicSizes.empty()) {
344 newOutputShape.push_back(rewriter.getIndexAttr(numStaticElems));
345 continue;
346 }
347
348 // There is at least one dynamic size, so we can initialize `result` to
349 // the first dynamic size.
350 Value result = dynamicSizes[0];
351 for (Value v : llvm::drop_begin(dynamicSizes))
352 result = arith::MulIOp::create(rewriter, loc, result, v,
353 arith::IntegerOverflowFlags::nsw);
354 if (numStaticElems != 1) {
355 result = arith::MulIOp::create(
356 rewriter, loc, result,
357 arith::ConstantIndexOp::create(rewriter, loc, numStaticElems),
358 arith::IntegerOverflowFlags::nsw);
359 }
360 newOutputShape.push_back(result);
361 }
362 rewriter.replaceOpWithNewOp<ExpandOpTy>(
363 collapseOp, resultType, expandOp.getSrc(), composedReassociation,
364 newOutputShape);
365 } else {
366 // Collapses/expansions that do not change the rank are not allowed. Use
367 // a cast instead.
368 assert(llvm::equal(srcType.getShape(), resultType.getShape()) &&
369 "expected same shape");
370 rewriter.replaceOpWithNewOp<CastOpTy>(collapseOp, resultType,
371 expandOp.getSrc());
372 }
373 return success();
374 }
375};
376
377template <typename ExpandOpTy, typename CollapseOpTy, typename CastOpTy>
378struct ComposeExpandOfCollapseOp : public OpRewritePattern<ExpandOpTy> {
379 using OpRewritePattern<ExpandOpTy>::OpRewritePattern;
380 LogicalResult matchAndRewrite(ExpandOpTy expandOp,
381 PatternRewriter &rewriter) const override {
382 auto collapseOp = expandOp.getSrc().template getDefiningOp<CollapseOpTy>();
383 if (!collapseOp)
384 return failure();
385
386 ShapedType srcType = collapseOp.getSrcType();
387 ShapedType resultType = expandOp.getResultType();
388
389 if (hasNonIdentityLayout(expandOp.getSrc().getType()) ||
390 hasNonIdentityLayout(collapseOp.getSrc().getType()) ||
391 hasNonIdentityLayout(collapseOp.getResult().getType())) {
392 if (srcType.hasStaticShape() &&
393 CastOpTy::areCastCompatible(srcType, resultType)) {
394 rewriter.replaceOpWithNewOp<CastOpTy>(expandOp, resultType,
395 collapseOp.getSrc());
396 return success();
397 }
398 return failure();
399 }
400
401 int64_t srcRank = srcType.getRank();
402 int64_t resultRank = resultType.getRank();
403 if (srcRank == resultRank)
404 return failure();
405
406 auto srcReassociation = collapseOp.getReassociationIndices();
407 auto resultReassociation = expandOp.getReassociationIndices();
408 if (srcRank > resultRank) {
409 auto composedReassociation = findCollapsingReassociation(
410 srcReassociation, resultReassociation, srcType.getShape(),
411 resultType.getShape());
412 if (!composedReassociation)
413 return failure();
414
415 rewriter.replaceOpWithNewOp<CollapseOpTy>(
416 expandOp, resultType, collapseOp.getSrc(), *composedReassociation);
417 return success();
418 }
419 auto composedReassociation =
420 findCollapsingReassociation(resultReassociation, srcReassociation,
421 resultType.getShape(), srcType.getShape());
422 if (!composedReassociation)
423 return failure();
424
426 expandOp.getStaticOutputShape(), expandOp.getOutputShape(), rewriter));
427 rewriter.replaceOpWithNewOp<ExpandOpTy>(
428 expandOp, resultType, collapseOp.getSrc(), *composedReassociation,
429 outputShape);
430 return success();
431 }
432
433private:
434 // Attempts to find a way to collapse `srcShape` to `resultShape` by
435 // collapsing subshapes defined by the reassociation indices.
436 std::optional<SmallVector<ReassociationIndices>> findCollapsingReassociation(
437 ArrayRef<ReassociationIndices> srcReassociation,
438 ArrayRef<ReassociationIndices> resultReassociation,
439 ArrayRef<int64_t> srcShape, ArrayRef<int64_t> resultShape) const {
440 SmallVector<ReassociationIndices, 4> composedReassociation;
441
442 if (srcReassociation.empty())
443 return {getReassociationIndicesForCollapse(srcShape, resultShape)};
444
445 for (auto item : llvm::zip(srcReassociation, resultReassociation)) {
446 auto &srcIndices = std::get<0>(item);
447 auto &resultIndices = std::get<1>(item);
448 auto srcSubShape = srcShape.slice(srcIndices.front(), srcIndices.size());
449 auto resultSubShape =
450 resultShape.slice(resultIndices.front(), resultIndices.size());
451
452 if (llvm::count_if(srcSubShape, ShapedType::isDynamic) >= 2 &&
453 llvm::count_if(resultSubShape, ShapedType::isDynamic) >= 2)
454 return std::nullopt;
455
456 if (srcSubShape.size() == resultSubShape.size()) {
457 if (srcSubShape != resultSubShape)
458 return std::nullopt;
459
460 for (auto index : llvm::seq<int64_t>(0, srcSubShape.size())) {
461 composedReassociation.emplace_back(1, srcIndices.front() + index);
462 }
463 continue;
464 }
465
466 // Find reassociation to collapse `srcSubShape` into `resultSubShape`.
467 auto subShapeReassociation =
468 getReassociationIndicesForCollapse(srcSubShape, resultSubShape);
469 if (!subShapeReassociation)
470 return std::nullopt;
471
472 // Remap the subshape indices back to the original srcShape.
473 for (auto &subshapeIndices : *subShapeReassociation) {
474 ReassociationIndices shapeIndices;
475 for (int64_t index : subshapeIndices)
476 shapeIndices.push_back(srcIndices.front() + index);
477 composedReassociation.push_back(shapeIndices);
478 }
479 }
480 return {std::move(composedReassociation)};
481 }
482};
483
484/// The input parameters `offsets`, `sizes`, `strides` specify a rectangular
485/// non rank-reducing slice of the collapse_shape output. Try to find which
486/// dimensions have been sliced and which dimensions are not sliced (offset = 0,
487/// size = dim, size = 1). Note that this conservative as it cannot detect if a
488/// dynamic size corresponds to the full tensor dimension or not.
489llvm::SmallBitVector getSlicedDimensions(ArrayRef<OpFoldResult> sliceInputShape,
490 ArrayRef<Range> sliceParams);
491
492/// Determine which dimensions are linearized by a `tensor.collapse_shape` op by
493/// inspecting its reassociation indices.
494llvm::SmallBitVector
496
497/// Given the parameters for both operations in a `CollapseShape->ExtractSlice`
498/// chain and reified source and result shapes of the CollapseShapeOp, this
499/// class provides two functions that assist with directly forming the result
500/// of the extract slice by "tiling the CollapseShapeOp by 1".
501//// Example:
502// clang-format off
503/// ```
504/// %0 = linalg.generic ... -> tensor<3x7x11x10xf32>
505/// %1 = tensor.collapse_shape %0 [[0, 1, 2], [3]] : ... to tensor<341x10xf32>
506/// %2 = tensor.extract_slice %1 [13, 0] [10, 10] [2, 1] : .... tensor<10x10xf32>
507/// ```
508/// This class helps build the below IR to replace %2:
509/// ```
510/// %dest = tensor.empty() : tensor<10x10xf32>
511/// %2 = scf.for %iv = %c0 to %c10 step %c1 iter_args(%arg0) -> tensor<10x10xf32> {
512/// %linear_index = affine.apply affine_map<(d0)[]->(d0*2 + 11)>(%iv)
513/// %3:3 = arith.delinearize_index %iv into (3, 7, 11)
514///
515/// // This function takes %3 (multiIndices) and the parameters for the slice below.
516/// %4 = tensor.extract_slice %0 [%3#0, %3#1, %3#2, 0] [1, 1, 1, 10] [1, 1, 1, 1] :
517/// tensor<3x7x11x10xf32> to tensor<1x1x1x10xf32>
518///
519/// %5 = tensor.collapse_shape %4 [[0, 1, 2], [3]] :
520/// tensor<1x1x1x10xf32> into tensor<1x10xf32>
521/// %6 = tensor.insert_slice %5 into %arg0 [%iv, 0] [1, 10] [1, 1] :
522/// tensor<1x10xf32> into tensor<10x10xf32>
523/// scf.yield %6 : tensor<10x10xf32>
524/// }
525/// ```
526// clang-format on
527class SliceFromCollapseHelper {
528public:
529 SliceFromCollapseHelper(ArrayRef<ReassociationIndices> reassociationIndices,
530 ArrayRef<OpFoldResult> collapseShapeInputShape,
531 ArrayRef<OpFoldResult> collapseShapeOutputShape,
532 ArrayRef<Range> extractSliceParams)
533 : reassociationIndices(reassociationIndices),
534 collapseShapeInputShape(collapseShapeInputShape),
535 collapseShapeOutputShape(collapseShapeOutputShape),
536 sliceParams(extractSliceParams),
537 linearizedDimensions(getLinearizedDimensions(reassociationIndices)),
538 slicedDimensions(getSlicedDimensions(collapseShapeOutputShape,
539 extractSliceParams)) {}
540
541 /// This function takes multi-indices and maps them to ExtractSlice parameters
542 /// in the index space of the CollapseShape's source tensor. This function's
543 /// signature can be described by `(D_0, D_1,.. D_{n-1}) -> (offsets, sizes,
544 /// strides)` where `n` the number of "tiled dimensions", which are the
545 /// dimensions of the output that are linearized by the collapse shape op and
546 /// are also sliced. Each `D_i` is a tuple that must represent a valid
547 /// multi-index for the `i-th` tiled dimension. In the example above, there is
548 /// only one tiled dimension (D_0) and `arith.delinearize_index` produces the
549 /// multi-index (%3) that would be passed to this function to generate the
550 /// parameters for the `tensor.extract_slice` op (%4).
551 SmallVector<Range> getExtractSliceParams(MLIRContext *ctx,
552 ArrayRef<ValueRange> multiIndices);
553
554 /// This function takes indices in the index space of the "tiled dimensions"
555 /// described above and returns a set of Range variables that describe how the
556 /// slice should be inserted into the destination. In the example above, `%iv`
557 /// would be passed to this function to generate the parameters for the
558 /// `tensor.insert_slice` op producing %6.
559 SmallVector<Range> getInsertSliceParams(MLIRContext *ctx,
560 ValueRange tileIndices);
561
562private:
563 SmallVector<ReassociationIndices> reassociationIndices;
564 SmallVector<OpFoldResult> collapseShapeInputShape;
565 SmallVector<OpFoldResult> collapseShapeOutputShape;
566 SmallVector<Range> sliceParams;
567 llvm::SmallBitVector linearizedDimensions;
568 llvm::SmallBitVector slicedDimensions;
569};
570
571/// Parameters required to simplify a collapsing reshape op with a rank-reducing
572/// slice operation. See `getSimplifyCollapseShapeWithRankReducingSliceInfo`.
573struct CollapseShapeRankReducingSliceSimplificationInfo {
574 /// The shape of the output of the rank-reducing slice.
575 RankedTensorType sliceResultType;
576 /// The reassociation indices for the new collapse shape op, if required. If
577 /// `std::nullopt`, the slice should replace the collapse shape op.
578 std::optional<SmallVector<ReassociationIndices>> newReassociationIndices;
579};
580
581/// A collapsing reshape operation can sometimes be simplified or eliminated by
582/// inserting a single rank-reducing slice operation between it and the source
583/// tensor. The slice op will either take the place of the source, allowing for
584/// a new, simpler reshape op to replace the original, or the reshape op will be
585/// completely replaced by the slice result.
586///
587/// This function returns the parameters required to implement this pattern. If
588/// the pattern is not applicable, then failure is returned.
589///
590/// ### Example:
591/// ```
592/// %result = tensor.collapse_shape %0 [[0, 1], [2, 3]]
593/// : tensor<?x1x30x10xf32> to tensor<?x300xf32>
594/// ```
595/// can be transformed to
596/// ```
597/// %tmp = tensor.extract_slice %0 [0, 0, 0, 0]
598/// [0, %dim1, 30, 30]
599/// [1, 1, 1 1]
600/// : tensor<?x1x30x10xf32> to tensor<?x30x10xf32>
601/// %result = tensor.collapse_shape %tmp [[0], [1, 2]]
602/// : tensor<?x30x10xf32> to tensor<?x300xf32>
603/// ```
604///
605/// ### Example:
606/// ```
607/// %result = tensor.collapse_shape %1 [[0, 1], [2]]
608/// : tensor<?x1x30xf32> to tensor<?x30xf32>
609/// ```
610/// can be transformed to
611/// ```
612/// %result = tensor.extract_slice %1 [0, 0, 0]
613/// [%dim2, 1, 30]
614/// [1, 1, 1]
615/// : tensor<?x1x30xf32> to tensor<?x30xf32>
616/// ```
617FailureOr<CollapseShapeRankReducingSliceSimplificationInfo>
618getSimplifyCollapseShapeWithRankReducingSliceInfo(
619 RankedTensorType sourceType,
620 ArrayRef<ReassociationIndices> reassociationIndices);
621
622struct PackingMetadata {
623 SmallVector<int64_t> insertPositions;
624 SmallVector<int64_t> outerPositions;
625 SmallVector<ReassociationIndices> reassociations;
626};
627
628/// Given a vector of `positions` indices representing desired packing insertion
629/// points into a target vector (i.e. pack/unpack.inner_dim_pos), compute the
630/// final positions in the target shape as well as the reshape reassociations.
631// Note: This should not be called with a large positions array (or the
632// implementation needs to be updated to use an N.log N sort instead of
633// repeated N^2 counts).
634PackingMetadata computePackingMetadata(int64_t packedRank,
635 ArrayRef<int64_t> innerDimPos);
636
637/// Try to remove a tensor operation if it would only reshape a constant.
638/// Removes the op and replaces the constant with a new constant of the result
639/// shape. When an optional cst attribute is passed, it is reshaped only if the
640/// splat value matches the value in the attribute.
641OpFoldResult reshapeConstantSource(DenseElementsAttr source, TensorType result,
642 std::optional<Attribute> cst = std::nullopt);
643} // namespace mlir
644
645#endif // MLIR_DIALECT_UTILS_RESHAPEOPSUTILS_H
return success()
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
ArrayAttr()
static RankedTensorType sliceResultType(Type operandType, GridOp grid, ArrayRef< GridAxis > gridAxes, int64_t sliceAxis)
IntegerAttr getIndexAttr(int64_t value)
Definition Builders.cpp:116
An attribute that represents a reference to a dense vector or tensor object.
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
This class represents a single result from folding an operation.
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
Tensor types represent multi-dimensional arrays, and have two variants: RankedTensorType and Unranked...
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
Definition ArithOps.cpp:398
Include the generated interface declarations.
SmallVector< OpFoldResult > getMixedValues(ArrayRef< int64_t > staticValues, ValueRange dynamicValues, MLIRContext *context)
Return a vector of OpFoldResults with the same size a staticValues, but all elements for which Shaped...
llvm::SmallBitVector getSlicedDimensions(ArrayRef< OpFoldResult > sliceInputShape, ArrayRef< Range > sliceParams)
The input parameters offsets, sizes, strides specify a rectangular non rank-reducing slice of the col...
ArrayRef< int64_t > ReassociationIndicesRef
constexpr StringRef getReassociationAttrName()
Attribute name for the ArrayAttr which encodes reassociation indices.
std::optional< int64_t > getConstantIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
InFlightDiagnostic emitError(Location loc)
Utility method to emit an error message using this location.
SmallVector< AffineMap, 4 > getSymbolLessAffineMaps(ArrayRef< ReassociationExprs > reassociation)
Constructs affine maps out of Array<Array<AffineExpr>>.
SmallVector< ReassociationIndices, 2 > convertReassociationMapsToIndices(ArrayRef< ReassociationExprs > reassociationExprs)
Convert Array<Array<AffineExpr>> to Array<Array<int64_t>>.
OpFoldResult foldReshapeOp(ReshapeOpTy reshapeOp, ArrayRef< Attribute > operands)
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...
std::optional< SmallVector< ReassociationIndices > > getReassociationIndicesForCollapse(ArrayRef< int64_t > sourceShape, ArrayRef< int64_t > targetShape)
Returns the reassociation maps to collapse sourceShape to targetShape if possible.
SmallVector< SmallVector< AffineExpr, 2 >, 2 > convertReassociationIndicesToExprs(MLIRContext *context, ArrayRef< ReassociationIndices > reassociationIndices)
Convert reassociation indices to affine expressions.
SmallVector< AffineExpr, 2 > ReassociationExprs
bool isReassociationValid(ArrayRef< AffineMap > reassociation, int *invalidIndex=nullptr)
Return true if the reassociation specification is valid, false otherwise.
std::optional< SmallVector< ReassociationIndices > > composeReassociationIndices(ArrayRef< ReassociationIndices > producerReassociations, ArrayRef< ReassociationIndices > consumerReassociations, MLIRContext *context)
Compose reassociation maps that are used in pair of reshape ops where one is a producer and other is ...
llvm::SmallBitVector getLinearizedDimensions(ArrayRef< ReassociationIndices > reassociationIndices)
Determine which dimensions are linearized by a tensor.collapse_shape op by inspecting its reassociati...
SmallVector< int64_t, 2 > ReassociationIndices
Definition Utils.h:27
ArrayAttr getReassociationIndicesAttribute(Builder &b, ArrayRef< ReassociationIndices > reassociation)
Wraps a list of reassociations in an ArrayAttr.
llvm::function_ref< Fn > function_ref
Definition LLVM.h:147
LogicalResult verifyReassociationIndicesNotEmpty(ReshapeOpTy op)
Verify that none of the reassociation groups is empty.
Common verifier for reshape-like types.
LogicalResult matchAndRewrite(CollapseOpTy collapseOp, PatternRewriter &rewriter) const override
LogicalResult matchAndRewrite(ExpandOpTy expandOp, PatternRewriter &rewriter) const override
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})