1 //===- ResolveShapedTypeResultDims.cpp - Resolve dim ops of result values -===//
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 pass resolves `memref.dim` operations of result values in terms of
10 // shapes of their operands using the `InferShapedTypeOpInterface`.
11 //
12 //===----------------------------------------------------------------------===//
13
14 #include "PassDetail.h"
15 #include "mlir/Dialect/Affine/IR/AffineOps.h"
16 #include "mlir/Dialect/Arithmetic/IR/Arithmetic.h"
17 #include "mlir/Dialect/MemRef/IR/MemRef.h"
18 #include "mlir/Dialect/MemRef/Transforms/Passes.h"
19 #include "mlir/Dialect/Tensor/IR/Tensor.h"
20 #include "mlir/Interfaces/InferTypeOpInterface.h"
21 #include "mlir/Transforms/GreedyPatternRewriteDriver.h"
22
23 using namespace mlir;
24
25 namespace {
26 /// Fold dim of an operation that implements the InferShapedTypeOpInterface
27 template <typename OpTy>
28 struct DimOfShapedTypeOpInterface : public OpRewritePattern<OpTy> {
29 using OpRewritePattern<OpTy>::OpRewritePattern;
30
matchAndRewrite__anon47d8ba0a0111::DimOfShapedTypeOpInterface31 LogicalResult matchAndRewrite(OpTy dimOp,
32 PatternRewriter &rewriter) const override {
33 OpResult dimValue = dimOp.getSource().template dyn_cast<OpResult>();
34 if (!dimValue)
35 return failure();
36 auto shapedTypeOp =
37 dyn_cast<InferShapedTypeOpInterface>(dimValue.getOwner());
38 if (!shapedTypeOp)
39 return failure();
40
41 Optional<int64_t> dimIndex = dimOp.getConstantIndex();
42 if (!dimIndex)
43 return failure();
44
45 SmallVector<Value> reifiedResultShapes;
46 if (failed(shapedTypeOp.reifyReturnTypeShapes(
47 rewriter, shapedTypeOp->getOperands(), reifiedResultShapes)))
48 return failure();
49
50 if (reifiedResultShapes.size() != shapedTypeOp->getNumResults())
51 return failure();
52
53 Value resultShape = reifiedResultShapes[dimValue.getResultNumber()];
54 auto resultShapeType = resultShape.getType().dyn_cast<RankedTensorType>();
55 if (!resultShapeType || !resultShapeType.getElementType().isa<IndexType>())
56 return failure();
57
58 Location loc = dimOp->getLoc();
59 rewriter.replaceOpWithNewOp<tensor::ExtractOp>(
60 dimOp, resultShape,
61 rewriter.createOrFold<arith::ConstantIndexOp>(loc, *dimIndex));
62 return success();
63 }
64 };
65
66 /// Fold dim of an operation that implements the InferShapedTypeOpInterface
67 template <typename OpTy>
68 struct DimOfReifyRankedShapedTypeOpInterface : public OpRewritePattern<OpTy> {
69 using OpRewritePattern<OpTy>::OpRewritePattern;
70
matchAndRewrite__anon47d8ba0a0111::DimOfReifyRankedShapedTypeOpInterface71 LogicalResult matchAndRewrite(OpTy dimOp,
72 PatternRewriter &rewriter) const override {
73 OpResult dimValue = dimOp.getSource().template dyn_cast<OpResult>();
74 if (!dimValue)
75 return failure();
76 auto rankedShapeTypeOp =
77 dyn_cast<ReifyRankedShapedTypeOpInterface>(dimValue.getOwner());
78 if (!rankedShapeTypeOp)
79 return failure();
80
81 Optional<int64_t> dimIndex = dimOp.getConstantIndex();
82 if (!dimIndex)
83 return failure();
84
85 SmallVector<SmallVector<Value>> reifiedResultShapes;
86 if (failed(
87 rankedShapeTypeOp.reifyResultShapes(rewriter, reifiedResultShapes)))
88 return failure();
89
90 if (reifiedResultShapes.size() != rankedShapeTypeOp->getNumResults())
91 return failure();
92
93 unsigned resultNumber = dimValue.getResultNumber();
94 auto sourceType = dimValue.getType().dyn_cast<RankedTensorType>();
95 if (reifiedResultShapes[resultNumber].size() !=
96 static_cast<size_t>(sourceType.getRank()))
97 return failure();
98
99 rewriter.replaceOp(dimOp, reifiedResultShapes[resultNumber][*dimIndex]);
100 return success();
101 }
102 };
103 } // namespace
104
105 //===----------------------------------------------------------------------===//
106 // Pass registration
107 //===----------------------------------------------------------------------===//
108
109 namespace {
110 struct ResolveRankedShapeTypeResultDimsPass final
111 : public ResolveRankedShapeTypeResultDimsBase<
112 ResolveRankedShapeTypeResultDimsPass> {
113 void runOnOperation() override;
114 };
115
116 struct ResolveShapedTypeResultDimsPass final
117 : public ResolveShapedTypeResultDimsBase<ResolveShapedTypeResultDimsPass> {
118 void runOnOperation() override;
119 };
120
121 } // namespace
122
populateResolveRankedShapeTypeResultDimsPatterns(RewritePatternSet & patterns)123 void memref::populateResolveRankedShapeTypeResultDimsPatterns(
124 RewritePatternSet &patterns) {
125 patterns.add<DimOfReifyRankedShapedTypeOpInterface<memref::DimOp>,
126 DimOfReifyRankedShapedTypeOpInterface<tensor::DimOp>>(
127 patterns.getContext());
128 }
129
populateResolveShapedTypeResultDimsPatterns(RewritePatternSet & patterns)130 void memref::populateResolveShapedTypeResultDimsPatterns(
131 RewritePatternSet &patterns) {
132 // TODO: Move tensor::DimOp pattern to the Tensor dialect.
133 patterns.add<DimOfShapedTypeOpInterface<memref::DimOp>,
134 DimOfShapedTypeOpInterface<tensor::DimOp>>(
135 patterns.getContext());
136 }
137
runOnOperation()138 void ResolveRankedShapeTypeResultDimsPass::runOnOperation() {
139 RewritePatternSet patterns(&getContext());
140 memref::populateResolveRankedShapeTypeResultDimsPatterns(patterns);
141 if (failed(applyPatternsAndFoldGreedily(getOperation()->getRegions(),
142 std::move(patterns))))
143 return signalPassFailure();
144 }
145
runOnOperation()146 void ResolveShapedTypeResultDimsPass::runOnOperation() {
147 RewritePatternSet patterns(&getContext());
148 memref::populateResolveRankedShapeTypeResultDimsPatterns(patterns);
149 memref::populateResolveShapedTypeResultDimsPatterns(patterns);
150 if (failed(applyPatternsAndFoldGreedily(getOperation()->getRegions(),
151 std::move(patterns))))
152 return signalPassFailure();
153 }
154
createResolveShapedTypeResultDimsPass()155 std::unique_ptr<Pass> memref::createResolveShapedTypeResultDimsPass() {
156 return std::make_unique<ResolveShapedTypeResultDimsPass>();
157 }
158
createResolveRankedShapeTypeResultDimsPass()159 std::unique_ptr<Pass> memref::createResolveRankedShapeTypeResultDimsPass() {
160 return std::make_unique<ResolveRankedShapeTypeResultDimsPass>();
161 }
162