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