14#ifndef MLIR_DIALECT_UTILS_RESHAPEOPSUTILS_H
15#define MLIR_DIALECT_UTILS_RESHAPEOPSUTILS_H
22#include "llvm/ADT/STLExtras.h"
23#include "llvm/ADT/StringRef.h"
48 ArrayRef<ReassociationIndices> producerReassociations,
49 ArrayRef<ReassociationIndices> consumerReassociations,
50 MLIRContext *context);
54 MLIRContext *context, ArrayRef<ReassociationIndices> reassociationIndices);
57SmallVector<AffineMap, 4>
63 ArrayRef<ReassociationIndices> reassociation);
67 ArrayRef<ReassociationExprs> reassociationExprs);
72std::optional<SmallVector<ReassociationIndices>>
77std::optional<SmallVector<ReassociationIndices>>
79 ArrayRef<int64_t> targetShape);
85 int *invalidIndex =
nullptr);
88template <
typename ReshapeOpTy>
91 op.getReassociationIndices(),
93 return op.emitOpError(
"reassociation indices must not be empty");
98template <
typename ReshapeOpTy,
typename InverseReshapeOpTy>
102 if (reshapeOp.getSrcType() == reshapeOp.getType())
103 return reshapeOp.getSrc();
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);
117 reshapeOp.getSrc().template getDefiningOp<InverseReshapeOpTy>();
120 auto srcType = reshapeSrcOp.getSrcType();
121 auto resultType = reshapeOp.getResultType();
122 if (srcType != resultType)
125 if (llvm::count_if(srcType.getShape(), ShapedType::isDynamic) < 2) {
126 return reshapeSrcOp.getSrc();
135 auto reassociations = reshapeOp.getReassociationIndices();
136 if (reassociations != reshapeSrcOp.getReassociationIndices())
140 if (srcType.getRank() < reshapeSrcOp.getResultType().getRank())
141 return reshapeSrcOp.getSrc();
142 if (llvm::all_of(reassociations, [&](
auto reInd) {
144 srcType.getShape().slice(reInd.front(), reInd.size());
145 return llvm::count_if(srcSlice, ShapedType::isDynamic) < 2;
147 return reshapeSrcOp.getSrc();
154template <
typename Op,
typename T>
155LogicalResult verifyReshapeLikeTypes(Op op, T expandedType, T collapsedType,
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 <<
'.';
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() <<
").";
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()
180 return op.emitOpError(
"expected reassociation map #")
181 << invalidIdx <<
" to be valid and contiguous.";
183 return reshapeLikeShapesAreCompatible(
184 [&](
const Twine &msg) {
return op->emitOpError(msg); },
185 collapsedType.getShape(), expandedType.getShape(),
186 op.getReassociationIndices(), isExpansion);
194LogicalResult reshapeLikeShapesAreCompatible(
200bool hasNonIdentityLayout(
Type type);
202enum class ReshapeOpKind { kExpand, kCollapse };
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 {
212 reshapeOp.getSrc().template getDefiningOp<ReshapeOpTy>();
216 ShapedType resultType = reshapeOp.getResultType();
218 if (hasNonIdentityLayout(srcReshapeOp.getSrc().getType()) ||
219 hasNonIdentityLayout(reshapeOp.getSrc().getType()) ||
220 hasNonIdentityLayout(reshapeOp.getResult().getType()))
223 std::optional<SmallVector<ReassociationIndices>> reassociationIndices =
225 reshapeOp.getReassociationIndices(),
226 rewriter.getContext());
227 if (!reassociationIndices)
230 if constexpr (opKind == ReshapeOpKind::kExpand) {
231 SmallVector<OpFoldResult> outputShape(
233 reshapeOp.getOutputShape(), rewriter));
234 rewriter.replaceOpWithNewOp<ReshapeOpTy>(
235 reshapeOp, resultType, srcReshapeOp.getSrc(), *reassociationIndices,
238 rewriter.replaceOpWithNewOp<ReshapeOpTy>(
239 reshapeOp, resultType, srcReshapeOp.getSrc(), *reassociationIndices);
273template <
typename CollapseOpTy,
typename ExpandOpTy,
typename CastOpTy,
274 typename DimOpTy,
typename TensorTy>
279 auto expandOp = collapseOp.getSrc().template getDefiningOp<ExpandOpTy>();
283 ShapedType srcType = expandOp.getSrcType();
284 ShapedType resultType = collapseOp.getResultType();
286 if (hasNonIdentityLayout(collapseOp.getSrc().getType()) ||
287 hasNonIdentityLayout(expandOp.getSrc().getType()) ||
288 hasNonIdentityLayout(expandOp.getResult().getType()))
291 int64_t srcRank = srcType.getRank();
292 int64_t resultRank = resultType.getRank();
293 if (srcType == resultType)
297 lowerRankReassociation;
299 if (srcRank > resultRank) {
300 higherRankReassociation = expandOp.getReassociationIndices();
301 lowerRankReassociation = collapseOp.getReassociationIndices();
303 higherRankReassociation = collapseOp.getReassociationIndices();
304 lowerRankReassociation = expandOp.getReassociationIndices();
307 size_t higherRankIndicesID = 0;
309 for (
const auto &lowerRankIndices : lowerRankReassociation) {
311 while (higherRankIndicesID < higherRankReassociation.size()) {
312 auto rightmostIndex =
313 higherRankReassociation[higherRankIndicesID].back();
314 if (rightmostIndex > lowerRankIndices.back())
316 composedIndices.push_back(higherRankIndicesID++);
317 if (rightmostIndex == lowerRankIndices.back())
320 composedReassociation.push_back(composedIndices);
322 if (srcRank > resultRank) {
324 collapseOp, resultType, expandOp.getSrc(), composedReassociation);
325 }
else if (srcRank < resultRank) {
329 expandOp.getMixedOutputShape();
332 collapseOp.getReassociationIndices()) {
338 numStaticElems *= maybeCst.value();
341 dynamicSizes.push_back(cast<Value>(size));
343 if (dynamicSizes.empty()) {
344 newOutputShape.push_back(rewriter.
getIndexAttr(numStaticElems));
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(
358 arith::IntegerOverflowFlags::nsw);
360 newOutputShape.push_back(
result);
363 collapseOp, resultType, expandOp.getSrc(), composedReassociation,
368 assert(llvm::equal(srcType.getShape(), resultType.getShape()) &&
369 "expected same shape");
377template <
typename ExpandOpTy,
typename CollapseOpTy,
typename CastOpTy>
382 auto collapseOp = expandOp.getSrc().template getDefiningOp<CollapseOpTy>();
386 ShapedType srcType = collapseOp.getSrcType();
387 ShapedType resultType = expandOp.getResultType();
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)) {
395 collapseOp.getSrc());
401 int64_t srcRank = srcType.getRank();
402 int64_t resultRank = resultType.getRank();
403 if (srcRank == resultRank)
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)
416 expandOp, resultType, collapseOp.getSrc(), *composedReassociation);
419 auto composedReassociation =
420 findCollapsingReassociation(resultReassociation, srcReassociation,
421 resultType.getShape(), srcType.getShape());
422 if (!composedReassociation)
426 expandOp.getStaticOutputShape(), expandOp.getOutputShape(), rewriter));
428 expandOp, resultType, collapseOp.getSrc(), *composedReassociation,
436 std::optional<SmallVector<ReassociationIndices>> findCollapsingReassociation(
442 if (srcReassociation.empty())
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());
452 if (llvm::count_if(srcSubShape, ShapedType::isDynamic) >= 2 &&
453 llvm::count_if(resultSubShape, ShapedType::isDynamic) >= 2)
456 if (srcSubShape.size() == resultSubShape.size()) {
457 if (srcSubShape != resultSubShape)
460 for (
auto index : llvm::seq<int64_t>(0, srcSubShape.size())) {
461 composedReassociation.emplace_back(1, srcIndices.front() + index);
467 auto subShapeReassociation =
469 if (!subShapeReassociation)
473 for (
auto &subshapeIndices : *subShapeReassociation) {
475 for (int64_t index : subshapeIndices)
476 shapeIndices.push_back(srcIndices.front() + index);
477 composedReassociation.push_back(shapeIndices);
480 return {std::move(composedReassociation)};
527class SliceFromCollapseHelper {
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),
539 extractSliceParams)) {}
551 SmallVector<Range> getExtractSliceParams(MLIRContext *ctx,
552 ArrayRef<ValueRange> multiIndices);
559 SmallVector<Range> getInsertSliceParams(MLIRContext *ctx,
563 SmallVector<ReassociationIndices> reassociationIndices;
564 SmallVector<OpFoldResult> collapseShapeInputShape;
565 SmallVector<OpFoldResult> collapseShapeOutputShape;
566 SmallVector<Range> sliceParams;
567 llvm::SmallBitVector linearizedDimensions;
568 llvm::SmallBitVector slicedDimensions;
573struct CollapseShapeRankReducingSliceSimplificationInfo {
578 std::optional<SmallVector<ReassociationIndices>> newReassociationIndices;
617FailureOr<CollapseShapeRankReducingSliceSimplificationInfo>
618getSimplifyCollapseShapeWithRankReducingSliceInfo(
619 RankedTensorType sourceType,
622struct PackingMetadata {
623 SmallVector<int64_t> insertPositions;
624 SmallVector<int64_t> outerPositions;
625 SmallVector<ReassociationIndices> reassociations;
634PackingMetadata computePackingMetadata(int64_t packedRank,
642 std::optional<Attribute> cst = std::nullopt);
static RankedTensorType sliceResultType(Type operandType, GridOp grid, ArrayRef< GridAxis > gridAxes, int64_t sliceAxis)
IntegerAttr getIndexAttr(int64_t value)
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...
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...
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
static ConstantIndexOp create(OpBuilder &builder, Location location, int64_t value)
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
ArrayAttr getReassociationIndicesAttribute(Builder &b, ArrayRef< ReassociationIndices > reassociation)
Wraps a list of reassociations in an ArrayAttr.
llvm::function_ref< Fn > function_ref
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={})