MLIR 24.0.0git
DecomposeMemRefs.cpp
Go to the documentation of this file.
1//===- DecomposeMemRefs.cpp - Decompose memrefs pass implementation -------===//
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 decompose memrefs pass.
10//
11//===----------------------------------------------------------------------===//
12
19#include "mlir/IR/AffineExpr.h"
20#include "mlir/IR/Builders.h"
23
24namespace mlir {
25#define GEN_PASS_DEF_GPUDECOMPOSEMEMREFSPASS
26#include "mlir/Dialect/GPU/Transforms/Passes.h.inc"
27} // namespace mlir
28
29using namespace mlir;
30
31static MemRefType inferCastResultType(Value source, OpFoldResult offset) {
32 auto sourceType = cast<BaseMemRefType>(source.getType());
33 SmallVector<int64_t> staticOffsets;
34 SmallVector<Value> dynamicOffsets;
35 dispatchIndexOpFoldResults(offset, dynamicOffsets, staticOffsets);
36 auto stridedLayout =
37 StridedLayoutAttr::get(source.getContext(), staticOffsets.front(), {});
38 return MemRefType::get({}, sourceType.getElementType(), stridedLayout,
39 sourceType.getMemorySpace());
40}
41
42static void setInsertionPointToStart(OpBuilder &builder, Value val) {
43 if (auto *parentOp = val.getDefiningOp()) {
44 builder.setInsertionPointAfter(parentOp);
45 } else {
47 }
48}
49
50static bool isInsideLaunch(Operation *op) {
51 return op->getParentOfType<gpu::LaunchOp>();
52}
53
54static std::tuple<Value, OpFoldResult, SmallVector<OpFoldResult>>
56 ArrayRef<OpFoldResult> subOffsets,
57 ArrayRef<OpFoldResult> subStrides = {}) {
58 auto sourceType = cast<MemRefType>(source.getType());
59 auto sourceRank = static_cast<unsigned>(sourceType.getRank());
60
61 memref::ExtractStridedMetadataOp newExtractStridedMetadata;
62 {
63 OpBuilder::InsertionGuard g(rewriter);
64 setInsertionPointToStart(rewriter, source);
65 newExtractStridedMetadata =
66 memref::ExtractStridedMetadataOp::create(rewriter, loc, source);
67 }
68
69 auto &&[sourceStrides, sourceOffset] = sourceType.getStridesAndOffset();
70
71 auto getDim = [&](int64_t dim, Value dimVal) -> OpFoldResult {
72 return ShapedType::isDynamic(dim) ? getAsOpFoldResult(dimVal)
73 : rewriter.getIndexAttr(dim);
74 };
75
76 OpFoldResult origOffset =
77 getDim(sourceOffset, newExtractStridedMetadata.getOffset());
78 ValueRange sourceStridesVals = newExtractStridedMetadata.getStrides();
79
80 SmallVector<OpFoldResult> origStrides;
81 origStrides.reserve(sourceRank);
82
84 strides.reserve(sourceRank);
85
86 AffineExpr s0 = rewriter.getAffineSymbolExpr(0);
87 AffineExpr s1 = rewriter.getAffineSymbolExpr(1);
88 for (auto i : llvm::seq(0u, sourceRank)) {
89 OpFoldResult origStride = getDim(sourceStrides[i], sourceStridesVals[i]);
90
91 if (!subStrides.empty()) {
93 rewriter, loc, s0 * s1, {subStrides[i], origStride}));
94 }
95
96 origStrides.emplace_back(origStride);
97 }
98
99 auto &&[expr, values] =
100 computeLinearIndex(origOffset, origStrides, subOffsets);
101 OpFoldResult finalOffset =
102 affine::makeComposedFoldedAffineApply(rewriter, loc, expr, values);
103 return {newExtractStridedMetadata.getBaseBuffer(), finalOffset, strides};
104}
105
106static Value getFlatMemref(OpBuilder &rewriter, Location loc, Value source,
107 ValueRange offsets) {
108 SmallVector<OpFoldResult> offsetsTemp = getAsOpFoldResult(offsets);
109 auto &&[base, offset, ignore] =
110 getFlatOffsetAndStrides(rewriter, loc, source, offsetsTemp);
111 MemRefType retType = inferCastResultType(base, offset);
112 return memref::ReinterpretCastOp::create(rewriter, loc, retType, base, offset,
115}
116
117static bool needFlatten(Value val) {
118 auto type = cast<MemRefType>(val.getType());
119 return type.getRank() != 0;
120}
121
122static bool checkLayout(Value val) {
123 auto type = cast<MemRefType>(val.getType());
124 return type.getLayout().isIdentity() ||
125 isa<StridedLayoutAttr>(type.getLayout());
126}
127
128namespace {
129struct FlattenLoad : public OpRewritePattern<memref::LoadOp> {
131
132 LogicalResult matchAndRewrite(memref::LoadOp op,
133 PatternRewriter &rewriter) const override {
134 if (!isInsideLaunch(op))
135 return rewriter.notifyMatchFailure(op, "not inside gpu.launch");
136
137 Value memref = op.getMemref();
138 if (!needFlatten(memref))
139 return rewriter.notifyMatchFailure(op, "nothing to do");
140
141 if (!checkLayout(memref))
142 return rewriter.notifyMatchFailure(op, "unsupported layout");
143
144 Location loc = op.getLoc();
145 Value flatMemref = getFlatMemref(rewriter, loc, memref, op.getIndices());
146 rewriter.replaceOpWithNewOp<memref::LoadOp>(
147 op, flatMemref, ValueRange{}, op.getNontemporalAttr(),
148 op.getAlignmentAttr(), op.getInvariantAttr());
149 return success();
150 }
151};
152
153struct FlattenStore : public OpRewritePattern<memref::StoreOp> {
155
156 LogicalResult matchAndRewrite(memref::StoreOp op,
157 PatternRewriter &rewriter) const override {
158 if (!isInsideLaunch(op))
159 return rewriter.notifyMatchFailure(op, "not inside gpu.launch");
160
161 Value memref = op.getMemref();
162 if (!needFlatten(memref))
163 return rewriter.notifyMatchFailure(op, "nothing to do");
164
165 if (!checkLayout(memref))
166 return rewriter.notifyMatchFailure(op, "unsupported layout");
167
168 Location loc = op.getLoc();
169 Value flatMemref = getFlatMemref(rewriter, loc, memref, op.getIndices());
170 Value value = op.getValue();
171 rewriter.replaceOpWithNewOp<memref::StoreOp>(
172 op, value, flatMemref, ValueRange{}, op.getNontemporalAttr(),
173 op.getAlignmentAttr());
174 return success();
175 }
176};
177
178struct FlattenSubview : public OpRewritePattern<memref::SubViewOp> {
180
181 LogicalResult matchAndRewrite(memref::SubViewOp op,
182 PatternRewriter &rewriter) const override {
183 if (!isInsideLaunch(op))
184 return rewriter.notifyMatchFailure(op, "not inside gpu.launch");
185
186 Value memref = op.getSource();
187 if (!needFlatten(memref))
188 return rewriter.notifyMatchFailure(op, "nothing to do");
189
190 if (!checkLayout(memref))
191 return rewriter.notifyMatchFailure(op, "unsupported layout");
192
193 Location loc = op.getLoc();
194 SmallVector<OpFoldResult> subOffsets = op.getMixedOffsets();
195 SmallVector<OpFoldResult> subSizes = op.getMixedSizes();
196 SmallVector<OpFoldResult> subStrides = op.getMixedStrides();
197 auto &&[base, finalOffset, strides] =
198 getFlatOffsetAndStrides(rewriter, loc, memref, subOffsets, subStrides);
199
200 auto srcType = cast<MemRefType>(memref.getType());
201 auto resultType = cast<MemRefType>(op.getType());
202 unsigned subRank = static_cast<unsigned>(resultType.getRank());
203
204 llvm::SmallBitVector droppedDims = op.getDroppedDims();
205
206 SmallVector<OpFoldResult> finalSizes;
207 finalSizes.reserve(subRank);
208
209 SmallVector<OpFoldResult> finalStrides;
210 finalStrides.reserve(subRank);
211
212 for (auto i : llvm::seq(0u, static_cast<unsigned>(srcType.getRank()))) {
213 if (droppedDims.test(i))
214 continue;
215
216 finalSizes.push_back(subSizes[i]);
217 finalStrides.push_back(strides[i]);
218 }
219
220 resultType = updateTypeFromMetadata(resultType, finalOffset, finalSizes,
221 finalStrides);
222 auto flattenedSubview = memref::ReinterpretCastOp::create(
223 rewriter, op.getLoc(), resultType, base, finalOffset, finalSizes,
224 finalStrides);
225 if (resultType == op.getType()) {
226 rewriter.replaceOp(op, flattenedSubview);
227 return success();
228 }
229 // Preserve the original result type expected by existing users.
230 rewriter.replaceOpWithNewOp<memref::CastOp>(op, op.getType(),
231 flattenedSubview);
232 return success();
233 }
234};
235
236struct GpuDecomposeMemrefsPass
237 : public impl::GpuDecomposeMemrefsPassBase<GpuDecomposeMemrefsPass> {
238
239 void runOnOperation() override {
240 RewritePatternSet patterns(&getContext());
241
243
244 if (failed(applyPatternsGreedily(getOperation(), std::move(patterns))))
245 return signalPassFailure();
246 }
247};
248
249} // namespace
250
252 patterns.insert<FlattenLoad, FlattenStore, FlattenSubview>(
253 patterns.getContext());
254}
return success()
static bool isInsideLaunch(Operation *op)
static bool needFlatten(Value val)
static MemRefType inferCastResultType(Value source, OpFoldResult offset)
static bool checkLayout(Value val)
static Value getFlatMemref(OpBuilder &rewriter, Location loc, Value source, ValueRange offsets)
static void setInsertionPointToStart(OpBuilder &builder, Value val)
static std::tuple< Value, OpFoldResult, SmallVector< OpFoldResult > > getFlatOffsetAndStrides(OpBuilder &rewriter, Location loc, Value source, ArrayRef< OpFoldResult > subOffsets, ArrayRef< OpFoldResult > subStrides={})
b getContext())
Base type for affine expression.
Definition AffineExpr.h:68
AffineExpr getAffineSymbolExpr(unsigned position)
Definition Builders.cpp:377
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
Definition Location.h:76
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 setInsertionPointToStart(Block *block)
Sets the insertion point to the start of the specified block.
Definition Builders.h:434
void setInsertionPointAfter(Operation *op)
Sets the insertion point to the node after the specified operation, which will cause subsequent inser...
Definition Builders.h:415
This class represents a single result from folding an operation.
Operation is the basic unit of execution within MLIR.
Definition Operation.h:87
OpTy getParentOfType()
Return the closest surrounding parent operation that is of type 'OpTy'.
Definition Operation.h:255
RewritePatternSet & insert(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
MLIRContext * getContext() const
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
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,...
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 different types of ranges over Values.
Definition ValueRange.h:389
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
MLIRContext * getContext() const
Utility to get the associated MLIRContext that this value is defined in.
Definition Value.h:108
Type getType() const
Return the type of this value.
Definition Value.h:105
Block * getParentBlock()
Return the Block in which this Value is defined.
Definition Value.cpp:46
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
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...
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
Include the generated interface declarations.
MemRefType updateTypeFromMetadata(MemRefType type, OpFoldResult offset, ArrayRef< OpFoldResult > sizes, ArrayRef< OpFoldResult > strides)
Returns a memref type matching the provided offset, size, and stride metadata.
std::pair< AffineExpr, SmallVector< OpFoldResult > > computeLinearIndex(OpFoldResult sourceOffset, ArrayRef< OpFoldResult > strides, ArrayRef< OpFoldResult > indices)
Compute linear index from provided strides and indices, assuming strided layout.
LogicalResult applyPatternsGreedily(Region &region, const FrozenRewritePatternSet &patterns, GreedyRewriteConfig config=GreedyRewriteConfig(), bool *changed=nullptr)
Rewrite ops in the given region, which must be isolated from above, by repeatedly applying the highes...
void dispatchIndexOpFoldResults(ArrayRef< OpFoldResult > ofrs, SmallVectorImpl< Value > &dynamicVec, SmallVectorImpl< int64_t > &staticVec)
Helper function to dispatch multiple OpFoldResults according to the behavior of dispatchIndexOpFoldRe...
void populateGpuDecomposeMemrefsPatterns(RewritePatternSet &patterns)
Collect a set of patterns to decompose memrefs ops.
OpFoldResult getAsOpFoldResult(Value val)
Given a value, try to extract a constant Attribute.
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})
Patterns must specify the root operation name they match against, and can also specify the benefit of...