#include "mlir/Transforms/Passes.h"
#include "mlir/IR/Operation.h"
#include "mlir/IR/SymbolTable.h"
#include "llvm/Support/Debug.h"
#include "llvm/Support/DebugLog.h"
#include "llvm/Support/InterleavedRange.h"
namespace mlir {
#define GEN_PASS_DEF_SYMBOLDCE
#include "mlir/Transforms/Passes.h.inc"
}
using namespace mlir;
#define DEBUG_TYPE "symbol-dce"
namespace {
struct SymbolDCE : public impl::SymbolDCEBase<SymbolDCE> {
void runOnOperation() override;
LogicalResult computeLiveness(Operation *symbolTableOp,
SymbolTableCollection &symbolTable,
bool symbolTableIsHidden,
DenseSet<Operation *> &liveSymbols);
};
}
void SymbolDCE::runOnOperation() {
Operation *symbolTableOp = getOperation();
if (!symbolTableOp->hasTrait<OpTrait::SymbolTable>()) {
symbolTableOp->emitOpError()
<< " was scheduled to run under SymbolDCE, but does not define a "
"symbol table";
return signalPassFailure();
}
bool symbolTableIsHidden = true;
SymbolOpInterface symbol = dyn_cast<SymbolOpInterface>(symbolTableOp);
if (symbolTableOp->getParentOp() && symbol)
symbolTableIsHidden = symbol.isPrivate();
DenseSet<Operation *> liveSymbols;
SymbolTableCollection symbolTable;
if (failed(computeLiveness(symbolTableOp, symbolTable, symbolTableIsHidden,
liveSymbols)))
return signalPassFailure();
symbolTableOp->walk([&](Operation *nestedSymbolTable) {
if (!nestedSymbolTable->hasTrait<OpTrait::SymbolTable>())
return;
for (auto &block : nestedSymbolTable->getRegion(0)) {
for (Operation &op : llvm::make_early_inc_range(block)) {
if (isa<SymbolOpInterface>(&op) && !liveSymbols.count(&op)) {
op.erase();
++numDCE;
}
}
}
});
}
LogicalResult SymbolDCE::computeLiveness(Operation *symbolTableOp,
SymbolTableCollection &symbolTable,
bool symbolTableIsHidden,
DenseSet<Operation *> &liveSymbols) {
LDBG() << "computeLiveness: "
<< OpWithFlags(symbolTableOp, OpPrintingFlags().skipRegions());
SmallVector<Operation *, 16> worklist;
for (auto &block : symbolTableOp->getRegion(0)) {
for (Operation &op : block) {
SymbolOpInterface symbol = dyn_cast<SymbolOpInterface>(&op);
if (!symbol) {
worklist.push_back(&op);
continue;
}
bool isDiscardable = (symbolTableIsHidden || symbol.isPrivate()) &&
symbol.canDiscardOnUseEmpty();
if (!isDiscardable && liveSymbols.insert(&op).second)
worklist.push_back(&op);
}
}
while (!worklist.empty()) {
Operation *op = worklist.pop_back_val();
LDBG() << "processing: "
<< OpWithFlags(op, OpPrintingFlags().skipRegions());
if (op->hasTrait<OpTrait::SymbolTable>()) {
SymbolOpInterface symbol = dyn_cast<SymbolOpInterface>(op);
bool symIsHidden = symbolTableIsHidden || !symbol || symbol.isPrivate();
LDBG() << "\tsymbol table: "
<< OpWithFlags(op, OpPrintingFlags().skipRegions())
<< " is hidden: " << symIsHidden;
if (failed(computeLiveness(op, symbolTable, symIsHidden, liveSymbols)))
return failure();
} else {
LDBG() << "\tnon-symbol table: "
<< OpWithFlags(op, OpPrintingFlags().skipRegions());
for (auto ®ion : op->getRegions())
for (auto &block : region.getBlocks())
for (Operation &op : block)
if (op.getNumRegions())
worklist.push_back(&op);
}
Operation *parentOp = op->getParentOp();
if (!parentOp->hasTrait<OpTrait::SymbolTable>())
continue;
std::optional<SymbolTable::UseRange> uses = SymbolTable::getSymbolUses(op);
if (!uses) {
return op->emitError()
<< "operation contains potentially unknown symbol table, meaning "
<< "that we can't reliable compute symbol uses";
}
SmallVector<Operation *, 4> resolvedSymbols;
LDBG() << "uses of " << OpWithFlags(op, OpPrintingFlags().skipRegions());
for (const SymbolTable::SymbolUse &use : *uses) {
LDBG() << "\tuse: " << use.getUser();
resolvedSymbols.clear();
if (failed(symbolTable.lookupSymbolIn(parentOp, use.getSymbolRef(),
resolvedSymbols)))
continue;
LDBG() << "\t\tresolved symbols: "
<< llvm::interleaved(resolvedSymbols, ", ");
for (Operation *resolvedSymbol : resolvedSymbols)
if (liveSymbols.insert(resolvedSymbol).second)
worklist.push_back(resolvedSymbol);
}
}
return success();
}
std::unique_ptr<Pass> mlir::createSymbolDCEPass() {
return std::make_unique<SymbolDCE>();
}