#include "mlir/Transforms/WalkPatternRewriteDriver.h"
#include "mlir/IR/MLIRContext.h"
#include "mlir/IR/Operation.h"
#include "mlir/IR/OperationSupport.h"
#include "mlir/IR/PatternMatch.h"
#include "mlir/IR/Verifier.h"
#include "mlir/IR/Visitors.h"
#include "mlir/Rewrite/PatternApplicator.h"
#include "llvm/ADT/STLExtras.h"
#include "llvm/Support/DebugLog.h"
#include "llvm/Support/ErrorHandling.h"
#define DEBUG_TYPE "walk-rewriter"
namespace mlir {
static void findReachableBlocks(Region ®ion,
DenseSet<Block *> &reachableBlocks) {
Block *entryBlock = ®ion.front();
reachableBlocks.insert(entryBlock);
SmallVector<Block *> worklist({entryBlock});
while (!worklist.empty()) {
Block *block = worklist.pop_back_val();
Operation *terminator = &block->back();
for (Block *successor : terminator->getSuccessors()) {
if (reachableBlocks.contains(successor))
continue;
worklist.push_back(successor);
reachableBlocks.insert(successor);
}
}
}
namespace {
struct WalkAndApplyPatternsAction final
: tracing::ActionImpl<WalkAndApplyPatternsAction> {
MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(WalkAndApplyPatternsAction)
using ActionImpl::ActionImpl;
static constexpr StringLiteral tag = "walk-and-apply-patterns";
void print(raw_ostream &os) const override { os << tag; }
};
#if MLIR_ENABLE_EXPENSIVE_PATTERN_API_CHECKS
struct ErasedOpsListener final : RewriterBase::ForwardingListener {
using RewriterBase::ForwardingListener::ForwardingListener;
void notifyOperationErased(Operation *op) override {
checkErasure(op);
ForwardingListener::notifyOperationErased(op);
}
void notifyBlockErased(Block *block) override {
checkErasure(block->getParentOp());
ForwardingListener::notifyBlockErased(block);
}
void checkErasure(Operation *op) const {
Operation *ancestorOp = op;
while (ancestorOp && ancestorOp != visitedOp)
ancestorOp = ancestorOp->getParentOp();
if (ancestorOp != visitedOp)
llvm::report_fatal_error(
"unsupported erasure in WalkPatternRewriter; "
"erasure is only supported for matched ops and their descendants");
}
Operation *visitedOp = nullptr;
};
#endif
}
void walkAndApplyPatterns(Operation *op,
const FrozenRewritePatternSet &patterns,
RewriterBase::Listener *listener) {
#if MLIR_ENABLE_EXPENSIVE_PATTERN_API_CHECKS
if (failed(verify(op)))
llvm::report_fatal_error("walk pattern rewriter input IR failed to verify");
#endif
MLIRContext *ctx = op->getContext();
PatternRewriter rewriter(ctx);
#if MLIR_ENABLE_EXPENSIVE_PATTERN_API_CHECKS
ErasedOpsListener erasedListener(listener);
rewriter.setListener(&erasedListener);
#else
rewriter.setListener(listener);
#endif
PatternApplicator applicator(patterns);
applicator.applyDefaultCostModel();
struct RegionReachableOpIterator {
RegionReachableOpIterator(Region *region) : region(region) {
regionIt = region->begin();
if (regionIt != region->end())
blockIt = regionIt->begin();
if (!llvm::hasSingleElement(*region))
findReachableBlocks(*region, reachableBlocks);
}
void advance() {
assert(regionIt != region->end());
hasVisitedRegions = false;
if (blockIt == regionIt->end()) {
++regionIt;
while (regionIt != region->end() &&
!reachableBlocks.contains(&*regionIt))
++regionIt;
if (regionIt != region->end())
blockIt = regionIt->begin();
return;
}
++blockIt;
if (blockIt != regionIt->end()) {
LDBG() << "Incrementing block iterator, next op: "
<< OpWithFlags(&*blockIt, OpPrintingFlags().skipRegions());
}
}
Region *region;
Region::iterator regionIt;
Block::iterator blockIt;
DenseSet<Block *> reachableBlocks;
bool hasVisitedRegions = false;
};
SmallVector<RegionReachableOpIterator> worklist;
LDBG() << "Starting walk-based pattern rewrite driver";
ctx->executeAction<WalkAndApplyPatternsAction>(
[&] {
for (Region ®ion : op->getRegions()) {
assert(worklist.empty());
if (region.empty())
continue;
worklist.push_back({®ion});
while (!worklist.empty()) {
RegionReachableOpIterator &it = worklist.back();
if (it.regionIt == it.region->end()) {
worklist.pop_back();
continue;
}
if (it.blockIt == it.regionIt->end()) {
it.advance();
continue;
}
Operation *op = &*it.blockIt;
if (!it.hasVisitedRegions) {
it.hasVisitedRegions = true;
for (Region &nestedRegion : llvm::reverse(op->getRegions())) {
if (nestedRegion.empty())
continue;
worklist.push_back({&nestedRegion});
}
}
if (&it != &worklist.back())
continue;
it.advance();
LDBG() << "Visiting op: "
<< OpWithFlags(op, OpPrintingFlags().skipRegions());
#if MLIR_ENABLE_EXPENSIVE_PATTERN_API_CHECKS
erasedListener.visitedOp = op;
#endif
if (succeeded(applicator.matchAndRewrite(op, rewriter)))
LDBG() << "\tOp matched and rewritten";
}
}
},
{op});
#if MLIR_ENABLE_EXPENSIVE_PATTERN_API_CHECKS
if (failed(verify(op)))
llvm::report_fatal_error(
"walk pattern rewriter result IR failed to verify");
#endif
}
}