17#include "llvm/ADT/EquivalenceClasses.h"
18#include "llvm/Support/DebugLog.h"
26#include "mlir/Interfaces/ControlFlowInterfaces.cpp.inc"
29 : producedOperandCount(0), forwardedOperands(std::move(forwardedOperands)) {
34 : producedOperandCount(producedOperandCount),
35 forwardedOperands(std::move(forwardedOperands)) {}
44std::optional<BlockArgument>
46 unsigned operandIndex,
Block *successor) {
47 LDBG() <<
"Getting branch successor argument for operand index "
48 << operandIndex <<
" in successor block";
52 if (forwardedOperands.empty()) {
53 LDBG() <<
"No forwarded operands, returning nullopt";
59 if (operandIndex < operandsStart ||
60 operandIndex >= (operandsStart + forwardedOperands.size())) {
61 LDBG() <<
"Operand index " << operandIndex <<
" out of range ["
62 << operandsStart <<
", "
63 << (operandsStart + forwardedOperands.size())
64 <<
"), returning nullopt";
71 LDBG() <<
"Computed argument index " << argIndex <<
" for successor block";
79 LDBG() <<
"Verifying branch successor operands for successor #" << succNo
80 <<
" in operation " << op->
getName();
83 unsigned operandCount = operands.
size();
85 LDBG() <<
"Branch has " << operandCount <<
" operands, target block has "
89 return op->
emitError() <<
"branch has " << operandCount
90 <<
" operands for successor #" << succNo
91 <<
", but target block has "
95 LDBG() <<
"Checking type compatibility for "
97 <<
" forwarded operands";
100 Type operandType = operands[i].getType();
102 LDBG() <<
"Checking type compatibility: operand type " << operandType
103 <<
" vs argument type " << argType;
105 if (!cast<BranchOpInterface>(op).areTypesCompatible(operandType, argType))
106 return op->
emitError() <<
"type mismatch for bb argument #" << i
107 <<
" of successor #" << succNo;
110 LDBG() <<
"Branch successor operand verification successful";
120 std::size_t expectedWeightsNum,
121 llvm::StringRef weightAnchorName,
122 llvm::StringRef weightRefName) {
126 if (weights.size() != expectedWeightsNum)
127 return op->
emitError() <<
"expects number of " << weightAnchorName
128 <<
" weights to match number of " << weightRefName
129 <<
": " << weights.size() <<
" vs "
130 << expectedWeightsNum;
132 if (llvm::all_of(weights, [](int32_t value) {
return value == 0; }))
133 return op->
emitError() <<
"branch weights cannot all be zero";
140 cast<WeightedBranchOpInterface>(op).getWeights();
151 cast<WeightedRegionBranchOpInterface>(op).getWeights();
161 auto regionInterface = cast<RegionBranchOpInterface>(op);
166 regionInterface.getAllRegionBranchPoints();
169 regionInterface.getSuccessorRegions(branchPoint, successors);
173 auto emitRegionEdgeError = [&]() {
175 regionInterface->emitOpError(
"along control flow edge from ");
176 if (branchPoint.isParent()) {
178 diag.attachNote(op->
getLoc()) <<
"region branch point";
181 << branchPoint.getTerminatorPredecessorOrNull()->getName();
183 branchPoint.getTerminatorPredecessorOrNull()->getLoc())
184 <<
"region branch point";
187 if (
Region *region = successor.getSuccessor()) {
188 diag <<
"Region #" << region->getRegionNumber();
190 diag <<
"Operation " << successor.getSuccessorOp()->getName();
197 regionInterface.getSuccessorOperands(branchPoint, successor);
198 ValueRange succInputs = regionInterface.getSuccessorInputs(successor);
199 if (succOperands.size() != succInputs.size()) {
200 return emitRegionEdgeError()
201 <<
": region branch point has " << succOperands.size()
202 <<
" operands, but region successor needs " << succInputs.size()
209 for (
const auto &typesIdx :
210 llvm::enumerate(llvm::zip(succOperandTypes, succInputTypes))) {
211 Type succOperandType = std::get<0>(typesIdx.value());
212 Type succInputType = std::get<1>(typesIdx.value());
213 if (!regionInterface.areTypesCompatible(succOperandType, succInputType))
214 return emitRegionEdgeError()
215 <<
": successor operand type #" << typesIdx.index() <<
" "
216 << succOperandType <<
" should match successor input type #"
217 << typesIdx.index() <<
" " << succInputType;
236 auto op = cast<RegionBranchOpInterface>(begin->
getParentOp());
237 LDBG() <<
"Starting region graph traversal from region #"
242 LDBG() <<
"Initialized visited array with " << op->getNumRegions()
247 auto enqueueAllSuccessors = [&](
Region *region) {
248 LDBG() <<
"Enqueuing successors for region #" << region->getRegionNumber();
250 for (
Block &block : *region) {
254 dyn_cast<RegionBranchTerminatorOpInterface>(block.back());
258 operandAttributes.resize(terminator->getNumOperands());
259 terminator.getSuccessorRegions(operandAttributes, successors);
260 LDBG() <<
"Found " << successors.size()
261 <<
" successors from terminator in block";
263 if (successor.isRegion()) {
264 worklist.push_back(successor.getSuccessor());
265 LDBG() <<
"Added region #"
266 << successor.getSuccessor()->getRegionNumber()
269 LDBG() <<
"Skipping operation successor";
274 enqueueAllSuccessors(begin);
275 LDBG() <<
"Initial worklist size: " << worklist.size();
278 while (!worklist.empty()) {
279 Region *nextRegion = worklist.pop_back_val();
281 <<
" from worklist (remaining: " << worklist.size() <<
")";
283 if (stopConditionFn(nextRegion, visited)) {
284 LDBG() <<
"Stop condition met for region #"
289 llvm::errs() <<
"Region " << *nextRegion <<
" has no parent op\n";
294 <<
" already visited, skipping";
300 enqueueAllSuccessors(nextRegion);
303 LDBG() <<
"Traversal completed, returning false";
311 "expected that both regions belong to the same op");
315 return nextRegion == r;
329 LDBG() <<
"Checking if operations are in mutually exclusive regions: "
330 << a->
getName() <<
" and " <<
b->getName();
332 assert(a &&
"expected non-empty operation");
333 assert(
b &&
"expected non-empty operation");
337 LDBG() <<
"Checking branch operation " << branchOp->getName();
340 if (!branchOp->isProperAncestor(
b)) {
341 LDBG() <<
"Operation b is not inside branchOp, checking next ancestor";
343 branchOp = branchOp->getParentOfType<RegionBranchOpInterface>();
347 LDBG() <<
"Both operations are inside branchOp, finding their regions";
351 Region *regionA =
nullptr, *regionB =
nullptr;
352 for (
Region &r : branchOp->getRegions()) {
353 if (r.findAncestorOpInRegion(*a)) {
354 assert(!regionA &&
"already found a region for a");
356 LDBG() <<
"Found region #" << r.
getRegionNumber() <<
" for operation a";
358 if (r.findAncestorOpInRegion(*
b)) {
359 assert(!regionB &&
"already found a region for b");
361 LDBG() <<
"Found region #" << r.getRegionNumber() <<
" for operation b";
364 assert(regionA && regionB &&
"could not find region of op");
366 LDBG() <<
"Region A: #" << regionA->
getRegionNumber() <<
", Region B: #"
367 << regionB->getRegionNumber();
371 bool regionsAreDistinct = (regionA != regionB);
375 LDBG() <<
"Regions distinct: " << regionsAreDistinct
376 <<
", A not reachable from B: " << aNotReachableFromB
377 <<
", B not reachable from A: " << bNotReachableFromA;
379 bool mutuallyExclusive =
380 regionsAreDistinct && aNotReachableFromB && bNotReachableFromA;
381 LDBG() <<
"Operations are mutually exclusive: " << mutuallyExclusive;
383 return mutuallyExclusive;
388 LDBG() <<
"No common RegionBranchOpInterface found, operations are not "
389 "mutually exclusive";
393bool RegionBranchOpInterface::isRepetitiveRegion(
unsigned index) {
394 LDBG() <<
"Checking if region #" <<
index <<
" is repetitive in operation "
395 << getOperation()->getName();
397 Region *region = &getOperation()->getRegion(
index);
400 LDBG() <<
"Region #" <<
index <<
" is repetitive: " << isRepetitive;
404bool RegionBranchOpInterface::hasLoop() {
405 LDBG() <<
"Checking if operation " << getOperation()->getName()
410 LDBG() <<
"Found " << entryRegions.size() <<
" entry regions";
413 if (successor.isRegion()) {
414 LDBG() <<
"Checking entry region #"
415 << successor.getSuccessor()->getRegionNumber() <<
" for loops";
422 return visited[nextRegion->getRegionNumber()];
426 LDBG() <<
"Found loop in entry region #"
427 << successor.getSuccessor()->getRegionNumber();
431 LDBG() <<
"Skipping operation successor";
435 LDBG() <<
"No loops found in operation";
443 return getEntrySuccessorOperands(dest);
448RegionBranchOpInterface::getNonSuccessorInputs(
RegionSuccessor successor) {
453 ValueRange successorInputs = getSuccessorInputs(successor);
454 if (!successorInputs.empty()) {
455 unsigned inputBegin =
457 ? cast<OpResult>(successorInputs.front()).getResultNumber()
458 : cast<BlockArgument>(successorInputs.front()).getArgNumber();
459 results.erase(results.begin() + inputBegin,
460 results.begin() + inputBegin + successorInputs.size());
474 branchOp.getSuccessorRegions(src, successors);
476 OperandRange operands = branchOp.getSuccessorOperands(src, dst);
477 assert(operands.size() == branchOp.getSuccessorInputs(dst).size() &&
478 "expected the same number of operands and inputs");
479 for (
const auto &[operand, input] : llvm::zip_equal(
481 mapping[&operand].push_back(input);
484void RegionBranchOpInterface::getSuccessorOperandInputMapping(
486 std::optional<RegionBranchPoint> src) {
487 if (src.has_value()) {
500 for (
const auto &[operand, inputs] : operandToInputs) {
501 for (
Value input : inputs)
502 inputToOperands[input].push_back(operand);
504 return inputToOperands;
507void RegionBranchOpInterface::getSuccessorInputOperandMapping(
515RegionBranchOpInterface::getAllRegionBranchPoints() {
518 for (
Region ®ion : getOperation()->getRegions()) {
519 for (
Block &block : region) {
522 if (
auto terminator =
523 dyn_cast<RegionBranchTerminatorOpInterface>(block.back()))
531 LDBG() <<
"Finding enclosing repetitive region for operation "
535 LDBG() <<
"Checking region #" << region->getRegionNumber()
536 <<
" in operation " << region->getParentOp()->getName();
539 if (
auto branchOp = dyn_cast<RegionBranchOpInterface>(op)) {
541 <<
"Found RegionBranchOpInterface, checking if region is repetitive";
542 if (branchOp.isRepetitiveRegion(region->getRegionNumber())) {
543 LDBG() <<
"Found repetitive region #" << region->getRegionNumber();
547 LDBG() <<
"Parent operation does not implement RegionBranchOpInterface";
551 LDBG() <<
"No enclosing repetitive region found";
556 LDBG() <<
"Finding enclosing repetitive region for value";
564 if (
auto branchOp = dyn_cast<RegionBranchOpInterface>(op)) {
566 <<
"Found RegionBranchOpInterface, checking if region is repetitive";
572 LDBG() <<
"Parent operation does not implement RegionBranchOpInterface";
577 LDBG() <<
"No enclosing repetitive region found for value";
586 assert((
b.getDefiningOp() == regionBranchOp ||
587 b.getParentRegion()->getParentOp() == regionBranchOp) &&
588 "b must be a region successor input");
601 if (isa<OpResult>(
b))
611 assert(isa<BlockArgument>(
b) &&
"b must be a block argument");
612 return isa<BlockArgument>(a) && cast<BlockArgument>(a).getOwner() ==
613 cast<BlockArgument>(
b).getOwner();
669 std::optional<unsigned> maxReachableValues = std::nullopt) {
670 assert(inputToOperands.contains(value) &&
"value must be a successor input");
671 llvm::SmallDenseSet<Value> visited;
673 worklist.push_back(value);
674 while (!worklist.empty()) {
675 Value next = worklist.pop_back_val();
676 auto it = inputToOperands.find(next);
677 if (it == inputToOperands.end()) {
679 if (maxReachableValues &&
result.size() > *maxReachableValues)
684 if (visited.insert(operand->
get()).second)
685 worklist.push_back(operand->
get());
718struct MakeRegionBranchOpSuccessorInputsDead :
public RewritePattern {
719 MakeRegionBranchOpSuccessorInputsDead(MLIRContext *context, StringRef name,
720 PatternBenefit benefit = 1)
721 : RewritePattern(name, benefit, context) {}
723 LogicalResult matchAndRewrite(Operation *op,
724 PatternRewriter &rewriter)
const override {
726 auto regionBranchOp = cast<RegionBranchOpInterface>(op);
728 regionBranchOp.getSuccessorInputOperandMapping(inputToOperands);
731 bool changed =
false;
732 const bool isIsolated = op->
hasTrait<OpTrait::IsIsolatedFromAbove>();
733 for (Value value : inputToOperands.keys()) {
735 if (value.use_empty())
739 llvm::SmallDenseSet<Value> reachableValues;
741 reachableValues, value, inputToOperands,
743 reachableValues.empty())
747 "successor inputs are supposed to be excluded");
763 Region *valueRegion = value.getParentRegion();
764 if (isIsolated && valueRegion->
getParentOp() == op &&
788template <
typename MappingTy,
typename KeyTy>
789static BitVector &lookupOrCreateBitVector(MappingTy &mapping, KeyTy key,
791 return mapping.try_emplace(key, size,
false).first->second;
804static llvm::EquivalenceClasses<Value> computeTiedSuccessorInputs(
806 llvm::EquivalenceClasses<Value> tiedSuccessorInputs;
807 for (
const auto &[operand, inputs] : operandToInputs) {
808 assert(!inputs.empty() &&
"expected non-empty inputs");
809 Value firstInput = inputs.front();
810 tiedSuccessorInputs.insert(firstInput);
811 for (
Value nextInput : llvm::drop_begin(inputs)) {
814 tiedSuccessorInputs.unionSets(firstInput, nextInput);
817 return tiedSuccessorInputs;
859struct RemoveDeadRegionBranchOpSuccessorInputs :
public RewritePattern {
860 RemoveDeadRegionBranchOpSuccessorInputs(MLIRContext *context, StringRef name,
861 PatternBenefit benefit = 1)
862 : RewritePattern(name, benefit, context) {}
864 LogicalResult matchAndRewrite(Operation *op,
865 PatternRewriter &rewriter)
const override {
869 auto regionBranchOp = cast<RegionBranchOpInterface>(op);
871 regionBranchOp.getSuccessorOperandInputMapping(operandToInputs);
872 llvm::EquivalenceClasses<Value> tiedSuccessorInputs =
873 computeTiedSuccessorInputs(operandToInputs);
876 SmallVector<Value> valuesToRemove;
878 BitVector resultsToRemove(regionBranchOp->getNumResults(),
false);
880 for (
auto it = tiedSuccessorInputs.begin(), e = tiedSuccessorInputs.end();
882 if (!(*it)->isLeader())
888 for (
auto memberIt = tiedSuccessorInputs.member_begin(**it);
889 memberIt != tiedSuccessorInputs.member_end(); ++memberIt) {
891 if (!memberIt->use_empty()) {
901 for (
auto memberIt = tiedSuccessorInputs.member_begin(**it);
902 memberIt != tiedSuccessorInputs.member_end(); ++memberIt) {
903 if (
auto arg = dyn_cast<BlockArgument>(*memberIt)) {
906 lookupOrCreateBitVector(blockArgsToRemove, arg.getOwner(),
907 arg.getOwner()->getNumArguments());
908 vector.set(arg.getArgNumber());
911 OpResult
result = cast<OpResult>(*memberIt);
912 assert(
result.getDefiningOp() == regionBranchOp &&
913 "result must be a region branch op result");
914 resultsToRemove.set(
result.getResultNumber());
916 valuesToRemove.push_back(*memberIt);
920 if (valuesToRemove.empty())
927 for (Value value : valuesToRemove) {
928 for (OpOperand *operand : inputsToOperands[value]) {
931 lookupOrCreateBitVector(operandsToRemove, operand->getOwner(),
932 operand->getOwner()->getNumOperands());
933 vector.set(operand->getOperandNumber());
938 for (
auto &pair : operandsToRemove) {
939 Operation *op = pair.first;
940 BitVector &operands = pair.second;
945 for (
auto &pair : blockArgsToRemove) {
946 Block *block = pair.first;
947 BitVector &blockArg = pair.second;
949 [&]() { block->eraseArguments(blockArg); });
953 if (resultsToRemove.any())
963 if (
auto arg = dyn_cast<BlockArgument>(value))
964 return arg.getOwner();
969static unsigned getArgOrResultNumber(
Value value) {
970 if (
auto opResult = llvm::dyn_cast<OpResult>(value))
971 return opResult.getResultNumber();
972 return llvm::cast<BlockArgument>(value).getArgNumber();
1006struct RemoveDuplicateSuccessorInputUses :
public RewritePattern {
1007 RemoveDuplicateSuccessorInputUses(MLIRContext *context, StringRef name,
1008 PatternBenefit benefit = 1)
1009 : RewritePattern(name, benefit, context) {}
1011 LogicalResult matchAndRewrite(Operation *op,
1012 PatternRewriter &rewriter)
const override {
1019 auto regionBranchOp = cast<RegionBranchOpInterface>(op);
1021 regionBranchOp.getSuccessorInputOperandMapping(inputsToOperands);
1022 SmallVector<Value> inputs = llvm::to_vector(inputsToOperands.keys());
1023 llvm::sort(inputs, [](Value a, Value
b) {
1024 return getArgOrResultNumber(a) < getArgOrResultNumber(
b);
1037 using SigEntry = std::pair<Operation *, Value>;
1038 using Signature = SmallVector<SigEntry>;
1039 auto sigEntryLess = [](
const SigEntry &a,
const SigEntry &
b) {
1040 if (a.first !=
b.first)
1041 return a.first <
b.first;
1047 using MapKey = std::pair<Signature, void *>;
1048 auto mapKeyLess = [&](
const MapKey &a,
const MapKey &
b) {
1049 if (a.second !=
b.second)
1050 return a.second <
b.second;
1051 return std::lexicographical_compare(a.first.begin(), a.first.end(),
1052 b.first.begin(),
b.first.end(),
1055 std::map<MapKey, Value,
decltype(mapKeyLess)> signatureToCanonical(
1057 bool changed =
false;
1060 for (Value input : inputs) {
1064 for (OpOperand *operand : inputsToOperands[input])
1065 sig.emplace_back(operand->getOwner(), operand->get());
1066 llvm::sort(sig, sigEntryLess);
1070 auto [it,
inserted] = signatureToCanonical.try_emplace(
1071 MapKey{std::move(sig), owner}, input);
1073 Value canonical = it->second;
1075 if (input.use_empty())
1089 return llvm::map_to_vector(values, [](
Value value) {
1100getSuccessorRegionsWithAttrs(RegionBranchOpInterface op,
1104 op.getEntrySuccessorRegions(extractConstants(op->getOperands()),
1108 RegionBranchTerminatorOpInterface terminator =
1110 terminator.getSuccessorRegions(extractConstants(terminator->getOperands()),
1136computeSingleAcyclicRegionBranchPath(RegionBranchOpInterface op) {
1137 llvm::SmallDenseSet<Region *> visited;
1144 getSuccessorRegionsWithAttrs(op, next);
1145 if (successors.size() != 1) {
1150 path.push_back(successors.front());
1151 if (successors.front().isOperation()) {
1155 Region *region = successors.front().getSuccessor();
1161 if (!visited.insert(region).second) {
1166 dyn_cast<RegionBranchTerminatorOpInterface>(®ion->
front().
back());
1174 llvm_unreachable(
"expected to return from loop");
1212 InlineRegionBranchOp(MLIRContext *context, StringRef name,
1215 : RewritePattern(name, benefit, context), replBuilderFn(replBuilderFn),
1216 matcherFn(matcherFn) {}
1218 LogicalResult matchAndRewrite(Operation *op,
1219 PatternRewriter &rewriter)
const override {
1221 if (
failed(matcherFn(op)))
1226 if (!op->
hasTrait<OpTrait::HasRecursiveMemoryEffects>())
1228 op,
"pattern not applicable to ops without recursive memory effects");
1231 auto regionBranchOp = cast<RegionBranchOpInterface>(op);
1232 SmallVector<RegionSuccessor> path =
1233 computeSingleAcyclicRegionBranchPath(regionBranchOp);
1236 op,
"failed to find acyclic region branch path");
1240 ArrayRef remainingPath = path;
1241 SmallVector<Value> successorOperands = llvm::to_vector(
1242 regionBranchOp.getEntrySuccessorOperands(remainingPath.front()));
1243 while (!remainingPath.empty()) {
1244 RegionSuccessor nextSuccessor = remainingPath.consume_front();
1246 regionBranchOp.getSuccessorInputs(nextSuccessor);
1247 assert(successorInputs.size() == successorOperands.size() &&
1251 unsigned firstSuccessorInputIdx = 0;
1252 if (!successorInputs.empty())
1253 firstSuccessorInputIdx =
1255 ? cast<OpResult>(successorInputs.front()).getResultNumber()
1256 : cast<BlockArgument>(successorInputs.front()).getArgNumber();
1258 unsigned numValues =
1263 SmallVector<Value> replacements;
1265 auto getValue = [&](
unsigned idx) {
1268 : Value(nextSuccessor.getSuccessor()->getArgument(idx));
1272 for (
unsigned i = 0; i < firstSuccessorInputIdx; ++i)
1273 replacements.push_back(
1274 replBuilderFn(rewriter, op->
getLoc(), getValue(i)));
1277 llvm::append_range(replacements, successorOperands);
1280 for (
unsigned i = replacements.size(); i < numValues; ++i)
1281 replacements.push_back(
1282 replBuilderFn(rewriter, op->
getLoc(), getValue(i)));
1288 op,
"path ends after a different operation");
1289 assert(remainingPath.empty() &&
"expected that the path ended");
1296 auto terminator = cast<RegionBranchTerminatorOpInterface>(
1301 successorOperands = llvm::to_vector(
1302 terminator.getSuccessorOperands(remainingPath.front()));
1306 llvm_unreachable(
"expected that path ends with an operation");
1316 patterns.
add<MakeRegionBranchOpSuccessorInputsDead,
1317 RemoveDuplicateSuccessorInputUses,
1318 RemoveDeadRegionBranchOpSuccessorInputs>(patterns.
getContext(),
1326 patterns.
add<InlineRegionBranchOp>(patterns.
getContext(), opName,
1327 replBuilderFn, matcherFn, benefit);
static LogicalResult verifyWeights(Operation *op, llvm::ArrayRef< int32_t > weights, std::size_t expectedWeightsNum, llvm::StringRef weightAnchorName, llvm::StringRef weightRefName)
static bool isDefinedBefore(Operation *regionBranchOp, Value a, Value b)
Return "true" if a can be used in lieu of b, where b is a region successor input and a is a "reachabl...
static void getSuccessorOperandInputMapping(RegionBranchOpInterface branchOp, RegionBranchSuccessorMapping &mapping, RegionBranchPoint src)
static bool traverseRegionGraph(Region *begin, StopConditionFn stopConditionFn)
Traverse the region graph starting at begin.
static LogicalResult computeReachableValuesFromSuccessorInput(llvm::SmallDenseSet< Value > &result, Value value, const RegionBranchInverseSuccessorMapping &inputToOperands, std::optional< unsigned > maxReachableValues=std::nullopt)
Compute all non-successor-input values that a successor input could have based on the given successor...
static RegionBranchInverseSuccessorMapping invertRegionBranchSuccessorMapping(const RegionBranchSuccessorMapping &operandToInputs)
function_ref< bool(Region *, ArrayRef< bool > visited)> StopConditionFn
Stop condition for traverseRegionGraph.
static bool isRegionReachable(Region *begin, Region *r)
Return true if region r is reachable from region begin according to the RegionBranchOpInterface (by t...
*if copies could not be generated due to yet unimplemented cases *copyInPlacementStart and copyOutPlacementStart in copyPlacementBlock *specify the insertion points where the incoming copies and outgoing should be inserted(the insertion happens right before the *insertion point). Since `begin` can itself be invalidated due to the memref *rewriting done from this method
*if copies could not be generated due to yet unimplemented cases *copyInPlacementStart and copyOutPlacementStart in copyPlacementBlock *specify the insertion points where the incoming copies and outgoing should be the output argument nBegin is set to its * replacement(set to `begin` if no invalidation happens). Since outgoing *copies could have been inserted at `end`
static MutableArrayRef< OpOperand > operandsToOpOperands(OperandRange &operands)
static Operation * getOwnerOfValue(Value value)
Attributes are known-constant values of operations.
Block represents an ordered list of Operations.
BlockArgument getArgument(unsigned i)
unsigned getNumArguments()
Operation * getParentOp()
Returns the closest surrounding operation that contains this block.
IRValueT get() const
Return the current value being used by this operand.
This class represents a diagnostic that is inflight and set to be reported.
This class provides a mutable adaptor for a range of operands.
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
This class represents an operand of an operation.
This class implements the operand iterators for the Operation class.
unsigned getBeginOperandIndex() const
Return the operand index of the first element of this range.
type_range getTypes() const
Operation is the basic unit of execution within MLIR.
bool hasTrait()
Returns true if the operation was registered with a particular trait, e.g.
unsigned getNumSuccessors()
Block * getBlock()
Returns the operation block that contains this operation.
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
unsigned getNumRegions()
Returns the number of regions held by this operation.
Location getLoc()
The source location the operation was defined or derived from.
Operation * getParentOp()
Returns the closest surrounding operation that contains this operation or nullptr if this is a top-le...
InFlightDiagnostic emitError(const Twine &message={})
Emit an error about fatal conditions with this operation, reporting up to any diagnostic handlers tha...
OpTy getParentOfType()
Return the closest surrounding parent operation that is of type 'OpTy'.
OperationName getName()
The name of an operation is the key identifier for it.
Block * getSuccessor(unsigned index)
result_range getResults()
Region * getParentRegion()
Returns the region to which the instruction belongs.
unsigned getNumResults()
Return the number of results held by this operation.
This class represents the benefit of a pattern match in a unitless scheme that ranges from 0 (very li...
This class represents a point being branched from in the methods of the RegionBranchOpInterface.
bool isParent() const
Returns true if branching from the parent op.
static constexpr RegionBranchPoint parent()
Returns an instance of RegionBranchPoint representing the parent operation.
RegionBranchTerminatorOpInterface getTerminatorPredecessorOrNull() const
Returns the terminator if branching from a region.
This class represents a successor of a region.
Region * getSuccessor() const
Return the given region successor.
bool isOperation() const
Return true if the successor is an operation.
Operation * getSuccessorOp() const
Return the given operation successor.
This class contains a list of basic blocks and a link to the parent operation it is attached to.
BlockArgListType getArguments()
unsigned getRegionNumber()
Return the number of this region in the parent operation.
unsigned getNumArguments()
Operation * getParentOp()
Return the parent operation this region is attached to.
bool hasOneBlock()
Return true if this region has exactly one block.
MLIRContext * getContext() const
RewritePatternSet & add(ConstructorArg &&arg, ConstructorArgs &&...args)
Add an instance of each of the pattern types 'Ts' to the pattern list with the given arguments.
RewritePattern is the common base class for all DAG to DAG replacements.
void eraseOperands(Operation *op, const BitVector &eraseIndices)
Erase the operands selected by eraseIndices and update operandSegmentSizes if the operation has AttrS...
virtual void replaceOp(Operation *op, ValueRange newValues)
Replace the results of the given (original) operation with the specified list of values (replacements...
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 inlineBlockBefore(Block *source, Block *dest, Block::iterator before, ValueRange argValues={})
Inline the operations of block 'source' into block 'dest' before the given position.
std::enable_if_t<!std::is_convertible< CallbackT, Twine >::value, LogicalResult > notifyMatchFailure(Location loc, CallbackT &&reasonCallback)
Used to notify the listener that the IR failed to be rewritten because of a match failure,...
void modifyOpInPlace(Operation *root, CallableT &&callable)
This method is a utility wrapper around an in-place modification of an 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.
SuccessorOperands(MutableOperandRange forwardedOperands)
Constructs a SuccessorOperands with no produced operands that simply forwards operands to the success...
unsigned getProducedOperandCount() const
Returns the amount of operands that are produced internally by the operation.
unsigned size() const
Returns the amount of operands passed to the successor.
OperandRange getForwardedOperands() const
Get the range of operands that are simply forwarded to the successor.
This class provides an abstraction over the various different ranges of value types.
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
This class provides an abstraction over the different types of ranges over Values.
type_range getTypes() const
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Type getType() const
Return the type of this value.
void * getAsOpaquePointer() const
Methods for supporting PointerLikeTypeTraits.
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Region * getParentRegion()
Return the Region in which this Value is defined.
std::optional< BlockArgument > getBranchSuccessorArgument(const SuccessorOperands &operands, unsigned operandIndex, Block *successor)
Return the BlockArgument corresponding to operand operandIndex in some successor if operandIndex is w...
LogicalResult verifyRegionBranchWeights(Operation *op)
Verify that the region weights attached to an operation implementing WeightedRegiobBranchOpInterface ...
LogicalResult verifyBranchSuccessorOperands(Operation *op, unsigned succNo, const SuccessorOperands &operands)
Verify that the given operands match those of the given successor block.
LogicalResult verifyRegionBranchOpInterface(Operation *op)
Verify that types match along control flow edges described the given op.
LogicalResult verifyBranchWeights(Operation *op)
Verify that the branch weights attached to an operation implementing WeightedBranchOpInterface are co...
InFlightDiagnostic & next(InFlightDiagnostic &diag)
Starts a new message part in an in-flight diagnostic.
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
DenseMap< OpOperand *, SmallVector< Value > > RegionBranchSuccessorMapping
A mapping from successor operands to successor inputs.
std::function< LogicalResult(Operation *)> PatternMatcherFn
Helper function for the region branch op inlining pattern that checks if the pattern is applicable to...
bool insideMutuallyExclusiveRegions(Operation *a, Operation *b)
Return true if a and b are in mutually exclusive regions as per RegionBranchOpInterface.
std::function< Value(OpBuilder &, Location, Value)> NonSuccessorInputReplacementBuilderFn
Helper function for the region branch op inlining pattern that builds replacement values for non-succ...
Region * getEnclosingRepetitiveRegion(Operation *op)
Return the first enclosing region of the given op that may be executed repetitively as per RegionBran...
DenseMap< Value, SmallVector< OpOperand * > > RegionBranchInverseSuccessorMapping
llvm::DenseMap< KeyT, ValueT, KeyInfoT, BucketT > DenseMap
void populateRegionBranchOpInterfaceInliningPattern(RewritePatternSet &patterns, StringRef opName, NonSuccessorInputReplacementBuilderFn replBuilderFn=detail::defaultReplBuilderFn, PatternMatcherFn matcherFn=detail::defaultMatcherFn, PatternBenefit benefit=1)
Populate a pattern that inlines the body of region branch ops when there is a single acyclic path thr...
void populateRegionBranchOpInterfaceCanonicalizationPatterns(RewritePatternSet &patterns, StringRef opName, PatternBenefit benefit=1)
Populate canonicalization patterns that simplify successor operands/inputs of region branch operation...
detail::constant_op_matcher m_Constant()
Matches a constant foldable operation.
llvm::function_ref< Fn > function_ref