70#include "llvm/ADT/StringExtras.h"
71#include "llvm/Support/Debug.h"
75#define GEN_PASS_DEF_ACCLOOPTILING
76#include "mlir/Dialect/OpenACC/Transforms/Passes.h.inc"
80#define DEBUG_TYPE "acc-loop-tile"
86 ACCLoopTilingImpl(
MLIRContext *context, int32_t defaultTileSize,
89 defaultTileSize(defaultTileSize), accSupport(accSupport) {}
94 LogicalResult checkTileSizeTypes(acc::LoopOp loop,
96 auto ivTypes = loop.getBody().getArgumentTypes();
97 for (
size_t i = 0; i < tileSizes.size() && i < ivTypes.size(); ++i) {
98 Type tileType = tileSizes[i].getType();
99 Type ivType = ivTypes[i];
103 if (constVal && *constVal < 0)
107 auto tileIntType = dyn_cast<IntegerType>(tileType);
108 auto ivIntType = dyn_cast<IntegerType>(ivType);
109 if (tileIntType && ivIntType) {
110 if (tileIntType.getWidth() > ivIntType.getWidth()) {
111 accSupport.
emitNYI(loop.getLoc(),
112 "tile size type (i" +
113 std::to_string(tileIntType.getWidth()) +
114 ") is wider than loop IV type (i" +
115 std::to_string(ivIntType.getWidth()) +
")");
123 void emitTilingRemarks(acc::LoopOp loop,
ArrayRef<Value> tileSizes)
const {
128 auto getTileSizeStr = [&](
Value v) -> std::string {
131 if (name.empty() || name ==
"-1")
136 for (
Value v : tileSizes)
137 tileStrs.push_back(getTileSizeStr(v));
138 return "Tiling " + std::to_string(tileSizes.size()) +
139 "-level loop nest with tile(" + llvm::join(tileStrs,
",") +
146 for (
Value tileSize : tileSizes) {
148 if (val && *val < 0) {
152 return "Picking default tile size " +
153 std::to_string(defaultTileSize) +
154 " for unknown tile size '*'";
161 LogicalResult matchAndRewrite(acc::LoopOp origLoop,
164 if (origLoop.getTileValues().empty())
168 origLoop.getTileValues().end());
173 if (origLoop.getCollapseAttr()) {
174 accSupport.
emitNYI(origLoop.getLoc(),
175 "a tile clause combined with a collapse clause on the "
181 if (failed(checkTileSizeTypes(origLoop, tileSizes)))
186 emitTilingRemarks(origLoop, tileSizes);
188 LLVM_DEBUG(llvm::dbgs() <<
"\nBefore tiling:\n" << *origLoop <<
"\n");
192 origLoop.getTileOperandsMutable().clear();
193 origLoop.removeTileOperandsSegmentsAttr();
194 origLoop.removeTileOperandsDeviceTypeAttr();
202 LLVM_DEBUG(llvm::dbgs() <<
"\nAfter tiling:\n " << *origLoop <<
"\n");
207 int32_t defaultTileSize;
213 using ACCLoopTilingBase<ACCLoopTiling>::ACCLoopTilingBase;
215 void runOnOperation()
override {
216 func::FuncOp funcOp = getOperation();
221 patterns.
insert<ACCLoopTilingImpl>(context, defaultTileSize, accSupport);
This class allows control over how the GreedyPatternRewriteDriver works.
GreedyRewriteConfig & setMaxIterations(int64_t iterations)
GreedyRewriteConfig & setUseTopDownTraversal(bool use=true)
MLIRContext is the top-level object for a collection of MLIR operations.
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
RewritePatternSet & insert(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
virtual void finalizeOpModification(Operation *op)
This method is used to signal the end of an in-place modification of the given operation.
virtual void startOpModification(Operation *op)
This method is used to notify the rewriter that an in-place operation modification is about to happen...
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...
remark::detail::InFlightRemark emitRemark(Operation *op, std::function< std::string()> messageFn, llvm::StringRef category="openacc")
Emit an OpenACC remark with lazy message generation.
std::string getVariableName(Value v, VariableNameConfig config={})
Get the variable name for a given value.
InFlightDiagnostic emitNYI(Location loc, const Twine &message)
Report a case that is not yet supported by the implementation.
mlir::acc::LoopOp tileACCLoops(mlir::acc::LoopOp tileLoop, const llvm::SmallVector< mlir::Value > &tileSizes, int32_t defaultTileSize, mlir::RewriterBase &rewriter)
Tile a single fused acc.loop that carries all associated induction variables (one IV per tile dimensi...
Include the generated interface declarations.
std::optional< int64_t > getConstantIntValue(OpFoldResult ofr)
If ofr is a constant integer or an IntegerAttr, return the integer.
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...
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...