49#include "llvm/ADT/STLExtras.h"
50#include "llvm/Support/Debug.h"
51#include "llvm/Support/DebugLog.h"
58#define DEBUG_TYPE "remove-dead-values"
61#define GEN_PASS_DEF_REMOVEDEADVALUESPASS
62#include "mlir/Transforms/Passes.h.inc"
78struct FunctionToCleanUp {
79 FunctionOpInterface funcOp;
80 BitVector nonLiveArgs;
81 BitVector nonLiveRets;
84struct ResultsToCleanup {
89struct OperandsToCleanup {
93 Operation *callee =
nullptr;
96 bool replaceWithPoison =
false;
99struct BlockArgsToCleanup {
101 BitVector nonLiveArgs;
104struct SuccessorOperandsToCleanup {
105 BranchOpInterface branch;
106 unsigned successorIndex;
107 BitVector nonLiveOperands;
110struct RDVFinalCleanupList {
111 SmallVector<Operation *> operations;
112 SmallVector<FunctionToCleanUp> functions;
113 SmallVector<OperandsToCleanup> operands;
114 SmallVector<ResultsToCleanup> results;
115 SmallVector<BlockArgsToCleanup> blocks;
116 SmallVector<SuccessorOperandsToCleanup> successorOperands;
125 for (
Value value : values) {
126 if (nonLiveSet.contains(value)) {
127 LDBG() <<
"Value " << value <<
" is already marked non-live (dead)";
133 LDBG() <<
"Value " << value
134 <<
" has no liveness info, conservatively considered live";
138 LDBG() <<
"Value " << value <<
" is live according to liveness analysis";
141 LDBG() <<
"Value " << value <<
" is dead according to liveness analysis";
150 BitVector lives(values.size(),
true);
152 for (
auto [
index, value] : llvm::enumerate(values)) {
153 if (nonLiveSet.contains(value)) {
155 LDBG() <<
"Value " << value
156 <<
" is already marked non-live (dead) at index " <<
index;
168 LDBG() <<
"Value " << value <<
" at index " <<
index
169 <<
" has no liveness info, conservatively considered live";
174 LDBG() <<
"Value " << value <<
" at index " <<
index
175 <<
" is dead according to liveness analysis";
177 LDBG() <<
"Value " << value <<
" at index " <<
index
178 <<
" is live according to liveness analysis";
189 const BitVector &nonLive) {
190 for (
auto [
index,
result] : llvm::enumerate(range)) {
193 nonLiveSet.insert(
result);
194 LDBG() <<
"Marking value " <<
result <<
" as non-live (dead) at index "
204 "expected the number of results in `op` and the size of `toErase` to "
206 for (
auto idx : toErase.set_bits())
226 RDVFinalCleanupList &cl) {
230 bool hasDeadOperand =
231 markLives(op->
getOperands(), nonLiveSet, la).flip().any();
232 if (hasDeadOperand) {
233 LDBG() <<
"Simple op has dead operands, so the op must be dead: "
236 assert(!hasLive(op->
getResults(), nonLiveSet, la) &&
237 "expected the op to have no live results");
238 cl.operations.push_back(op);
239 collectNonLiveValues(nonLiveSet, op->
getResults(),
245 LDBG() <<
"Simple op is not memory effect free or has live results, "
253 <<
"Simple op has all dead results and is memory effect free, scheduling "
256 cl.operations.push_back(op);
257 collectNonLiveValues(nonLiveSet, op->
getResults(),
272static void processFuncOp(FunctionOpInterface funcOp,
275 RDVFinalCleanupList &cl) {
276 LDBG() <<
"Processing function op: "
280 if (funcOp.isPublic() || funcOp.isExternal() ||
282 LDBG() <<
"Function is public, external, or has unknown users, skipping: "
283 << funcOp.getOperation()->getName();
287 if (!llvm::all_of(users, llvm::IsaPred<CallOpInterface>)) {
295 BitVector nonLiveArgs = markLives(arguments, nonLiveSet, la);
296 nonLiveArgs = nonLiveArgs.flip();
299 for (
auto [
index, arg] : llvm::enumerate(arguments))
300 if (arg && nonLiveArgs[
index])
301 nonLiveSet.insert(arg);
311 cl.operands.push_back({callOp, BitVector(callOp->getNumOperands(),
false),
312 funcOp.getOperation()});
338 size_t numReturns = funcOp.getNumResults();
339 BitVector nonLiveRets(numReturns,
true);
343 BitVector liveCallRets = markLives(
344 cast<CallOpInterface>(callOp).getForwardedResults(), nonLiveSet, la);
345 nonLiveRets &= liveCallRets.flip();
351 for (
Block &block : funcOp.getBlocks()) {
352 Operation *returnOp = block.getTerminator();
356 cl.operands.push_back({returnOp, nonLiveRets});
360 cl.functions.push_back({funcOp, nonLiveArgs, nonLiveRets});
370 cast<CallOpInterface>(callOp).getForwardedResults();
371 BitVector nonLiveCallResults(callOp->getNumResults(),
false);
372 for (
int index : nonLiveRets.set_bits())
373 nonLiveCallResults.set(forwardedResults[
index].getResultNumber());
374 cl.results.push_back({callOp, nonLiveCallResults});
375 collectNonLiveValues(nonLiveSet, callOp->getResults(), nonLiveCallResults);
398static void processRegionBranchOp(RegionBranchOpInterface regionBranchOp,
401 RDVFinalCleanupList &cl) {
402 LDBG() <<
"Processing region branch op: "
412 !hasLive(regionBranchOp->getResults(), nonLiveSet, la)) {
413 cl.operations.push_back(regionBranchOp.getOperation());
437 regionBranchOp.getSuccessorOperandInputMapping(operandToSuccessorInputs);
440 for (
auto [opOperand, successorInputs] : operandToSuccessorInputs) {
443 auto markOperandDead = [&opOperand = opOperand, &deadOperandsPerOp]() {
448 BitVector &deadOperands =
450 .try_emplace(opOperand->getOwner(),
451 opOperand->getOwner()->getNumOperands(),
false)
453 deadOperands.set(opOperand->getOperandNumber());
457 if (!hasLive(opOperand->get(), nonLiveSet, la)) {
464 if (!hasLive(successorInputs, nonLiveSet, la))
468 for (
auto [op, deadOperands] : deadOperandsPerOp) {
469 cl.operands.push_back(
470 {op, deadOperands,
nullptr,
true});
488 RDVFinalCleanupList &cl) {
489 LDBG() <<
"Processing branch op: " << *branchOp;
492 BitVector deadNonForwardedOperands =
493 markLives(branchOp->getOperands(), nonLiveSet, la).flip();
494 unsigned numSuccessors = branchOp->getNumSuccessors();
495 for (
unsigned succIdx = 0; succIdx < numSuccessors; ++succIdx) {
497 branchOp.getSuccessorOperands(succIdx);
500 deadNonForwardedOperands[opOperand.getOperandNumber()] =
false;
502 if (deadNonForwardedOperands.any()) {
503 cl.operations.push_back(branchOp.getOperation());
507 for (
unsigned succIdx = 0; succIdx < numSuccessors; ++succIdx) {
512 branchOp.getSuccessorOperands(succIdx);
514 for (
unsigned operandIdx = 0; operandIdx < successorOperands.
size();
516 operandValues.push_back(successorOperands[operandIdx]);
520 BitVector successorNonLive =
521 markLives(operandValues, nonLiveSet, la).flip();
522 collectNonLiveValues(nonLiveSet, successorBlock->
getArguments(),
526 cl.blocks.push_back({successorBlock, successorNonLive});
527 cl.successorOperands.push_back({branchOp, succIdx, successorNonLive});
536 return ub::PoisonOp::create(
b, value.
getLoc(), value.
getType()).getResult();
542 void notifyOperationErased(Operation *op)
override {
543 if (
auto poisonOp = dyn_cast<ub::PoisonOp>(op))
544 poisonOps.erase(poisonOp);
546 void notifyOperationInserted(Operation *op,
547 OpBuilder::InsertPoint previous)
override {
548 if (
auto poisonOp = dyn_cast<ub::PoisonOp>(op))
549 poisonOps.insert(poisonOp);
557static void cleanUpDeadVals(
MLIRContext *ctx, RDVFinalCleanupList &list) {
558 LDBG() <<
"Starting cleanup of dead values...";
563 TrackingListener listener;
569 LDBG() <<
"Replacing dead operands with poison in " << list.operands.size()
571 for (OperandsToCleanup &o : list.operands) {
572 if (!o.replaceWithPoison || !o.nonLive.any())
575 os <<
"Replacing non-live operands [";
576 llvm::interleaveComma(o.nonLive.set_bits(), os);
577 os <<
"] with poison in operation: "
582 for (
auto deadIdx : o.nonLive.set_bits()) {
584 assert(operand &&
"expected non-null operand for poison replacement");
585 o.op->
setOperand(deadIdx, createPoisonedValue(rewriter, operand));
591 LDBG() <<
"Cleaning up " << list.blocks.size() <<
" block argument lists";
592 for (
auto &
b : list.blocks) {
594 if (
b.b->getNumArguments() !=
b.nonLiveArgs.size())
597 os <<
"Erasing non-live arguments [";
598 llvm::interleaveComma(
b.nonLiveArgs.set_bits(), os);
599 os <<
"] from block #" <<
b.b->computeBlockNumber() <<
" in region #"
600 <<
b.b->getParent()->getRegionNumber() <<
" of operation "
606 for (
int i =
b.nonLiveArgs.size() - 1; i >= 0; --i) {
607 if (!
b.nonLiveArgs[i])
609 b.b->getArgument(i).dropAllUses();
610 b.b->eraseArgument(i);
615 LDBG() <<
"Cleaning up " << list.successorOperands.size()
616 <<
" successor operand lists";
617 for (
auto &op : list.successorOperands) {
619 op.branch.getSuccessorOperands(op.successorIndex);
621 if (successorOperands.
size() != op.nonLiveOperands.size())
624 os <<
"Erasing non-live successor operands [";
625 llvm::interleaveComma(op.nonLiveOperands.set_bits(), os);
626 os <<
"] from successor " << op.successorIndex <<
" of branch: "
631 for (
int i = successorOperands.
size() - 1; i >= 0; --i) {
632 if (!op.nonLiveOperands[i])
634 successorOperands.
erase(i);
639 LDBG() <<
"Cleaning up " << list.functions.size() <<
" functions";
644 for (
auto &f : list.functions) {
645 LDBG() <<
"Cleaning up function: " << f.funcOp.getName() <<
" ("
646 << f.funcOp.getOperation() <<
")";
648 os <<
" Erasing non-live arguments [";
649 llvm::interleaveComma(f.nonLiveArgs.set_bits(), os);
651 os <<
" Erasing non-live return values [";
652 llvm::interleaveComma(f.nonLiveRets.set_bits(), os);
656 for (
auto deadIdx : f.nonLiveArgs.set_bits())
657 f.funcOp.getArgument(deadIdx).dropAllUses();
661 if (succeeded(f.funcOp.eraseArguments(f.nonLiveArgs))) {
663 if (f.nonLiveArgs.any())
664 erasedFuncArgs.try_emplace(f.funcOp.getOperation(), f.nonLiveArgs);
666 LDBG() <<
"Failed to erase arguments for function: "
667 << f.funcOp.getName();
669 (
void)f.funcOp.eraseResults(f.nonLiveRets);
673 LDBG() <<
"Cleaning up " << list.operands.size() <<
" operand lists";
674 for (OperandsToCleanup &o : list.operands) {
675 if (o.replaceWithPoison)
680 bool handledAsCall =
false;
681 if (o.callee && isa<CallOpInterface>(o.op)) {
682 auto call = cast<CallOpInterface>(o.op);
683 auto it = erasedFuncArgs.find(o.callee);
684 if (it != erasedFuncArgs.end()) {
685 const BitVector &deadArgIdxs = it->second;
689 for (
unsigned argIdx : llvm::reverse(deadArgIdxs.set_bits()))
695 if (o.nonLive.any()) {
697 int operandOffset = call.getArgOperands().getBeginOperandIndex();
698 for (
int argIdx : deadArgIdxs.set_bits()) {
699 int operandNumber = operandOffset + argIdx;
700 if (operandNumber <
static_cast<int>(o.nonLive.size()))
701 o.nonLive.reset(operandNumber);
704 handledAsCall =
true;
711 if (!handledAsCall && o.nonLive.any()) {
713 os <<
"Erasing non-live operands [";
714 llvm::interleaveComma(o.nonLive.set_bits(), os);
715 os <<
"] from operation: "
724 LDBG() <<
"Cleaning up " << list.results.size() <<
" result lists";
725 for (
auto &r : list.results) {
727 os <<
"Erasing non-live results [";
728 llvm::interleaveComma(r.nonLive.set_bits(), os);
729 os <<
"] from operation: "
733 dropUsesAndEraseResults(rewriter, r.op, r.nonLive);
737 LDBG() <<
"Cleaning up " << list.operations.size() <<
" operations";
739 LDBG() <<
"Erasing operation: "
745 ub::UnreachableOp::create(rewriter, op->
getLoc());
758 for (
Value opResult : opResults) {
761 if (opResult.use_empty())
765 Value poisonedValue = createPoisonedValue(rewriter, opResult);
774 for (ub::PoisonOp poisonOp : listener.poisonOps) {
775 if (poisonOp.use_empty())
779 LDBG() <<
"Finished cleanup of dead values";
782struct RemoveDeadValues
784 using impl::RemoveDeadValuesPassBase<
785 RemoveDeadValues>::RemoveDeadValuesPassBase;
786 void runOnOperation()
override;
790void RemoveDeadValues::runOnOperation() {
791 auto &la = getAnalysis<RunLivenessAnalysis>();
792 Operation *root = getOperation();
798 SymbolTableCollection symbolTableCollection;
799 SymbolUserMap symbolUserMap(symbolTableCollection, root);
807 RDVFinalCleanupList finalCleanupList;
809 root->
walk([&](Operation *op) {
815 if (
auto funcOp = dyn_cast<FunctionOpInterface>(op)) {
816 processFuncOp(funcOp, symbolUserMap, la, deadVals, finalCleanupList);
817 }
else if (
auto regionBranchOp = dyn_cast<RegionBranchOpInterface>(op)) {
818 processRegionBranchOp(regionBranchOp, la, deadVals, finalCleanupList);
819 }
else if (
auto branchOp = dyn_cast<BranchOpInterface>(op)) {
820 processBranchOp(branchOp, la, deadVals, finalCleanupList);
821 }
else if (op->
hasTrait<::mlir::OpTrait::IsTerminator>()) {
824 }
else if (isa<CallOpInterface>(op)) {
828 processSimpleOp(op, la, deadVals, finalCleanupList);
833 cleanUpDeadVals(context, finalCleanupList);
839 RewritePatternSet owningPatterns(context);
841 root->
walk([&](RegionBranchOpInterface regionBranchOp) {
842 Operation *op = regionBranchOp.getOperation();
846 if (populatedPatterns.insert(*info).second)
847 info->getCanonicalizationPatterns(owningPatterns, context);
849 FrozenRewritePatternSet patterns(std::move(owningPatterns));
854 SmallVector<Operation *> opsToCanonicalize;
855 region.walk([&](RegionBranchOpInterface regionBranchOp) {
856 opsToCanonicalize.push_back(regionBranchOp.getOperation());
860 GreedyRewriteConfig config;
863 root->
emitError(
"greedy pattern rewrite failed to converge");
864 return signalPassFailure();
Block represents an ordered list of Operations.
BlockArgListType getArguments()
Block * getSuccessor(unsigned i)
GreedyRewriteConfig & setScope(Region *scope)
This class coordinates rewriting a piece of IR outside of a pattern rewrite, providing a way to keep ...
MLIRContext is the top-level object for a collection of MLIR operations.
This class provides a mutable adaptor for a range of operands.
void erase(unsigned subStart, unsigned subLen=1)
Erase the operands within the given sub-range.
This class helps build Operations.
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
This class represents an operand of an operation.
Set of flags used to control the behavior of the various IR print methods (e.g.
This class provides the API for ops that are known to be terminators.
A wrapper class that allows for printing an operation with a set of flags, useful to act as a "stream...
Operation is the basic unit of execution within MLIR.
Value getOperand(unsigned idx)
bool hasTrait()
Returns true if the operation was registered with a particular trait, e.g.
void dropAllUses()
Drop all uses of results of this operation.
void setOperand(unsigned idx, Value value)
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Location getLoc()
The source location the operation was defined or derived from.
std::optional< RegisteredOperationName > getRegisteredInfo()
If this operation has a registered operation description, return it.
unsigned getNumOperands()
InFlightDiagnostic emitError(const Twine &message={})
Emit an error about fatal conditions with this operation, reporting up to any diagnostic handlers tha...
MutableArrayRef< Region > getRegions()
Returns the regions held by this operation.
operand_range getOperands()
Returns an iterator on the underlying Value's.
std::enable_if_t< llvm::function_traits< std::decay_t< FnT > >::num_args==1, RetT > walk(FnT &&callback)
Walk the operation by calling the callback for each nested operation (including this one),...
result_range getResults()
MLIRContext * getContext()
Return the context this operation is associated with.
unsigned getNumResults()
Return the number of results held by this operation.
This class implements the result iterators for the Operation class.
This class coordinates the application of a rewrite on a set of IR, providing a way for clients to tr...
void eraseOperands(Operation *op, const BitVector &eraseIndices)
Erase the operands selected by eraseIndices and update operandSegmentSizes if the operation has AttrS...
virtual void eraseOp(Operation *op)
This method erases an operation that is known to have no uses.
Operation * eraseOpResults(Operation *op, const BitVector &eraseIndices)
Erase the specified results of the given operation.
virtual void replaceAllUsesWith(Value from, Value to)
Find uses of from and replace them with to.
This class models how operands are forwarded to block arguments in control flow.
void erase(unsigned subStart, unsigned subLen=1)
Erase operands forwarded to the successor.
MutableOperandRange getMutableForwardedOperands() const
Get the range of operands that are simply forwarded to the successor.
unsigned size() const
Returns the amount of operands passed to the successor.
This class represents a map of symbols to users, and provides efficient implementations of symbol que...
bool areAllUsesVisible(Operation *symbol) const
Return true if all uses of the symbol within the IR are within this scope.
ArrayRef< Operation * > getUsers(Operation *symbol) const
Return the users of the provided symbol operation within this map's scope.
This class provides an abstraction over the different types of ranges over Values.
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
bool use_empty() const
Returns true if this value has no uses.
void dropAllUses()
Drop all uses of this object from their respective owners.
Type getType() const
Return the type of this value.
Location getLoc() const
Return the location of this value.
Include the generated interface declarations.
DenseMap< OpOperand *, SmallVector< Value > > RegionBranchSuccessorMapping
A mapping from successor operands to successor inputs.
llvm::DenseSet< ValueT, ValueInfoT > DenseSet
bool isMemoryEffectFree(Operation *op)
Returns true if the given operation is free of memory effects.
LogicalResult applyOpPatternsGreedily(ArrayRef< Operation * > ops, const FrozenRewritePatternSet &patterns, GreedyRewriteConfig config=GreedyRewriteConfig(), bool *changed=nullptr, bool *allErased=nullptr)
Rewrite the specified ops by repeatedly applying the highest benefit patterns in a greedy worklist dr...
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
This trait indicates that a terminator operation is "return-like".
This lattice represents, for a given value, whether or not it is "live".
Runs liveness analysis on the IR defined by op.
const Liveness * getLiveness(Value val)