25#define GEN_PASS_DEF_GPUDECOMPOSEMEMREFSPASS
26#include "mlir/Dialect/GPU/Transforms/Passes.h.inc"
32 auto sourceType = cast<BaseMemRefType>(source.
getType());
37 StridedLayoutAttr::get(source.
getContext(), staticOffsets.front(), {});
38 return MemRefType::get({}, sourceType.getElementType(), stridedLayout,
39 sourceType.getMemorySpace());
54static std::tuple<Value, OpFoldResult, SmallVector<OpFoldResult>>
58 auto sourceType = cast<MemRefType>(source.
getType());
59 auto sourceRank =
static_cast<unsigned>(sourceType.getRank());
61 memref::ExtractStridedMetadataOp newExtractStridedMetadata;
65 newExtractStridedMetadata =
66 memref::ExtractStridedMetadataOp::create(rewriter, loc, source);
69 auto &&[sourceStrides, sourceOffset] = sourceType.getStridesAndOffset();
73 : rewriter.getIndexAttr(dim);
77 getDim(sourceOffset, newExtractStridedMetadata.getOffset());
78 ValueRange sourceStridesVals = newExtractStridedMetadata.getStrides();
81 origStrides.reserve(sourceRank);
84 strides.reserve(sourceRank);
88 for (
auto i : llvm::seq(0u, sourceRank)) {
89 OpFoldResult origStride = getDim(sourceStrides[i], sourceStridesVals[i]);
91 if (!subStrides.empty()) {
93 rewriter, loc, s0 * s1, {subStrides[i], origStride}));
96 origStrides.emplace_back(origStride);
99 auto &&[expr, values] =
103 return {newExtractStridedMetadata.getBaseBuffer(), finalOffset, strides};
109 auto &&[base, offset, ignore] =
112 return memref::ReinterpretCastOp::create(rewriter, loc, retType, base, offset,
118 auto type = cast<MemRefType>(val.
getType());
119 return type.getRank() != 0;
123 auto type = cast<MemRefType>(val.
getType());
124 return type.getLayout().isIdentity() ||
125 isa<StridedLayoutAttr>(type.getLayout());
132 LogicalResult matchAndRewrite(memref::LoadOp op,
133 PatternRewriter &rewriter)
const override {
137 Value memref = op.getMemref();
144 Location loc = op.getLoc();
145 Value flatMemref =
getFlatMemref(rewriter, loc, memref, op.getIndices());
147 op, flatMemref,
ValueRange{}, op.getNontemporalAttr(),
148 op.getAlignmentAttr(), op.getInvariantAttr());
156 LogicalResult matchAndRewrite(memref::StoreOp op,
157 PatternRewriter &rewriter)
const override {
161 Value memref = op.getMemref();
168 Location loc = op.getLoc();
169 Value flatMemref =
getFlatMemref(rewriter, loc, memref, op.getIndices());
170 Value value = op.getValue();
172 op, value, flatMemref,
ValueRange{}, op.getNontemporalAttr(),
173 op.getAlignmentAttr());
181 LogicalResult matchAndRewrite(memref::SubViewOp op,
182 PatternRewriter &rewriter)
const override {
186 Value memref = op.getSource();
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] =
200 auto srcType = cast<MemRefType>(memref.
getType());
201 auto resultType = cast<MemRefType>(op.getType());
202 unsigned subRank =
static_cast<unsigned>(resultType.getRank());
204 llvm::SmallBitVector droppedDims = op.getDroppedDims();
206 SmallVector<OpFoldResult> finalSizes;
207 finalSizes.reserve(subRank);
209 SmallVector<OpFoldResult> finalStrides;
210 finalStrides.reserve(subRank);
212 for (
auto i : llvm::seq(0u,
static_cast<unsigned>(srcType.getRank()))) {
213 if (droppedDims.test(i))
216 finalSizes.push_back(subSizes[i]);
217 finalStrides.push_back(strides[i]);
222 auto flattenedSubview = memref::ReinterpretCastOp::create(
223 rewriter, op.getLoc(), resultType, base, finalOffset, finalSizes,
225 if (resultType == op.getType()) {
226 rewriter.
replaceOp(op, flattenedSubview);
236struct GpuDecomposeMemrefsPass
237 :
public impl::GpuDecomposeMemrefsPassBase<GpuDecomposeMemrefsPass> {
239 void runOnOperation()
override {
245 return signalPassFailure();
252 patterns.
insert<FlattenLoad, FlattenStore, FlattenSubview>(
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={})
Base type for affine expression.
AffineExpr getAffineSymbolExpr(unsigned position)
This class defines the main interface for locations in MLIR and acts as a non-nullable wrapper around...
RAII guard to reset the insertion point of the builder when destroyed.
This class helps build Operations.
void setInsertionPointToStart(Block *block)
Sets the insertion point to the start of the specified block.
void setInsertionPointAfter(Operation *op)
Sets the insertion point to the node after the specified operation, which will cause subsequent inser...
This class represents a single result from folding an operation.
Operation is the basic unit of execution within MLIR.
OpTy getParentOfType()
Return the closest surrounding parent operation that is of type 'OpTy'.
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.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
MLIRContext * getContext() const
Utility to get the associated MLIRContext that this value is defined in.
Type getType() const
Return the type of this value.
Block * getParentBlock()
Return the Block in which this Value is defined.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
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...
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 ®ion, 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...