#include "mlir/Dialect/SCF/IR/SCF.h"
#include "mlir/Dialect/SCF/Utils/Utils.h"
#include "mlir/IR/Builders.h"
#include "mlir/Pass/Pass.h"
using namespace mlir;
namespace {
class SimpleParametricLoopTilingPass
: public PassWrapper<SimpleParametricLoopTilingPass, OperationPass<>> {
public:
MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(SimpleParametricLoopTilingPass)
StringRef getArgument() const final {
return "test-extract-fixed-outer-loops";
}
StringRef getDescription() const final {
return "test application of parametric tiling to the outer loops so that "
"the ranges of outer loops become static";
}
SimpleParametricLoopTilingPass() = default;
SimpleParametricLoopTilingPass(const SimpleParametricLoopTilingPass &) {}
explicit SimpleParametricLoopTilingPass(ArrayRef<int64_t> outerLoopSizes) {
sizes = outerLoopSizes;
}
void runOnOperation() override {
getOperation()->walk([this](scf::ForOp op) {
if (op->getParentRegion()->getParentOfType<scf::ForOp>())
return;
extractFixedOuterLoops(op, sizes);
});
}
ListOption<int64_t> sizes{
*this, "test-outer-loop-sizes",
llvm::cl::desc(
"fixed number of iterations that the outer loops should have")};
};
}
namespace mlir {
namespace test {
void registerSimpleParametricTilingPass() {
PassRegistration<SimpleParametricLoopTilingPass>();
}
}
}