MLIR 24.0.0git
ControlFlowInterfaces.cpp
Go to the documentation of this file.
1//===- ControlFlowInterfaces.cpp - ControlFlow Interfaces -----------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8
9#include <map>
10#include <utility>
11
13#include "mlir/IR/Matchers.h"
14#include "mlir/IR/Operation.h"
17#include "llvm/ADT/EquivalenceClasses.h"
18#include "llvm/Support/DebugLog.h"
19
20using namespace mlir;
21
22//===----------------------------------------------------------------------===//
23// ControlFlowInterfaces
24//===----------------------------------------------------------------------===//
25
26#include "mlir/Interfaces/ControlFlowInterfaces.cpp.inc"
27
29 : producedOperandCount(0), forwardedOperands(std::move(forwardedOperands)) {
30}
31
32SuccessorOperands::SuccessorOperands(unsigned int producedOperandCount,
33 MutableOperandRange forwardedOperands)
34 : producedOperandCount(producedOperandCount),
35 forwardedOperands(std::move(forwardedOperands)) {}
36
37//===----------------------------------------------------------------------===//
38// BranchOpInterface
39//===----------------------------------------------------------------------===//
40
41/// Returns the `BlockArgument` corresponding to operand `operandIndex` in some
42/// successor if 'operandIndex' is within the range of 'operands', or
43/// std::nullopt if `operandIndex` isn't a successor operand index.
44std::optional<BlockArgument>
46 unsigned operandIndex, Block *successor) {
47 LDBG() << "Getting branch successor argument for operand index "
48 << operandIndex << " in successor block";
49
50 OperandRange forwardedOperands = operands.getForwardedOperands();
51 // Check that the operands are valid.
52 if (forwardedOperands.empty()) {
53 LDBG() << "No forwarded operands, returning nullopt";
54 return std::nullopt;
55 }
56
57 // Check to ensure that this operand is within the range.
58 unsigned operandsStart = forwardedOperands.getBeginOperandIndex();
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";
65 return std::nullopt;
66 }
67
68 // Index the successor.
69 unsigned argIndex =
70 operands.getProducedOperandCount() + operandIndex - operandsStart;
71 LDBG() << "Computed argument index " << argIndex << " for successor block";
72 return successor->getArgument(argIndex);
73}
74
75/// Verify that the given operands match those of the given successor block.
76LogicalResult
78 const SuccessorOperands &operands) {
79 LDBG() << "Verifying branch successor operands for successor #" << succNo
80 << " in operation " << op->getName();
81
82 // Check the count.
83 unsigned operandCount = operands.size();
84 Block *destBB = op->getSuccessor(succNo);
85 LDBG() << "Branch has " << operandCount << " operands, target block has "
86 << destBB->getNumArguments() << " arguments";
87
88 if (operandCount != destBB->getNumArguments())
89 return op->emitError() << "branch has " << operandCount
90 << " operands for successor #" << succNo
91 << ", but target block has "
92 << destBB->getNumArguments();
93
94 // Check the types.
95 LDBG() << "Checking type compatibility for "
96 << (operandCount - operands.getProducedOperandCount())
97 << " forwarded operands";
98 for (unsigned i = operands.getProducedOperandCount(); i != operandCount;
99 ++i) {
100 Type operandType = operands[i].getType();
101 Type argType = destBB->getArgument(i).getType();
102 LDBG() << "Checking type compatibility: operand type " << operandType
103 << " vs argument type " << argType;
104
105 if (!cast<BranchOpInterface>(op).areTypesCompatible(operandType, argType))
106 return op->emitError() << "type mismatch for bb argument #" << i
107 << " of successor #" << succNo;
108 }
109
110 LDBG() << "Branch successor operand verification successful";
111 return success();
112}
113
114//===----------------------------------------------------------------------===//
115// WeightedBranchOpInterface
116//===----------------------------------------------------------------------===//
117
118static LogicalResult verifyWeights(Operation *op,
120 std::size_t expectedWeightsNum,
121 llvm::StringRef weightAnchorName,
122 llvm::StringRef weightRefName) {
123 if (weights.empty())
124 return success();
125
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;
131
132 if (llvm::all_of(weights, [](int32_t value) { return value == 0; }))
133 return op->emitError() << "branch weights cannot all be zero";
134
135 return success();
136}
137
140 cast<WeightedBranchOpInterface>(op).getWeights();
141 return verifyWeights(op, weights, op->getNumSuccessors(), "branch",
142 "successors");
143}
144
145//===----------------------------------------------------------------------===//
146// WeightedRegionBranchOpInterface
147//===----------------------------------------------------------------------===//
148
151 cast<WeightedRegionBranchOpInterface>(op).getWeights();
152 return verifyWeights(op, weights, op->getNumRegions(), "region", "regions");
153}
154
155//===----------------------------------------------------------------------===//
156// RegionBranchOpInterface
157//===----------------------------------------------------------------------===//
158
159/// Verify that types match along control flow edges described the given op.
161 auto regionInterface = cast<RegionBranchOpInterface>(op);
162
163 // Verify all control flow edges from region branch points to region
164 // successors.
165 SmallVector<RegionBranchPoint> regionBranchPoints =
166 regionInterface.getAllRegionBranchPoints();
167 for (const RegionBranchPoint &branchPoint : regionBranchPoints) {
169 regionInterface.getSuccessorRegions(branchPoint, successors);
170 for (const RegionSuccessor &successor : successors) {
171 // Helper function that print the region branch point and the region
172 // successor.
173 auto emitRegionEdgeError = [&]() {
175 regionInterface->emitOpError("along control flow edge from ");
176 if (branchPoint.isParent()) {
177 diag << "parent";
178 diag.attachNote(op->getLoc()) << "region branch point";
179 } else {
180 diag << "Operation "
181 << branchPoint.getTerminatorPredecessorOrNull()->getName();
182 diag.attachNote(
183 branchPoint.getTerminatorPredecessorOrNull()->getLoc())
184 << "region branch point";
185 }
186 diag << " to ";
187 if (Region *region = successor.getSuccessor()) {
188 diag << "Region #" << region->getRegionNumber();
189 } else {
190 diag << "Operation " << successor.getSuccessorOp()->getName();
191 }
192 return diag;
193 };
194
195 // Verify number of successor operands and successor inputs.
196 OperandRange succOperands =
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()
203 << " inputs";
204 }
205
206 // Verify that the types are compatible.
207 TypeRange succInputTypes = succInputs.getTypes();
208 TypeRange succOperandTypes = succOperands.getTypes();
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;
218 }
219 }
220 }
221 return success();
222}
223
224/// Stop condition for `traverseRegionGraph`. The traversal is interrupted if
225/// this function returns "true" for a successor region. The first parameter is
226/// the successor region. The second parameter indicates all already visited
227/// regions.
229
230/// Traverse the region graph starting at `begin`. The traversal is interrupted
231/// if `stopCondition` evaluates to "true" for a successor region. In that case,
232/// this function returns "true". Otherwise, if the traversal was not
233/// interrupted, this function returns "false".
234static bool traverseRegionGraph(Region *begin,
235 StopConditionFn stopConditionFn) {
236 auto op = cast<RegionBranchOpInterface>(begin->getParentOp());
237 LDBG() << "Starting region graph traversal from region #"
238 << begin->getRegionNumber() << " in operation " << op->getName();
239
240 SmallVector<bool> visited(op->getNumRegions(), false);
241 visited[begin->getRegionNumber()] = true;
242 LDBG() << "Initialized visited array with " << op->getNumRegions()
243 << " regions";
244
245 // Retrieve all successors of the region and enqueue them in the worklist.
246 SmallVector<Region *> worklist;
247 auto enqueueAllSuccessors = [&](Region *region) {
248 LDBG() << "Enqueuing successors for region #" << region->getRegionNumber();
249 SmallVector<Attribute> operandAttributes(op->getNumOperands());
250 for (Block &block : *region) {
251 if (block.empty())
252 continue;
253 auto terminator =
254 dyn_cast<RegionBranchTerminatorOpInterface>(block.back());
255 if (!terminator)
256 continue;
258 operandAttributes.resize(terminator->getNumOperands());
259 terminator.getSuccessorRegions(operandAttributes, successors);
260 LDBG() << "Found " << successors.size()
261 << " successors from terminator in block";
262 for (RegionSuccessor successor : successors) {
263 if (successor.isRegion()) {
264 worklist.push_back(successor.getSuccessor());
265 LDBG() << "Added region #"
266 << successor.getSuccessor()->getRegionNumber()
267 << " to worklist";
268 } else {
269 LDBG() << "Skipping operation successor";
270 }
271 }
272 }
273 };
274 enqueueAllSuccessors(begin);
275 LDBG() << "Initial worklist size: " << worklist.size();
276
277 // Process all regions in the worklist via DFS.
278 while (!worklist.empty()) {
279 Region *nextRegion = worklist.pop_back_val();
280 LDBG() << "Processing region #" << nextRegion->getRegionNumber()
281 << " from worklist (remaining: " << worklist.size() << ")";
282
283 if (stopConditionFn(nextRegion, visited)) {
284 LDBG() << "Stop condition met for region #"
285 << nextRegion->getRegionNumber() << ", returning true";
286 return true;
287 }
288 if (!nextRegion->getParentOp()) {
289 llvm::errs() << "Region " << *nextRegion << " has no parent op\n";
290 return false;
291 }
292 if (visited[nextRegion->getRegionNumber()]) {
293 LDBG() << "Region #" << nextRegion->getRegionNumber()
294 << " already visited, skipping";
295 continue;
296 }
297 visited[nextRegion->getRegionNumber()] = true;
298 LDBG() << "Marking region #" << nextRegion->getRegionNumber()
299 << " as visited";
300 enqueueAllSuccessors(nextRegion);
301 }
302
303 LDBG() << "Traversal completed, returning false";
304 return false;
305}
306
307/// Return `true` if region `r` is reachable from region `begin` according to
308/// the RegionBranchOpInterface (by taking a branch).
309static bool isRegionReachable(Region *begin, Region *r) {
310 assert(begin->getParentOp() == r->getParentOp() &&
311 "expected that both regions belong to the same op");
312 return traverseRegionGraph(begin,
313 [&](Region *nextRegion, ArrayRef<bool> visited) {
314 // Interrupt traversal if `r` was reached.
315 return nextRegion == r;
316 });
317}
318
319/// Return `true` if `a` and `b` are in mutually exclusive regions.
320///
321/// 1. Find the first common of `a` and `b` (ancestor) that implements
322/// RegionBranchOpInterface.
323/// 2. Determine the regions `regionA` and `regionB` in which `a` and `b` are
324/// contained.
325/// 3. Check if `regionA` and `regionB` are mutually exclusive. They are
326/// mutually exclusive if they are not reachable from each other as per
327/// RegionBranchOpInterface::getSuccessorRegions.
329 LDBG() << "Checking if operations are in mutually exclusive regions: "
330 << a->getName() << " and " << b->getName();
331
332 assert(a && "expected non-empty operation");
333 assert(b && "expected non-empty operation");
334
335 auto branchOp = a->getParentOfType<RegionBranchOpInterface>();
336 while (branchOp) {
337 LDBG() << "Checking branch operation " << branchOp->getName();
338
339 // Check if b is inside branchOp. (We already know that a is.)
340 if (!branchOp->isProperAncestor(b)) {
341 LDBG() << "Operation b is not inside branchOp, checking next ancestor";
342 // Check next enclosing RegionBranchOpInterface.
343 branchOp = branchOp->getParentOfType<RegionBranchOpInterface>();
344 continue;
345 }
346
347 LDBG() << "Both operations are inside branchOp, finding their regions";
348
349 // b is contained in branchOp. Retrieve the regions in which `a` and `b`
350 // are contained.
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");
355 regionA = &r;
356 LDBG() << "Found region #" << r.getRegionNumber() << " for operation a";
357 }
358 if (r.findAncestorOpInRegion(*b)) {
359 assert(!regionB && "already found a region for b");
360 regionB = &r;
361 LDBG() << "Found region #" << r.getRegionNumber() << " for operation b";
362 }
363 }
364 assert(regionA && regionB && "could not find region of op");
365
366 LDBG() << "Region A: #" << regionA->getRegionNumber() << ", Region B: #"
367 << regionB->getRegionNumber();
368
369 // `a` and `b` are in mutually exclusive regions if both regions are
370 // distinct and neither region is reachable from the other region.
371 bool regionsAreDistinct = (regionA != regionB);
372 bool aNotReachableFromB = !isRegionReachable(regionA, regionB);
373 bool bNotReachableFromA = !isRegionReachable(regionB, regionA);
374
375 LDBG() << "Regions distinct: " << regionsAreDistinct
376 << ", A not reachable from B: " << aNotReachableFromB
377 << ", B not reachable from A: " << bNotReachableFromA;
378
379 bool mutuallyExclusive =
380 regionsAreDistinct && aNotReachableFromB && bNotReachableFromA;
381 LDBG() << "Operations are mutually exclusive: " << mutuallyExclusive;
382
383 return mutuallyExclusive;
384 }
385
386 // Could not find a common RegionBranchOpInterface among a's and b's
387 // ancestors.
388 LDBG() << "No common RegionBranchOpInterface found, operations are not "
389 "mutually exclusive";
390 return false;
391}
392
393bool RegionBranchOpInterface::isRepetitiveRegion(unsigned index) {
394 LDBG() << "Checking if region #" << index << " is repetitive in operation "
395 << getOperation()->getName();
396
397 Region *region = &getOperation()->getRegion(index);
398 bool isRepetitive = isRegionReachable(region, region);
399
400 LDBG() << "Region #" << index << " is repetitive: " << isRepetitive;
401 return isRepetitive;
402}
403
404bool RegionBranchOpInterface::hasLoop() {
405 LDBG() << "Checking if operation " << getOperation()->getName()
406 << " has loops";
407
408 SmallVector<RegionSuccessor> entryRegions;
409 getSuccessorRegions(RegionBranchPoint::parent(), entryRegions);
410 LDBG() << "Found " << entryRegions.size() << " entry regions";
411
412 for (RegionSuccessor successor : entryRegions) {
413 if (successor.isRegion()) {
414 LDBG() << "Checking entry region #"
415 << successor.getSuccessor()->getRegionNumber() << " for loops";
416
417 bool hasLoop =
418 traverseRegionGraph(successor.getSuccessor(),
419 [](Region *nextRegion, ArrayRef<bool> visited) {
420 // Interrupt traversal if the region was already
421 // visited.
422 return visited[nextRegion->getRegionNumber()];
423 });
424
425 if (hasLoop) {
426 LDBG() << "Found loop in entry region #"
427 << successor.getSuccessor()->getRegionNumber();
428 return true;
429 }
430 } else {
431 LDBG() << "Skipping operation successor";
432 }
433 }
434
435 LDBG() << "No loops found in operation";
436 return false;
437}
438
440RegionBranchOpInterface::getSuccessorOperands(RegionBranchPoint src,
441 RegionSuccessor dest) {
442 if (src.isParent())
443 return getEntrySuccessorOperands(dest);
444 return src.getTerminatorPredecessorOrNull().getSuccessorOperands(dest);
445}
446
448RegionBranchOpInterface::getNonSuccessorInputs(RegionSuccessor successor) {
449 SmallVector<Value> results = llvm::to_vector(
450 successor.isOperation()
451 ? ValueRange(successor.getSuccessorOp()->getResults())
452 : ValueRange(successor.getSuccessor()->getArguments()));
453 ValueRange successorInputs = getSuccessorInputs(successor);
454 if (!successorInputs.empty()) {
455 unsigned inputBegin =
456 successor.isOperation()
457 ? cast<OpResult>(successorInputs.front()).getResultNumber()
458 : cast<BlockArgument>(successorInputs.front()).getArgNumber();
459 results.erase(results.begin() + inputBegin,
460 results.begin() + inputBegin + successorInputs.size());
461 }
462 return results;
463}
464
466 return MutableArrayRef<OpOperand>(operands.getBase(), operands.size());
467}
468
469static void
470getSuccessorOperandInputMapping(RegionBranchOpInterface branchOp,
472 RegionBranchPoint src) {
474 branchOp.getSuccessorRegions(src, successors);
475 for (RegionSuccessor dst : 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(
480 operandsToOpOperands(operands), branchOp.getSuccessorInputs(dst)))
481 mapping[&operand].push_back(input);
482 }
483}
484void RegionBranchOpInterface::getSuccessorOperandInputMapping(
486 std::optional<RegionBranchPoint> src) {
487 if (src.has_value()) {
488 ::getSuccessorOperandInputMapping(*this, mapping, src.value());
489 } else {
490 // No region branch point specified: populate the mapping for all possible
491 // region branch points.
492 for (RegionBranchPoint branchPoint : getAllRegionBranchPoints())
493 ::getSuccessorOperandInputMapping(*this, mapping, branchPoint);
494 }
495}
496
498 const RegionBranchSuccessorMapping &operandToInputs) {
500 for (const auto &[operand, inputs] : operandToInputs) {
501 for (Value input : inputs)
502 inputToOperands[input].push_back(operand);
503 }
504 return inputToOperands;
505}
506
507void RegionBranchOpInterface::getSuccessorInputOperandMapping(
509 RegionBranchSuccessorMapping operandToInputs;
510 getSuccessorOperandInputMapping(operandToInputs);
511 mapping = invertRegionBranchSuccessorMapping(operandToInputs);
512}
513
515RegionBranchOpInterface::getAllRegionBranchPoints() {
517 branchPoints.push_back(RegionBranchPoint::parent());
518 for (Region &region : getOperation()->getRegions()) {
519 for (Block &block : region) {
520 if (block.empty())
521 continue;
522 if (auto terminator =
523 dyn_cast<RegionBranchTerminatorOpInterface>(block.back()))
524 branchPoints.push_back(RegionBranchPoint(terminator));
525 }
526 }
527 return branchPoints;
528}
529
531 LDBG() << "Finding enclosing repetitive region for operation "
532 << op->getName();
533
534 while (Region *region = op->getParentRegion()) {
535 LDBG() << "Checking region #" << region->getRegionNumber()
536 << " in operation " << region->getParentOp()->getName();
537
538 op = region->getParentOp();
539 if (auto branchOp = dyn_cast<RegionBranchOpInterface>(op)) {
540 LDBG()
541 << "Found RegionBranchOpInterface, checking if region is repetitive";
542 if (branchOp.isRepetitiveRegion(region->getRegionNumber())) {
543 LDBG() << "Found repetitive region #" << region->getRegionNumber();
544 return region;
545 }
546 } else {
547 LDBG() << "Parent operation does not implement RegionBranchOpInterface";
548 }
549 }
550
551 LDBG() << "No enclosing repetitive region found";
552 return nullptr;
553}
554
556 LDBG() << "Finding enclosing repetitive region for value";
557
558 Region *region = value.getParentRegion();
559 while (region) {
560 LDBG() << "Checking region #" << region->getRegionNumber()
561 << " in operation " << region->getParentOp()->getName();
562
563 Operation *op = region->getParentOp();
564 if (auto branchOp = dyn_cast<RegionBranchOpInterface>(op)) {
565 LDBG()
566 << "Found RegionBranchOpInterface, checking if region is repetitive";
567 if (branchOp.isRepetitiveRegion(region->getRegionNumber())) {
568 LDBG() << "Found repetitive region #" << region->getRegionNumber();
569 return region;
570 }
571 } else {
572 LDBG() << "Parent operation does not implement RegionBranchOpInterface";
573 }
574 region = op->getParentRegion();
575 }
576
577 LDBG() << "No enclosing repetitive region found for value";
578 return nullptr;
579}
580
581/// Return "true" if `a` can be used in lieu of `b`, where `b` is a region
582/// successor input and `a` is a "reachable value" of `b`. Reachable values
583/// are successor operand values that are (maybe transitively) forwarded to
584/// `b`.
585static bool isDefinedBefore(Operation *regionBranchOp, Value a, Value b) {
586 assert((b.getDefiningOp() == regionBranchOp ||
587 b.getParentRegion()->getParentOp() == regionBranchOp) &&
588 "b must be a region successor input");
589
590 // Case 1: `a` is defined inside of the region branch op. `a` must be
591 // directly nested in the region branch op. Otherwise, it could not have
592 // been among the reachable values for a region successor input.
593 if (a.getParentRegion()->getParentOp() == regionBranchOp) {
594 // Case 1.1: If `b` is a result of the region branch op, `a` is not in
595 // scope for `b`.
596 // Example:
597 // %b = region_op({
598 // ^bb0(%a1: ...):
599 // %a2 = ...
600 // })
601 if (isa<OpResult>(b))
602 return false;
603
604 // Case 1.2: `b` is an entry block argument of a region. `a` is in scope
605 // for `b` only if it is also an entry block argument of the same region.
606 // Example:
607 // region_op({
608 // ^bb0(%b: ..., %a: ...):
609 // ...
610 // })
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();
614 }
615
616 // Case 2: `a` is defined outside of the region branch op. In that case, we
617 // can safely assume that `a` was defined before `b`. Otherwise, it could not
618 // be among the reachable values for a region successor input.
619 // Example:
620 // { <- %a1 parent region begins here.
621 // ^bb0(%a1: ...):
622 // %a2 = ...
623 // %b1 = reigon_op({
624 // ^bb1(%b2: ...):
625 // ...
626 // })
627 // }
628 return true;
629}
630
631/// Compute all non-successor-input values that a successor input could have
632/// based on the given successor input to successor operand mapping.
633///
634/// Starting with the given value, trace back all predecessor values (i.e.,
635/// preceding successor operands) and add them to the set of reachable values.
636/// If the successor operand is again a successor input, do not add it to the
637/// result set, but instead continue the traversal.
638///
639/// If `maxReachableValues` is set, the traversal is aborted early and
640/// `failure` is returned as soon as the number of reachable values exceeds
641/// the limit. Otherwise, `success` is returned and the result set contains
642/// all reachable values.
643///
644/// Example 1:
645/// %r = scf.if ... {
646/// scf.yield %a : ...
647/// } else {
648/// scf.yield %b : ...
649/// }
650/// reachableValues(%r) = {%a, %b}
651///
652/// Example 2:
653/// %r = scf.for ... iter_args(%arg0 = %0) -> ... {
654/// scf.yield %arg0 : ...
655/// }
656/// reachableValues(%arg0) = {%0}
657/// reachableValues(%r) = {%0}
658///
659/// Example 3:
660/// %r = scf.for ... iter_args(%arg0 = %0) -> ... {
661/// ...
662/// scf.yield %1 : ...
663/// }
664/// reachableValues(%arg0) = {%0, %1}
665/// reachableValues(%r) = {%0, %1}
667 llvm::SmallDenseSet<Value> &result, Value value,
668 const RegionBranchInverseSuccessorMapping &inputToOperands,
669 std::optional<unsigned> maxReachableValues = std::nullopt) {
670 assert(inputToOperands.contains(value) && "value must be a successor input");
671 llvm::SmallDenseSet<Value> visited;
672 SmallVector<Value> worklist;
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()) {
678 result.insert(next);
679 if (maxReachableValues && result.size() > *maxReachableValues)
680 return failure();
681 continue;
682 }
683 for (OpOperand *operand : it->second)
684 if (visited.insert(operand->get()).second)
685 worklist.push_back(operand->get());
686 }
687 // Note: The result does not contain any successor inputs. (Therefore,
688 // `value` is also guaranteed to be excluded.)
689 return success();
690}
691
692namespace {
693/// Try to make successor inputs dead by replacing their uses with values that
694/// are not successor inputs. This pattern enables additional canonicalization
695/// opportunities for RemoveDeadRegionBranchOpSuccessorInputs.
696///
697/// Example:
698///
699/// %r0, %r1 = scf.for ... iter_args(%arg0 = %0, %arg1 = %1) -> ... {
700/// scf.yield %arg1, %arg1 : ...
701/// }
702/// use(%r0, %r1)
703///
704/// reachableValues(%r0) = {%0, %1}
705/// reachableValues(%r1) = {%1} ==> replace uses of %r1 with %1.
706/// reachableValues(%arg0) = {%0, %1}
707/// reachableValues(%arg1) = {%1} ==> replace uses of %arg1 with %1.
708///
709/// IR after pattern application:
710///
711/// %r0, %r1 = scf.for ... iter_args(%arg0 = %0, %arg1 = %1) -> ... {
712/// scf.yield %1, %1 : ...
713/// }
714/// use(%r0, %1)
715///
716/// Note that %r1 and %arg1 are dead now. The IR can now be further
717/// canonicalized by RemoveDeadRegionBranchOpSuccessorInputs.
718struct MakeRegionBranchOpSuccessorInputsDead : public RewritePattern {
719 MakeRegionBranchOpSuccessorInputsDead(MLIRContext *context, StringRef name,
720 PatternBenefit benefit = 1)
721 : RewritePattern(name, benefit, context) {}
722
723 LogicalResult matchAndRewrite(Operation *op,
724 PatternRewriter &rewriter) const override {
725 // Compute the mapping of successor inputs to successor operands.
726 auto regionBranchOp = cast<RegionBranchOpInterface>(op);
728 regionBranchOp.getSuccessorInputOperandMapping(inputToOperands);
729
730 // Try to replace the uses of each successor input one-by-one.
731 bool changed = false;
732 const bool isIsolated = op->hasTrait<OpTrait::IsIsolatedFromAbove>();
733 for (Value value : inputToOperands.keys()) {
734 // Nothing to do for successor inputs that are already dead.
735 if (value.use_empty())
736 continue;
737 // Nothing to do for successor inputs that may have multiple reachable
738 // values.
739 llvm::SmallDenseSet<Value> reachableValues;
741 reachableValues, value, inputToOperands,
742 /*maxReachableValues=*/1)) ||
743 reachableValues.empty())
744 continue;
745 Value replacement = *reachableValues.begin();
746 assert(replacement != value &&
747 "successor inputs are supposed to be excluded");
748 // A value inside an isolated region cannot be replaced with a value from
749 // another region, even if the replacement dominates its uses. Op results
750 // are in the parent region and do not cross this isolation boundary.
751 // Successor inputs are direct region arguments or results of this op.
752 // Valid input IR already prevents captures across nested isolation
753 // boundaries, so only this op's boundary needs an additional check.
754 // Example (isolated_op has IsolatedFromAbove):
755 // %r = isolated_op %x {
756 // ^bb0(%arg: ...):
757 // use(%arg)
758 // yield %arg
759 // }
760 // use(%r)
761 // Replacing %arg with %x would introduce a capture inside isolated_op.
762 // Replacing %r with %x changes only the outer use and is allowed.
763 Region *valueRegion = value.getParentRegion();
764 if (isIsolated && valueRegion->getParentOp() == op &&
765 replacement.getParentRegion() != valueRegion)
766 continue;
767 // Do not replace `value` with the found reachable value if doing so
768 // would violate dominance. Example:
769 // %r = scf.execute_region ... {
770 // %a = ...
771 // scf.yield %a : ...
772 // }
773 // use(%r)
774 // In the above example, reachableValues(%r) = {%a}, but %a cannot be
775 // used as a replacement for %r due to dominance / scope.
776 if (!isDefinedBefore(regionBranchOp, replacement, value))
777 continue;
778 rewriter.replaceAllUsesWith(value, replacement);
779 changed = true;
780 }
781 return success(changed);
782 }
783};
784
785/// Lookup a bit vector in the given mapping (DenseMap). If the key was not
786/// found, create a new bit vector with the given size and initialize it with
787/// false.
788template <typename MappingTy, typename KeyTy>
789static BitVector &lookupOrCreateBitVector(MappingTy &mapping, KeyTy key,
790 unsigned size) {
791 return mapping.try_emplace(key, size, false).first->second;
792}
793
794/// Compute tied successor inputs. Tied successor inputs are successor inputs
795/// that come as a set. If you erase one value from a set, you must erase all
796/// values from the set. Otherwise, the op would become structurally invalid.
797/// Each successor input appears in exactly one set.
798///
799/// Example:
800/// %r0, %r1 = scf.for ... iter_args(%arg0 = %0, %arg1 = %1) -> ... {
801/// ...
802/// }
803/// There are two sets: {{%r0, %arg0}, {%r1, %arg1}}.
804static llvm::EquivalenceClasses<Value> computeTiedSuccessorInputs(
805 const RegionBranchSuccessorMapping &operandToInputs) {
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)) {
812 // As we explore more successor operand to successor input mappings,
813 // existing sets may get merged.
814 tiedSuccessorInputs.unionSets(firstInput, nextInput);
815 }
816 }
817 return tiedSuccessorInputs;
818}
819
820/// Remove dead successor inputs from region branch ops. A successor input is
821/// dead if it has no uses. Successor inputs come in sets of tied values: if
822/// you remove one value from a set, you must remove all values from the set.
823/// Furthermore, successor operands must also be removed. (Op operands are not
824/// part of the set, but the set is built based on the successor operand to
825/// successor input mapping.)
826///
827/// Example 1:
828/// %r0, %r1 = scf.for ... iter_args(%arg0 = %0, %arg1 = %1) -> ... {
829/// scf.yield %0, %arg1 : ...
830/// }
831/// use(%0, %1)
832///
833/// There are two sets: {{%r0, %arg0}, {%r1, %arg1}}. All values in the first
834/// set are dead, so %arg0 and %r0 can be removed, but not %r1 and %arg1. The
835/// resulting IR is as follows:
836///
837/// %r1 = scf.for ... iter_args(%arg1 = %1) -> ... {
838/// scf.yield %arg1 : ...
839/// }
840/// use(%0, %1)
841///
842/// Example 2:
843/// %r0, %r1 = scf.while (%arg0 = %0) {
844/// scf.condition(...) %arg0, %arg0 : ...
845/// } do {
846/// ^bb0(%arg1: ..., %arg2: ...):
847/// scf.yield %arg1 : ...
848/// }
849/// There are three sets: {{%r0, %arg1}, {%r1, %arg2}, {%r0}}.
850///
851/// Example 3:
852/// %r1, %r2 = scf.if ... {
853/// scf.yield %0, %1 : ...
854/// } else {
855/// scf.yield %2, %3 : ...
856/// }
857/// There are two sets: {{%r1}, {%r2}}. Each set has one value, so there each
858/// value can be removed independently of the other values.
859struct RemoveDeadRegionBranchOpSuccessorInputs : public RewritePattern {
860 RemoveDeadRegionBranchOpSuccessorInputs(MLIRContext *context, StringRef name,
861 PatternBenefit benefit = 1)
862 : RewritePattern(name, benefit, context) {}
863
864 LogicalResult matchAndRewrite(Operation *op,
865 PatternRewriter &rewriter) const override {
866 // Compute tied values: values that must come as a set. If you remove one,
867 // you must remove all. If a successor op operand is forwarded to two
868 // successor inputs %a and %b, both %a and %b are in the same set.
869 auto regionBranchOp = cast<RegionBranchOpInterface>(op);
870 RegionBranchSuccessorMapping operandToInputs;
871 regionBranchOp.getSuccessorOperandInputMapping(operandToInputs);
872 llvm::EquivalenceClasses<Value> tiedSuccessorInputs =
873 computeTiedSuccessorInputs(operandToInputs);
874
875 // Determine which values to remove and group them by block and operation.
876 SmallVector<Value> valuesToRemove;
877 DenseMap<Block *, BitVector> blockArgsToRemove;
878 BitVector resultsToRemove(regionBranchOp->getNumResults(), false);
879 // Iterate over all sets of tied successor inputs.
880 for (auto it = tiedSuccessorInputs.begin(), e = tiedSuccessorInputs.end();
881 it != e; ++it) {
882 if (!(*it)->isLeader())
883 continue;
884
885 // Value can be removed if it is dead and all other tied values are also
886 // dead.
887 bool allDead = true;
888 for (auto memberIt = tiedSuccessorInputs.member_begin(**it);
889 memberIt != tiedSuccessorInputs.member_end(); ++memberIt) {
890 // Iterate over all values in the set and check their liveness.
891 if (!memberIt->use_empty()) {
892 allDead = false;
893 break;
894 }
895 }
896 if (!allDead)
897 continue;
898
899 // The entire set is dead. Group values by block and operation to
900 // simplify removal.
901 for (auto memberIt = tiedSuccessorInputs.member_begin(**it);
902 memberIt != tiedSuccessorInputs.member_end(); ++memberIt) {
903 if (auto arg = dyn_cast<BlockArgument>(*memberIt)) {
904 // Set blockArgsToRemove[block][arg_number] = true.
905 BitVector &vector =
906 lookupOrCreateBitVector(blockArgsToRemove, arg.getOwner(),
907 arg.getOwner()->getNumArguments());
908 vector.set(arg.getArgNumber());
909 } else {
910 // Set resultsToRemove[result_number] = true.
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());
915 }
916 valuesToRemove.push_back(*memberIt);
917 }
918 }
919
920 if (valuesToRemove.empty())
921 return rewriter.notifyMatchFailure(op, "no values to remove");
922
923 // Find operands that must be removed together with the values.
924 RegionBranchInverseSuccessorMapping inputsToOperands =
925 invertRegionBranchSuccessorMapping(operandToInputs);
927 for (Value value : valuesToRemove) {
928 for (OpOperand *operand : inputsToOperands[value]) {
929 // Set operandsToRemove[op][operand_number] = true.
930 BitVector &vector =
931 lookupOrCreateBitVector(operandsToRemove, operand->getOwner(),
932 operand->getOwner()->getNumOperands());
933 vector.set(operand->getOperandNumber());
934 }
935 }
936
937 // Erase operands.
938 for (auto &pair : operandsToRemove) {
939 Operation *op = pair.first;
940 BitVector &operands = pair.second;
941 rewriter.eraseOperands(op, operands);
942 }
943
944 // Erase block arguments.
945 for (auto &pair : blockArgsToRemove) {
946 Block *block = pair.first;
947 BitVector &blockArg = pair.second;
948 rewriter.modifyOpInPlace(block->getParentOp(),
949 [&]() { block->eraseArguments(blockArg); });
950 }
951
952 // Erase op results.
953 if (resultsToRemove.any())
954 rewriter.eraseOpResults(regionBranchOp, resultsToRemove);
955
956 return success();
957 }
958};
959
960/// Return the "owner" of a value: the parent block for block arguments, the
961/// defining op for op results.
962static void *getOwnerOfValue(Value value) {
963 if (auto arg = dyn_cast<BlockArgument>(value))
964 return arg.getOwner();
965 return value.getDefiningOp();
966}
967
968/// Get the block argument or op result number of the given value.
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();
973}
974
975/// Find duplicate successor inputs and make all dead except for one. Two
976/// successor inputs are "duplicate" if their corresponding successor operands
977/// have the same values. This pattern enables additional canonicalization
978/// opportunities for RemoveDeadRegionBranchOpSuccessorInputs.
979///
980/// Example:
981/// %r0, %r1 = scf.for ... iter_args(%arg0 = %0, %arg1 = %0) -> ... {
982/// use(%arg0, %arg1)
983/// ...
984/// scf.yield %x, %x : ...
985/// }
986/// use(%r0, %r1)
987///
988/// Operands of successor input %r0: [%0, %x]
989/// Operands of successor input %r1: [%0, %x] ==> DUPLICATE!
990/// Replace %r1 with %r0.
991///
992/// Operands of successor input %arg0: [%0, %x]
993/// Operands of successor input %arg1: [%0, %x] ==> DUPLICATE!
994/// Replace %arg1 with %arg0. (We have to make sure that we make same decision
995/// as for the other tied successor inputs above. Otherwise, a set of tied
996/// successor inputs may not become entirely dead.)
997///
998/// The resulting IR is as follows:
999/// %r0, %r1 = scf.for ... iter_args(%arg0 = %0, %arg1 = %0) -> ... {
1000/// use(%arg0, %arg0)
1001/// ...
1002/// scf.yield %x, %x : ...
1003/// }
1004/// use(%r0, %r0) // Note: We don't want use(%r1, %r1), which is also correct,
1005/// // but does not help with further canonicalizations.
1006struct RemoveDuplicateSuccessorInputUses : public RewritePattern {
1007 RemoveDuplicateSuccessorInputUses(MLIRContext *context, StringRef name,
1008 PatternBenefit benefit = 1)
1009 : RewritePattern(name, benefit, context) {}
1010
1011 LogicalResult matchAndRewrite(Operation *op,
1012 PatternRewriter &rewriter) const override {
1013 // Collect all successor inputs and sort them. When dropping the uses of a
1014 // successor input, we'd like to also drop the uses of the same tied
1015 // successor inputs. Otherwise, a set of tied successor inputs may not
1016 // become entirely dead, which is required for
1017 // RemoveDeadRegionBranchOpSuccessorInputs to be able to erase them.
1018 // (Sorting is not required for correctness.)
1019 auto regionBranchOp = cast<RegionBranchOpInterface>(op);
1020 RegionBranchInverseSuccessorMapping inputsToOperands;
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);
1025 });
1026
1027 // Group inputs by their operand "signature" to find duplicates. Two
1028 // successor inputs are duplicates if each predecessor (region branch point)
1029 // forwards the same value for both. Let n = number of successor inputs and
1030 // k = number of predecessors per input. Instead of comparing every pair of
1031 // inputs (O(n² * k)), we build a signature for each input and group them
1032 // via a std::map.
1033 //
1034 // A signature is a sorted list of (predecessor, forwarded value) pairs.
1035 // Within each group, all but the first (canonical) input are replaced with
1036 // the canonical one.
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;
1042 return a.second.getAsOpaquePointer() < b.second.getAsOpaquePointer();
1043 };
1044 // The map key is (signature, owner). Two inputs are duplicates only if they
1045 // have the same signature AND the same owner (block or defining op). This
1046 // ensures we track one canonical per owner group.
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(),
1053 sigEntryLess);
1054 };
1055 std::map<MapKey, Value, decltype(mapKeyLess)> signatureToCanonical(
1056 mapKeyLess);
1057 bool changed = false;
1058 // Total complexity: O(n * k * max(log k, log n)). For each input, sorting
1059 // the signature costs O(k log k) and the std::map lookup costs O(k log n).
1060 for (Value input : inputs) {
1061 // Gather the predecessor value for each predecessor (region branch
1062 // point) and sort them to form this input's signature.
1063 Signature sig;
1064 for (OpOperand *operand : inputsToOperands[input])
1065 sig.emplace_back(operand->getOwner(), operand->get());
1066 llvm::sort(sig, sigEntryLess);
1067
1068 void *owner = getOwnerOfValue(input);
1069
1070 auto [it, inserted] = signatureToCanonical.try_emplace(
1071 MapKey{std::move(sig), owner}, input);
1072 if (!inserted) {
1073 Value canonical = it->second;
1074 // Nothing to do if input is already dead.
1075 if (input.use_empty())
1076 continue;
1077 rewriter.replaceAllUsesWith(input, canonical);
1078 changed = true;
1079 }
1080 }
1081 return success(changed);
1082 }
1083};
1084
1085/// Given a range of values, return a vector of attributes of the same size,
1086/// where the i-th attribute is the constant value of the i-th value. If a
1087/// value is not constant, the corresponding attribute is null.
1088static SmallVector<Attribute> extractConstants(ValueRange values) {
1089 return llvm::map_to_vector(values, [](Value value) {
1090 Attribute attr;
1091 matchPattern(value, m_Constant(&attr));
1092 return attr;
1093 });
1094}
1095
1096/// Return all successor regions when branching from the given region branch
1097/// point. This helper functions extracts all constant operand values and
1098/// passes them to the `RegionBranchOpInterface`.
1100getSuccessorRegionsWithAttrs(RegionBranchOpInterface op,
1101 RegionBranchPoint point) {
1103 if (point.isParent()) {
1104 op.getEntrySuccessorRegions(extractConstants(op->getOperands()),
1105 successors);
1106 return successors;
1107 }
1108 RegionBranchTerminatorOpInterface terminator =
1110 terminator.getSuccessorRegions(extractConstants(terminator->getOperands()),
1111 successors);
1112 return successors;
1113}
1114
1115/// Find the single acyclic path through the given region branch op. Return an
1116/// empty vector if no such path or multiple such paths exist.
1117///
1118/// Example: "scf.if %true" has a single path:
1119/// parent => then_region => op
1120///
1121/// Example: "scf.if %cond" has multiple paths:
1122/// (1) parent => then_region => ancestor op
1123/// (2) parent => else_region => ancestor op
1124///
1125/// Example: "scf.while with scf.condition(%false)" has a single path:
1126/// parent => before_region => ancestor op
1127///
1128/// Example: "scf.for with 0 iterations" has a single path: parent => op
1129///
1130/// Note: Each path starts from the op. The initial parent branch point is
1131/// omitted from the result.
1132///
1133/// Note: This function also returns an "empty" path when a region with multiple
1134/// blocks was found.
1136computeSingleAcyclicRegionBranchPath(RegionBranchOpInterface op) {
1137 llvm::SmallDenseSet<Region *> visited;
1139
1140 // Path starts with "parent".
1142 do {
1143 SmallVector<RegionSuccessor> successors =
1144 getSuccessorRegionsWithAttrs(op, next);
1145 if (successors.size() != 1) {
1146 // There are multiple region successors. I.e., there are multiple paths
1147 // through the region branch op.
1148 return {};
1149 }
1150 path.push_back(successors.front());
1151 if (successors.front().isOperation()) {
1152 // Found path that ends after an operation.
1153 return path;
1154 }
1155 Region *region = successors.front().getSuccessor();
1156 if (!region->hasOneBlock()) {
1157 // Entering a region with multiple blocks. Such regions are not supported
1158 // at the moment.
1159 return {};
1160 }
1161 if (!visited.insert(region).second) {
1162 // We have already visited this region. I.e., we have found a cycle.
1163 return {};
1164 }
1165 auto terminator =
1166 dyn_cast<RegionBranchTerminatorOpInterface>(&region->front().back());
1167 if (!terminator) {
1168 // Region has no RegionBranchTerminatorOpInterface terminator. E.g., the
1169 // terminator could be a "ub.unreachable" op. Such IR is not supported.
1170 return {};
1171 }
1172 next = RegionBranchPoint(terminator);
1173 } while (true);
1174 llvm_unreachable("expected to return from loop");
1175}
1176
1177/// Inline the body of the matched region branch op into the enclosing block if
1178/// there is exactly one acyclic path through the region branch op, starting
1179/// from "parent", and if that path ends with "parent".
1180///
1181/// Example: This pattern can inline "scf.for" operations that are guaranteed to
1182/// have a single iteration, as indicated by the region branch path "parent =>
1183/// region => parent". "scf.for" operations have a non-successor-input: the loop
1184/// induction variable. Non-successor-input values have op-specific semantics
1185/// and cannot be reasoned about through the `RegionBranchOpInterface`. A
1186/// replacement value for non-successor-inputs is injected by the user-specified
1187/// lambda: in the case of the loop induction variable of an "scf.for", the
1188/// lower bound of the loop is used as a replacement value.
1189///
1190/// Before pattern application:
1191/// %r = scf.for %iv = %c5 to %c6 step %c1 iter_args(%arg0 = %0) {
1192/// %1 = "producer"(%arg0, %iv)
1193/// scf.yield %1
1194/// }
1195/// "user"(%r)
1196///
1197/// After pattern application:
1198/// %1 = "producer"(%0, %c5)
1199/// "user"(%1)
1200///
1201/// This pattern is limited to the following cases:
1202/// - Only regions with a single block are supported. This could be generalized.
1203/// - Region branch ops with side effects are not supported. (Recursive side
1204/// effects are fine.)
1205///
1206/// Note: This pattern queries the region dataflow from the
1207/// `RegionBranchOpInterface`. Replacement values are for block arguments / op
1208/// results are determined based on region dataflow. In case of
1209/// non-successor-inputs (whose values are not modeled by the
1210/// `RegionBranchOpInterface`), a user-specified lambda is queried.
1211struct InlineRegionBranchOp : public RewritePattern {
1212 InlineRegionBranchOp(MLIRContext *context, StringRef name,
1214 PatternMatcherFn matcherFn, PatternBenefit benefit = 1)
1215 : RewritePattern(name, benefit, context), replBuilderFn(replBuilderFn),
1216 matcherFn(matcherFn) {}
1217
1218 LogicalResult matchAndRewrite(Operation *op,
1219 PatternRewriter &rewriter) const override {
1220 // Check if the pattern is applicable to the given operation.
1221 if (failed(matcherFn(op)))
1222 return rewriter.notifyMatchFailure(op, "pattern not applicable");
1223
1224 // Patterns without recursive memory effects could have side effects, so
1225 // it is not safe to fold such ops away.
1226 if (!op->hasTrait<OpTrait::HasRecursiveMemoryEffects>())
1227 return rewriter.notifyMatchFailure(
1228 op, "pattern not applicable to ops without recursive memory effects");
1229
1230 // Find the single acyclic path through the region branch op.
1231 auto regionBranchOp = cast<RegionBranchOpInterface>(op);
1232 SmallVector<RegionSuccessor> path =
1233 computeSingleAcyclicRegionBranchPath(regionBranchOp);
1234 if (path.empty())
1235 return rewriter.notifyMatchFailure(
1236 op, "failed to find acyclic region branch path");
1237
1238 // Inline all regions on the path into the enclosing block.
1239 rewriter.setInsertionPoint(op);
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();
1245 ValueRange successorInputs =
1246 regionBranchOp.getSuccessorInputs(nextSuccessor);
1247 assert(successorInputs.size() == successorOperands.size() &&
1248 "size mismatch");
1249 // Find the index of the first block argument / op result that is a
1250 // succesor input.
1251 unsigned firstSuccessorInputIdx = 0;
1252 if (!successorInputs.empty())
1253 firstSuccessorInputIdx =
1254 nextSuccessor.isOperation()
1255 ? cast<OpResult>(successorInputs.front()).getResultNumber()
1256 : cast<BlockArgument>(successorInputs.front()).getArgNumber();
1257 // Query the total number of block arguments / op results.
1258 unsigned numValues =
1259 nextSuccessor.isOperation()
1260 ? nextSuccessor.getSuccessorOp()->getNumResults()
1261 : nextSuccessor.getSuccessor()->getNumArguments();
1262 // Compute replacement values for all block arguments / op results.
1263 SmallVector<Value> replacements;
1264 // Helper function to get the i-th block argument / op result.
1265 auto getValue = [&](unsigned idx) {
1266 return nextSuccessor.isOperation()
1267 ? Value(nextSuccessor.getSuccessorOp()->getResult(idx))
1268 : Value(nextSuccessor.getSuccessor()->getArgument(idx));
1269 };
1270 // Compute replacement values for all non-successor-input values that
1271 // precede the first successor input.
1272 for (unsigned i = 0; i < firstSuccessorInputIdx; ++i)
1273 replacements.push_back(
1274 replBuilderFn(rewriter, op->getLoc(), getValue(i)));
1275 // Use the successor operands of the predecessor as replacement values for
1276 // the successor inputs.
1277 llvm::append_range(replacements, successorOperands);
1278 // Compute replacement values for all block arguments / op results that
1279 // succeed the first successor input.
1280 for (unsigned i = replacements.size(); i < numValues; ++i)
1281 replacements.push_back(
1282 replBuilderFn(rewriter, op->getLoc(), getValue(i)));
1283 if (nextSuccessor.isOperation()) {
1284 // The path ends after the region branch op. Replace it with the
1285 // computed replacement values.
1286 if (nextSuccessor.getSuccessorOp() != op)
1287 return rewriter.notifyMatchFailure(
1288 op, "path ends after a different operation");
1289 assert(remainingPath.empty() && "expected that the path ended");
1290 rewriter.replaceOp(op, replacements);
1291 return success();
1292 }
1293 // We are inside of a region: query the successor operands from the
1294 // terminator, inline the region into the enclosing block, and erase the
1295 // terminator.
1296 auto terminator = cast<RegionBranchTerminatorOpInterface>(
1297 &nextSuccessor.getSuccessor()->front().back());
1298 rewriter.inlineBlockBefore(&nextSuccessor.getSuccessor()->front(),
1299 op->getBlock(), op->getIterator(),
1300 replacements);
1301 successorOperands = llvm::to_vector(
1302 terminator.getSuccessorOperands(remainingPath.front()));
1303 rewriter.eraseOp(terminator);
1304 }
1305
1306 llvm_unreachable("expected that path ends with an operation");
1307 }
1308
1310 PatternMatcherFn matcherFn;
1311};
1312} // namespace
1313
1315 RewritePatternSet &patterns, StringRef opName, PatternBenefit benefit) {
1316 patterns.add<MakeRegionBranchOpSuccessorInputsDead,
1317 RemoveDuplicateSuccessorInputUses,
1318 RemoveDeadRegionBranchOpSuccessorInputs>(patterns.getContext(),
1319 opName, benefit);
1320}
1321
1323 RewritePatternSet &patterns, StringRef opName,
1325 PatternMatcherFn matcherFn, PatternBenefit benefit) {
1326 patterns.add<InlineRegionBranchOp>(patterns.getContext(), opName,
1327 replBuilderFn, matcherFn, benefit);
1328}
return success()
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...
b
Return true if permutation is a valid permutation of the outer_dims_perm (case OuterOrInnerPerm::Oute...
*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.
Definition Attributes.h:25
Block represents an ordered list of Operations.
Definition Block.h:34
BlockArgument getArgument(unsigned i)
Definition Block.h:154
unsigned getNumArguments()
Definition Block.h:153
Operation & back()
Definition Block.h:177
Operation * getParentOp()
Returns the closest surrounding operation that contains this block.
Definition Block.cpp:31
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.
Definition ValueRange.h:119
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
Definition Builders.h:401
This class represents an operand of an operation.
Definition Value.h:254
This class implements the operand iterators for the Operation class.
Definition ValueRange.h:44
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.
Definition Operation.h:87
bool hasTrait()
Returns true if the operation was registered with a particular trait, e.g.
Definition Operation.h:801
unsigned getNumSuccessors()
Definition Operation.h:758
Block * getBlock()
Returns the operation block that contains this operation.
Definition Operation.h:230
OpResult getResult(unsigned idx)
Get the 'idx'th result of this operation.
Definition Operation.h:432
unsigned getNumRegions()
Returns the number of regions held by this operation.
Definition Operation.h:726
Location getLoc()
The source location the operation was defined or derived from.
Definition Operation.h:240
Operation * getParentOp()
Returns the closest surrounding operation that contains this operation or nullptr if this is a top-le...
Definition Operation.h:251
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'.
Definition Operation.h:255
OperationName getName()
The name of an operation is the key identifier for it.
Definition Operation.h:115
Block * getSuccessor(unsigned index)
Definition Operation.h:760
result_range getResults()
Definition Operation.h:440
Region * getParentRegion()
Returns the region to which the instruction belongs.
Definition Operation.h:247
unsigned getNumResults()
Return the number of results held by this operation.
Definition Operation.h:429
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.
Definition Region.h:26
Block & front()
Definition Region.h:65
BlockArgListType getArguments()
Definition Region.h:94
unsigned getRegionNumber()
Return the number of this region in the parent operation.
Definition Region.cpp:62
unsigned getNumArguments()
Definition Region.h:136
Operation * getParentOp()
Return the parent operation this region is attached to.
Definition Region.h:198
bool hasOneBlock()
Return true if this region has exactly one block.
Definition Region.h:68
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.
Definition TypeRange.h:40
Instances of the Type class are uniqued, have an immutable identifier and an optional mutable compone...
Definition Types.h:74
This class provides an abstraction over the different types of ranges over Values.
Definition ValueRange.h:389
type_range getTypes() const
This class represents an instance of an SSA value in the MLIR system, representing a computable value...
Definition Value.h:96
Type getType() const
Return the type of this value.
Definition Value.h:105
void * getAsOpaquePointer() const
Methods for supporting PointerLikeTypeTraits.
Definition Value.h:233
Operation * getDefiningOp() const
If this value is the result of an operation, return the operation that defines it.
Definition Value.cpp:18
Region * getParentRegion()
Return the Region in which this Value is defined.
Definition Value.cpp:39
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.
detail::InFlightRemark failed(Location loc, RemarkOpts opts)
Report an optimization remark that failed.
Definition Remarks.h:734
Include the generated interface declarations.
bool matchPattern(Value value, const Pattern &pattern)
Entry point for matching a pattern over a Value.
Definition Matchers.h:490
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
Definition LLVM.h:120
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.
Definition Matchers.h:369
llvm::function_ref< Fn > function_ref
Definition LLVM.h:147