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>> { 20b5e22e6dSMehdi Amini StringRef getArgument() const final { return "test-symbol-uses"; } 21b5e22e6dSMehdi Amini StringRef getDescription() const final { 22b5e22e6dSMehdi Amini return "Test detection of symbol uses"; 23b5e22e6dSMehdi Amini } 24ab9e5598SRiver Riddle WalkResult operateOnSymbol(Operation *symbol, ModuleOp module, 256fca03f0SRiver Riddle SmallVectorImpl<FuncOp> &deadFunctions) { 26ac91e673SRiver Riddle // Test computing uses on a non symboltable op. 27b3a6ae83SRiver Riddle Optional<SymbolTable::UseRange> symbolUses = 286fca03f0SRiver Riddle SymbolTable::getSymbolUses(symbol); 29b3a6ae83SRiver Riddle 30b3a6ae83SRiver Riddle // Test the conservative failure case. 31b3a6ae83SRiver Riddle if (!symbolUses) { 326fca03f0SRiver Riddle symbol->emitRemark() 336fca03f0SRiver Riddle << "symbol contains an unknown nested operation that " 346fca03f0SRiver Riddle "'may' define a new symbol table"; 356fca03f0SRiver Riddle return WalkResult::interrupt(); 36b3a6ae83SRiver Riddle } 37b3a6ae83SRiver Riddle if (unsigned numUses = llvm::size(*symbolUses)) 386fca03f0SRiver Riddle symbol->emitRemark() << "symbol contains " << numUses 39ac91e673SRiver Riddle << " nested references"; 40ac91e673SRiver Riddle 41b3a6ae83SRiver Riddle // Test the functionality of symbolKnownUseEmpty. 42ab9e5598SRiver Riddle if (SymbolTable::symbolKnownUseEmpty(symbol, &module.getBodyRegion())) { 436fca03f0SRiver Riddle FuncOp funcSymbol = dyn_cast<FuncOp>(symbol); 446fca03f0SRiver Riddle if (funcSymbol && funcSymbol.isExternal()) 456fca03f0SRiver Riddle deadFunctions.push_back(funcSymbol); 466fca03f0SRiver Riddle 476fca03f0SRiver Riddle symbol->emitRemark() << "symbol has no uses"; 486fca03f0SRiver Riddle return WalkResult::advance(); 49ac91e673SRiver Riddle } 50ac91e673SRiver Riddle 51b3a6ae83SRiver Riddle // Test the functionality of getSymbolUses. 52ab9e5598SRiver Riddle symbolUses = SymbolTable::getSymbolUses(symbol, &module.getBodyRegion()); 53b3a6ae83SRiver Riddle assert(symbolUses.hasValue() && "expected no unknown operations"); 54b3a6ae83SRiver Riddle for (SymbolTable::SymbolUse symbolUse : *symbolUses) { 556fca03f0SRiver Riddle // Check that we can resolve back to our symbol. 5603edd6d6SRiver Riddle if (SymbolTable::lookupNearestSymbolFrom( 576fca03f0SRiver Riddle symbolUse.getUser()->getParentOp(), symbolUse.getSymbolRef())) { 58ac91e673SRiver Riddle symbolUse.getUser()->emitRemark() 596fca03f0SRiver Riddle << "found use of symbol : " << symbolUse.getSymbolRef() << " : " 606fca03f0SRiver Riddle << symbol->getAttr(SymbolTable::getSymbolAttrName()); 61b3a6ae83SRiver Riddle } 626fca03f0SRiver Riddle } 636fca03f0SRiver Riddle symbol->emitRemark() << "symbol has " << llvm::size(*symbolUses) << " uses"; 646fca03f0SRiver Riddle return WalkResult::advance(); 65ac91e673SRiver Riddle } 66b8cd0c14STres Popp 67722f909fSRiver Riddle void runOnOperation() override { 68722f909fSRiver Riddle auto module = getOperation(); 696fca03f0SRiver Riddle 706fca03f0SRiver Riddle // Walk nested symbols. 716fca03f0SRiver Riddle SmallVector<FuncOp, 4> deadFunctions; 726fca03f0SRiver Riddle module.getBodyRegion().walk([&](Operation *nestedOp) { 737c221a7dSRiver Riddle if (isa<SymbolOpInterface>(nestedOp)) 746fca03f0SRiver Riddle return operateOnSymbol(nestedOp, module, deadFunctions); 756fca03f0SRiver Riddle return WalkResult::advance(); 766fca03f0SRiver Riddle }); 776fca03f0SRiver Riddle 78ab9e5598SRiver Riddle SymbolTable table(module); 796fca03f0SRiver Riddle for (Operation *op : deadFunctions) { 80b8cd0c14STres Popp // In order to test the SymbolTable::erase method, also erase completely 81b8cd0c14STres Popp // useless functions. 826fca03f0SRiver Riddle auto name = SymbolTable::getSymbolName(op); 836fca03f0SRiver Riddle assert(table.lookup(name) && "expected no unknown operations"); 846fca03f0SRiver Riddle table.erase(op); 856fca03f0SRiver Riddle assert(!table.lookup(name) && 86b8cd0c14STres Popp "expected erased operation to be unknown now"); 87*41d4aa7dSChris Lattner module.emitRemark() << name.getValue() << " function successfully erased"; 88b8cd0c14STres Popp } 89ac91e673SRiver Riddle } 90ac91e673SRiver Riddle }; 91ef43b565SRiver Riddle 92ef43b565SRiver Riddle /// This is a symbol test pass that tests the symbol use replacement 93ef43b565SRiver Riddle /// functionality provided by the symbol table. 94722f909fSRiver Riddle struct SymbolReplacementPass 9580aca1eaSRiver Riddle : public PassWrapper<SymbolReplacementPass, OperationPass<ModuleOp>> { 96b5e22e6dSMehdi Amini StringRef getArgument() const final { return "test-symbol-rauw"; } 97b5e22e6dSMehdi Amini StringRef getDescription() const final { 98b5e22e6dSMehdi Amini return "Test replacement of symbol uses"; 99b5e22e6dSMehdi Amini } 100722f909fSRiver Riddle void runOnOperation() override { 1014a7aed4eSRiver Riddle ModuleOp module = getOperation(); 102ef43b565SRiver Riddle 1034a7aed4eSRiver Riddle // Don't try to replace if we can't collect symbol uses. 1044a7aed4eSRiver Riddle if (!SymbolTable::getSymbolUses(&module.getBodyRegion())) 1054a7aed4eSRiver Riddle return; 1064a7aed4eSRiver Riddle 1074a7aed4eSRiver Riddle SymbolTableCollection symbolTable; 1084a7aed4eSRiver Riddle SymbolUserMap symbolUsers(symbolTable, module); 1096fca03f0SRiver Riddle module.getBodyRegion().walk([&](Operation *nestedOp) { 1106fca03f0SRiver Riddle StringAttr newName = nestedOp->getAttrOfType<StringAttr>("sym.new_name"); 111ef43b565SRiver Riddle if (!newName) 1126fca03f0SRiver Riddle return; 113*41d4aa7dSChris Lattner symbolUsers.replaceAllUsesWith(nestedOp, newName); 114*41d4aa7dSChris Lattner SymbolTable::setSymbolName(nestedOp, newName); 1156fca03f0SRiver Riddle }); 116ef43b565SRiver Riddle } 117ef43b565SRiver Riddle }; 118ac91e673SRiver Riddle } // end anonymous namespace 119ac91e673SRiver Riddle 120c6477050SMehdi Amini namespace mlir { 121c6477050SMehdi Amini void registerSymbolTestPasses() { 122b5e22e6dSMehdi Amini PassRegistration<SymbolUsesPass>(); 123ef43b565SRiver Riddle 124b5e22e6dSMehdi Amini PassRegistration<SymbolReplacementPass>(); 125c6477050SMehdi Amini } 126c6477050SMehdi Amini } // namespace mlir 127