#include "TestDialect.h"
#include "TestOps.h"
#include "mlir/Analysis/CallGraph.h"
#include "mlir/Dialect/Func/IR/FuncOps.h"
#include "mlir/Dialect/SCF/IR/SCF.h"
#include "mlir/IR/BuiltinOps.h"
#include "mlir/IR/IRMapping.h"
#include "mlir/Pass/Pass.h"
#include "mlir/Transforms/Inliner.h"
#include "mlir/Transforms/InliningUtils.h"
#include "llvm/ADT/StringSet.h"
using namespace mlir;
using namespace test;
namespace {
struct InlinerCallback
: public PassWrapper<InlinerCallback, OperationPass<func::FuncOp>> {
MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(InlinerCallback)
StringRef getArgument() const final { return "test-inline-callback"; }
StringRef getDescription() const final {
return "Test inlining region calls with call back functions";
}
void getDependentDialects(DialectRegistry ®istry) const override {
registry.insert<scf::SCFDialect>();
}
static LogicalResult runPipelineHelper(Pass &pass, OpPassManager &pipeline,
Operation *op) {
return mlir::cast<InlinerCallback>(pass).runPipeline(pipeline, op);
}
static void testDoClone(OpBuilder &builder, Region *src, Block *inlineBlock,
Block *postInsertBlock, IRMapping &mapper,
bool shouldCloneInlinedRegion) {
mlir::Operation &call = inlineBlock->back();
builder.setInsertionPointAfter(&call);
auto executeRegionOp = mlir::scf::ExecuteRegionOp::create(
builder, call.getLoc(), call.getResultTypes());
mlir::Region ®ion = executeRegionOp.getRegion();
src->cloneInto(®ion, mapper);
inlineBlock->splitBlock(executeRegionOp.getOperation());
for (mlir::Block &block : region) {
for (mlir::Operation &op : llvm::make_early_inc_range(block)) {
if (test::TestReturnOp returnOp =
llvm::dyn_cast<test::TestReturnOp>(&op)) {
mlir::OpBuilder returnBuilder(returnOp);
mlir::scf::YieldOp::create(returnBuilder, returnOp.getLoc(),
returnOp.getOperands());
returnOp.erase();
}
}
}
builder.setInsertionPointAfter(executeRegionOp);
test::TestReturnOp::create(builder, executeRegionOp.getLoc(),
executeRegionOp.getResults());
}
void runOnOperation() override {
InlinerConfig config;
CallGraph &cg = getAnalysis<CallGraph>();
func::FuncOp function = getOperation();
auto profitabilityCb = [&](const mlir::Inliner::ResolvedCall &call) {
return true;
};
config.setCloneCallback([](OpBuilder &builder, Region *src,
Block *inlineBlock, Block *postInsertBlock,
IRMapping &mapper,
bool shouldCloneInlinedRegion) {
return testDoClone(builder, src, inlineBlock, postInsertBlock, mapper,
shouldCloneInlinedRegion);
});
config.setCanHandleMultipleBlocks();
Inliner inliner(function, cg, *this, getAnalysisManager(),
runPipelineHelper, config, profitabilityCb);
SmallVector<func::CallIndirectOp> callers;
function.walk(
[&](func::CallIndirectOp caller) { callers.push_back(caller); });
InlinerInterface interface(&getContext());
for (auto caller : callers) {
auto callee = dyn_cast_or_null<FunctionalRegionOp>(
caller.getCallee().getDefiningOp());
if (!callee)
continue;
if (failed(inlineRegion(
interface, config.getCloneCallback(), &callee.getBody(), caller,
caller.getArgOperands(), caller.getResults(), caller.getLoc(),
!callee.getResult().hasOneUse())))
continue;
caller.erase();
if (callee.use_empty())
callee.erase();
}
}
};
}
namespace mlir {
namespace test {
void registerInlinerCallback() { PassRegistration<InlinerCallback>(); }
}
}