1ac91e673SRiver Riddle //===- TestSymbolUses.cpp - Pass to test symbol uselists ------------------===//
2ac91e673SRiver Riddle //
330857107SMehdi Amini // Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
456222a06SMehdi Amini // See https://llvm.org/LICENSE.txt for license information.
556222a06SMehdi Amini // SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6ac91e673SRiver Riddle //
756222a06SMehdi Amini //===----------------------------------------------------------------------===//
8ac91e673SRiver Riddle
9b8cd0c14STres Popp #include "TestDialect.h"
1065fcddffSRiver Riddle #include "mlir/IR/BuiltinOps.h"
11ac91e673SRiver Riddle #include "mlir/Pass/Pass.h"
12ac91e673SRiver Riddle
13ac91e673SRiver Riddle using namespace mlir;
14ac91e673SRiver Riddle
15ac91e673SRiver Riddle namespace {
16ac91e673SRiver Riddle /// This is a symbol test pass that tests the symbol uselist functionality
17b8cd0c14STres Popp /// provided by the symbol table along with erasing from the symbol table.
1880aca1eaSRiver Riddle struct SymbolUsesPass
1980aca1eaSRiver Riddle : public PassWrapper<SymbolUsesPass, OperationPass<ModuleOp>> {
MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID__anon4969bba80111::SymbolUsesPass205e50dd04SRiver Riddle MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(SymbolUsesPass)
215e50dd04SRiver Riddle
22b5e22e6dSMehdi Amini StringRef getArgument() const final { return "test-symbol-uses"; }
getDescription__anon4969bba80111::SymbolUsesPass23b5e22e6dSMehdi Amini StringRef getDescription() const final {
24b5e22e6dSMehdi Amini return "Test detection of symbol uses";
25b5e22e6dSMehdi Amini }
operateOnSymbol__anon4969bba80111::SymbolUsesPass26ab9e5598SRiver Riddle WalkResult operateOnSymbol(Operation *symbol, ModuleOp module,
2758ceae95SRiver Riddle SmallVectorImpl<func::FuncOp> &deadFunctions) {
28ac91e673SRiver Riddle // Test computing uses on a non symboltable op.
29b3a6ae83SRiver Riddle Optional<SymbolTable::UseRange> symbolUses =
306fca03f0SRiver Riddle SymbolTable::getSymbolUses(symbol);
31b3a6ae83SRiver Riddle
32b3a6ae83SRiver Riddle // Test the conservative failure case.
33b3a6ae83SRiver Riddle if (!symbolUses) {
346fca03f0SRiver Riddle symbol->emitRemark()
356fca03f0SRiver Riddle << "symbol contains an unknown nested operation that "
366fca03f0SRiver Riddle "'may' define a new symbol table";
376fca03f0SRiver Riddle return WalkResult::interrupt();
38b3a6ae83SRiver Riddle }
39b3a6ae83SRiver Riddle if (unsigned numUses = llvm::size(*symbolUses))
406fca03f0SRiver Riddle symbol->emitRemark() << "symbol contains " << numUses
41ac91e673SRiver Riddle << " nested references";
42ac91e673SRiver Riddle
43b3a6ae83SRiver Riddle // Test the functionality of symbolKnownUseEmpty.
44ab9e5598SRiver Riddle if (SymbolTable::symbolKnownUseEmpty(symbol, &module.getBodyRegion())) {
4558ceae95SRiver Riddle func::FuncOp funcSymbol = dyn_cast<func::FuncOp>(symbol);
466fca03f0SRiver Riddle if (funcSymbol && funcSymbol.isExternal())
476fca03f0SRiver Riddle deadFunctions.push_back(funcSymbol);
486fca03f0SRiver Riddle
496fca03f0SRiver Riddle symbol->emitRemark() << "symbol has no uses";
506fca03f0SRiver Riddle return WalkResult::advance();
51ac91e673SRiver Riddle }
52ac91e673SRiver Riddle
53b3a6ae83SRiver Riddle // Test the functionality of getSymbolUses.
54ab9e5598SRiver Riddle symbolUses = SymbolTable::getSymbolUses(symbol, &module.getBodyRegion());
55*5413bf1bSKazu Hirata assert(symbolUses && "expected no unknown operations");
56b3a6ae83SRiver Riddle for (SymbolTable::SymbolUse symbolUse : *symbolUses) {
576fca03f0SRiver Riddle // Check that we can resolve back to our symbol.
5803edd6d6SRiver Riddle if (SymbolTable::lookupNearestSymbolFrom(
596fca03f0SRiver Riddle symbolUse.getUser()->getParentOp(), symbolUse.getSymbolRef())) {
60ac91e673SRiver Riddle symbolUse.getUser()->emitRemark()
616fca03f0SRiver Riddle << "found use of symbol : " << symbolUse.getSymbolRef() << " : "
626fca03f0SRiver Riddle << symbol->getAttr(SymbolTable::getSymbolAttrName());
63b3a6ae83SRiver Riddle }
646fca03f0SRiver Riddle }
656fca03f0SRiver Riddle symbol->emitRemark() << "symbol has " << llvm::size(*symbolUses) << " uses";
666fca03f0SRiver Riddle return WalkResult::advance();
67ac91e673SRiver Riddle }
68b8cd0c14STres Popp
runOnOperation__anon4969bba80111::SymbolUsesPass69722f909fSRiver Riddle void runOnOperation() override {
70722f909fSRiver Riddle auto module = getOperation();
716fca03f0SRiver Riddle
726fca03f0SRiver Riddle // Walk nested symbols.
7358ceae95SRiver Riddle SmallVector<func::FuncOp, 4> deadFunctions;
746fca03f0SRiver Riddle module.getBodyRegion().walk([&](Operation *nestedOp) {
757c221a7dSRiver Riddle if (isa<SymbolOpInterface>(nestedOp))
766fca03f0SRiver Riddle return operateOnSymbol(nestedOp, module, deadFunctions);
776fca03f0SRiver Riddle return WalkResult::advance();
786fca03f0SRiver Riddle });
796fca03f0SRiver Riddle
80ab9e5598SRiver Riddle SymbolTable table(module);
816fca03f0SRiver Riddle for (Operation *op : deadFunctions) {
82b8cd0c14STres Popp // In order to test the SymbolTable::erase method, also erase completely
83b8cd0c14STres Popp // useless functions.
846fca03f0SRiver Riddle auto name = SymbolTable::getSymbolName(op);
856fca03f0SRiver Riddle assert(table.lookup(name) && "expected no unknown operations");
866fca03f0SRiver Riddle table.erase(op);
876fca03f0SRiver Riddle assert(!table.lookup(name) &&
88b8cd0c14STres Popp "expected erased operation to be unknown now");
8941d4aa7dSChris Lattner module.emitRemark() << name.getValue() << " function successfully erased";
90b8cd0c14STres Popp }
91ac91e673SRiver Riddle }
92ac91e673SRiver Riddle };
93ef43b565SRiver Riddle
94ef43b565SRiver Riddle /// This is a symbol test pass that tests the symbol use replacement
95ef43b565SRiver Riddle /// functionality provided by the symbol table.
96722f909fSRiver Riddle struct SymbolReplacementPass
9780aca1eaSRiver Riddle : public PassWrapper<SymbolReplacementPass, OperationPass<ModuleOp>> {
MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID__anon4969bba80111::SymbolReplacementPass985e50dd04SRiver Riddle MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(SymbolReplacementPass)
995e50dd04SRiver Riddle
100b5e22e6dSMehdi Amini StringRef getArgument() const final { return "test-symbol-rauw"; }
getDescription__anon4969bba80111::SymbolReplacementPass101b5e22e6dSMehdi Amini StringRef getDescription() const final {
102b5e22e6dSMehdi Amini return "Test replacement of symbol uses";
103b5e22e6dSMehdi Amini }
runOnOperation__anon4969bba80111::SymbolReplacementPass104722f909fSRiver Riddle void runOnOperation() override {
1054a7aed4eSRiver Riddle ModuleOp module = getOperation();
106ef43b565SRiver Riddle
1074a7aed4eSRiver Riddle // Don't try to replace if we can't collect symbol uses.
1084a7aed4eSRiver Riddle if (!SymbolTable::getSymbolUses(&module.getBodyRegion()))
1094a7aed4eSRiver Riddle return;
1104a7aed4eSRiver Riddle
1114a7aed4eSRiver Riddle SymbolTableCollection symbolTable;
1124a7aed4eSRiver Riddle SymbolUserMap symbolUsers(symbolTable, module);
1136fca03f0SRiver Riddle module.getBodyRegion().walk([&](Operation *nestedOp) {
1146fca03f0SRiver Riddle StringAttr newName = nestedOp->getAttrOfType<StringAttr>("sym.new_name");
115ef43b565SRiver Riddle if (!newName)
1166fca03f0SRiver Riddle return;
11741d4aa7dSChris Lattner symbolUsers.replaceAllUsesWith(nestedOp, newName);
11841d4aa7dSChris Lattner SymbolTable::setSymbolName(nestedOp, newName);
1196fca03f0SRiver Riddle });
120ef43b565SRiver Riddle }
121ef43b565SRiver Riddle };
122be0a7e9fSMehdi Amini } // namespace
123ac91e673SRiver Riddle
124c6477050SMehdi Amini namespace mlir {
registerSymbolTestPasses()125c6477050SMehdi Amini void registerSymbolTestPasses() {
126b5e22e6dSMehdi Amini PassRegistration<SymbolUsesPass>();
127ef43b565SRiver Riddle
128b5e22e6dSMehdi Amini PassRegistration<SymbolReplacementPass>();
129c6477050SMehdi Amini }
130c6477050SMehdi Amini } // namespace mlir
131