1ac91e673SRiver Riddle //===- TestSymbolUses.cpp - Pass to test symbol uselists ------------------===// 2ac91e673SRiver Riddle // 356222a06SMehdi Amini // Part of the MLIR 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" 10ac91e673SRiver Riddle #include "mlir/IR/Function.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. 18ac91e673SRiver Riddle struct SymbolUsesPass : public ModulePass<SymbolUsesPass> { 19*6fca03f0SRiver Riddle WalkResult operateOnSymbol(Operation *symbol, Operation *module, 20*6fca03f0SRiver Riddle SmallVectorImpl<FuncOp> &deadFunctions) { 21ac91e673SRiver Riddle // Test computing uses on a non symboltable op. 22b3a6ae83SRiver Riddle Optional<SymbolTable::UseRange> symbolUses = 23*6fca03f0SRiver Riddle SymbolTable::getSymbolUses(symbol); 24b3a6ae83SRiver Riddle 25b3a6ae83SRiver Riddle // Test the conservative failure case. 26b3a6ae83SRiver Riddle if (!symbolUses) { 27*6fca03f0SRiver Riddle symbol->emitRemark() 28*6fca03f0SRiver Riddle << "symbol contains an unknown nested operation that " 29*6fca03f0SRiver Riddle "'may' define a new symbol table"; 30*6fca03f0SRiver Riddle return WalkResult::interrupt(); 31b3a6ae83SRiver Riddle } 32b3a6ae83SRiver Riddle if (unsigned numUses = llvm::size(*symbolUses)) 33*6fca03f0SRiver Riddle symbol->emitRemark() << "symbol contains " << numUses 34ac91e673SRiver Riddle << " nested references"; 35ac91e673SRiver Riddle 36b3a6ae83SRiver Riddle // Test the functionality of symbolKnownUseEmpty. 37*6fca03f0SRiver Riddle if (SymbolTable::symbolKnownUseEmpty(symbol, module)) { 38*6fca03f0SRiver Riddle FuncOp funcSymbol = dyn_cast<FuncOp>(symbol); 39*6fca03f0SRiver Riddle if (funcSymbol && funcSymbol.isExternal()) 40*6fca03f0SRiver Riddle deadFunctions.push_back(funcSymbol); 41*6fca03f0SRiver Riddle 42*6fca03f0SRiver Riddle symbol->emitRemark() << "symbol has no uses"; 43*6fca03f0SRiver Riddle return WalkResult::advance(); 44ac91e673SRiver Riddle } 45ac91e673SRiver Riddle 46b3a6ae83SRiver Riddle // Test the functionality of getSymbolUses. 47*6fca03f0SRiver Riddle symbolUses = SymbolTable::getSymbolUses(symbol, module); 48b3a6ae83SRiver Riddle assert(symbolUses.hasValue() && "expected no unknown operations"); 49b3a6ae83SRiver Riddle for (SymbolTable::SymbolUse symbolUse : *symbolUses) { 50*6fca03f0SRiver Riddle // Check that we can resolve back to our symbol. 51*6fca03f0SRiver Riddle if (Operation *op = SymbolTable::lookupNearestSymbolFrom( 52*6fca03f0SRiver Riddle symbolUse.getUser()->getParentOp(), symbolUse.getSymbolRef())) { 53ac91e673SRiver Riddle symbolUse.getUser()->emitRemark() 54*6fca03f0SRiver Riddle << "found use of symbol : " << symbolUse.getSymbolRef() << " : " 55*6fca03f0SRiver Riddle << symbol->getAttr(SymbolTable::getSymbolAttrName()); 56b3a6ae83SRiver Riddle } 57*6fca03f0SRiver Riddle } 58*6fca03f0SRiver Riddle symbol->emitRemark() << "symbol has " << llvm::size(*symbolUses) << " uses"; 59*6fca03f0SRiver Riddle return WalkResult::advance(); 60ac91e673SRiver Riddle } 61b8cd0c14STres Popp 62*6fca03f0SRiver Riddle void runOnModule() override { 63*6fca03f0SRiver Riddle auto module = getModule(); 64*6fca03f0SRiver Riddle 65*6fca03f0SRiver Riddle // Walk nested symbols. 66*6fca03f0SRiver Riddle SmallVector<FuncOp, 4> deadFunctions; 67*6fca03f0SRiver Riddle module.getBodyRegion().walk([&](Operation *nestedOp) { 68*6fca03f0SRiver Riddle if (SymbolTable::isSymbol(nestedOp)) 69*6fca03f0SRiver Riddle return operateOnSymbol(nestedOp, module, deadFunctions); 70*6fca03f0SRiver Riddle return WalkResult::advance(); 71*6fca03f0SRiver Riddle }); 72*6fca03f0SRiver Riddle 73*6fca03f0SRiver Riddle for (Operation *op : deadFunctions) { 74b8cd0c14STres Popp // In order to test the SymbolTable::erase method, also erase completely 75b8cd0c14STres Popp // useless functions. 76b8cd0c14STres Popp SymbolTable table(module); 77*6fca03f0SRiver Riddle auto name = SymbolTable::getSymbolName(op); 78*6fca03f0SRiver Riddle assert(table.lookup(name) && "expected no unknown operations"); 79*6fca03f0SRiver Riddle table.erase(op); 80*6fca03f0SRiver Riddle assert(!table.lookup(name) && 81b8cd0c14STres Popp "expected erased operation to be unknown now"); 82*6fca03f0SRiver Riddle module.emitRemark() << name << " function successfully erased"; 83b8cd0c14STres Popp } 84ac91e673SRiver Riddle } 85ac91e673SRiver Riddle }; 86ef43b565SRiver Riddle 87ef43b565SRiver Riddle /// This is a symbol test pass that tests the symbol use replacement 88ef43b565SRiver Riddle /// functionality provided by the symbol table. 89ef43b565SRiver Riddle struct SymbolReplacementPass : public ModulePass<SymbolReplacementPass> { 90ef43b565SRiver Riddle void runOnModule() override { 91ef43b565SRiver Riddle auto module = getModule(); 92ef43b565SRiver Riddle 93*6fca03f0SRiver Riddle // Walk nested functions and modules. 94*6fca03f0SRiver Riddle module.getBodyRegion().walk([&](Operation *nestedOp) { 95*6fca03f0SRiver Riddle StringAttr newName = nestedOp->getAttrOfType<StringAttr>("sym.new_name"); 96ef43b565SRiver Riddle if (!newName) 97*6fca03f0SRiver Riddle return; 98*6fca03f0SRiver Riddle if (succeeded(SymbolTable::replaceAllSymbolUses( 99*6fca03f0SRiver Riddle nestedOp, newName.getValue(), module))) 100*6fca03f0SRiver Riddle SymbolTable::setSymbolName(nestedOp, newName.getValue()); 101*6fca03f0SRiver Riddle }); 102ef43b565SRiver Riddle } 103ef43b565SRiver Riddle }; 104ac91e673SRiver Riddle } // end anonymous namespace 105ac91e673SRiver Riddle 106ac91e673SRiver Riddle static PassRegistration<SymbolUsesPass> pass("test-symbol-uses", 107ac91e673SRiver Riddle "Test detection of symbol uses"); 108ef43b565SRiver Riddle 109ef43b565SRiver Riddle static PassRegistration<SymbolReplacementPass> 110ef43b565SRiver Riddle rauwPass("test-symbol-rauw", "Test replacement of symbol uses"); 111