28#define GEN_PASS_DEF_SCFTOCONTROLFLOWPASS
29#include "mlir/Conversion/Passes.h.inc"
37struct SCFToControlFlowPass
40 void runOnOperation()
override;
105 using OpRewritePattern<ForOp>::OpRewritePattern;
107 LogicalResult matchAndRewrite(ForOp forOp,
108 PatternRewriter &rewriter)
const override;
198 using OpRewritePattern<IfOp>::OpRewritePattern;
200 LogicalResult matchAndRewrite(IfOp ifOp,
201 PatternRewriter &rewriter)
const override;
205 using OpRewritePattern<ExecuteRegionOp>::OpRewritePattern;
207 LogicalResult matchAndRewrite(ExecuteRegionOp op,
208 PatternRewriter &rewriter)
const override;
212 using OpRewritePattern<mlir::scf::ParallelOp>::OpRewritePattern;
214 LogicalResult matchAndRewrite(mlir::scf::ParallelOp parallelOp,
215 PatternRewriter &rewriter)
const override;
281 PatternRewriter &rewriter)
const override;
289 using OpRewritePattern<WhileOp>::OpRewritePattern;
291 LogicalResult matchAndRewrite(WhileOp whileOp,
292 PatternRewriter &rewriter)
const override;
299 LogicalResult matchAndRewrite(IndexSwitchOp op,
300 PatternRewriter &rewriter)
const override;
308 using OpRewritePattern<mlir::scf::ForallOp>::OpRewritePattern;
310 LogicalResult matchAndRewrite(mlir::scf::ForallOp forallOp,
311 PatternRewriter &rewriter)
const override;
320 return isa<LLVM::LLVMDialect>(attr.getValue().getDialect());
333LogicalResult ForLowering::matchAndRewrite(ForOp forOp,
335 Location loc = forOp.getLoc();
342 auto *endBlock = rewriter.
splitBlock(initBlock, initPosition);
348 auto *conditionBlock = &forOp.getRegion().
front();
349 auto *firstBodyBlock =
350 rewriter.
splitBlock(conditionBlock, conditionBlock->begin());
351 auto *lastBodyBlock = &forOp.getRegion().
back();
353 auto iv = conditionBlock->getArgument(0);
358 Operation *terminator = lastBodyBlock->getTerminator();
360 auto step = forOp.getStep();
361 auto stepped = arith::AddIOp::create(rewriter, loc, iv, step).getResult();
365 SmallVector<Value, 8> loopCarried;
366 loopCarried.push_back(stepped);
369 cf::BranchOp::create(rewriter, loc, conditionBlock, loopCarried);
376 Value lowerBound = forOp.getLowerBound();
377 Value upperBound = forOp.getUpperBound();
378 if (!lowerBound || !upperBound)
383 SmallVector<Value, 8> destOperands;
384 destOperands.push_back(lowerBound);
385 llvm::append_range(destOperands, forOp.getInitArgs());
386 cf::BranchOp::create(rewriter, loc, conditionBlock, destOperands);
390 arith::CmpIPredicate predicate = forOp.getUnsignedCmp()
391 ? arith::CmpIPredicate::ult
392 : arith::CmpIPredicate::slt;
394 arith::CmpIOp::create(rewriter, loc, predicate, iv, upperBound);
396 cf::CondBranchOp::create(rewriter, loc, comparison, firstBodyBlock,
397 ArrayRef<Value>(), endBlock, ArrayRef<Value>());
401 rewriter.
replaceOp(forOp, conditionBlock->getArguments().drop_front());
405LogicalResult IfLowering::matchAndRewrite(IfOp ifOp,
406 PatternRewriter &rewriter)
const {
407 auto loc = ifOp.getLoc();
414 auto *remainingOpsBlock = rewriter.
splitBlock(condBlock, opPosition);
415 Block *continueBlock;
416 if (ifOp.getNumResults() == 0) {
417 continueBlock = remainingOpsBlock;
420 rewriter.
createBlock(remainingOpsBlock, ifOp.getResultTypes(),
421 SmallVector<Location>(ifOp.getNumResults(), loc));
422 cf::BranchOp::create(rewriter, loc, remainingOpsBlock);
427 auto &thenRegion = ifOp.getThenRegion();
428 auto *thenBlock = &thenRegion.
front();
429 Operation *thenTerminator = thenRegion.back().getTerminator();
432 cf::BranchOp::create(rewriter, loc, continueBlock, thenTerminatorOperands);
433 rewriter.
eraseOp(thenTerminator);
439 auto *elseBlock = continueBlock;
440 auto &elseRegion = ifOp.getElseRegion();
441 if (!elseRegion.empty()) {
442 elseBlock = &elseRegion.front();
443 Operation *elseTerminator = elseRegion.back().getTerminator();
446 cf::BranchOp::create(rewriter, loc, continueBlock, elseTerminatorOperands);
447 rewriter.
eraseOp(elseTerminator);
452 cf::CondBranchOp::create(rewriter, loc, ifOp.getCondition(), thenBlock,
453 ArrayRef<Value>(), elseBlock,
462ExecuteRegionLowering::matchAndRewrite(ExecuteRegionOp op,
463 PatternRewriter &rewriter)
const {
464 auto loc = op.getLoc();
468 auto *remainingOpsBlock = rewriter.
splitBlock(condBlock, opPosition);
470 auto ®ion = op.getRegion();
472 cf::BranchOp::create(rewriter, loc, ®ion.front());
474 for (
Block &block : region) {
475 if (
auto terminator = dyn_cast<scf::YieldOp>(block.getTerminator())) {
478 cf::BranchOp::create(rewriter, loc, remainingOpsBlock,
486 SmallVector<Value> vals;
487 SmallVector<Location> argLocs(op.getNumResults(), op->getLoc());
489 remainingOpsBlock->addArguments(op->getResultTypes(), argLocs))
496ParallelLowering::matchAndRewrite(ParallelOp parallelOp,
497 PatternRewriter &rewriter)
const {
498 Location loc = parallelOp.getLoc();
499 auto reductionOp = dyn_cast<ReduceOp>(parallelOp.getBody()->getTerminator());
509 SmallVector<Value, 4> iterArgs = llvm::to_vector<4>(parallelOp.getInitVals());
510 SmallVector<Value, 4> ivs;
511 ivs.reserve(parallelOp.getNumLoops());
513 SmallVector<Value, 4> loopResults(iterArgs);
514 ForOp innermostForOp;
515 for (
auto [iv, lower, upper, step] :
516 llvm::zip(parallelOp.getInductionVars(), parallelOp.getLowerBound(),
517 parallelOp.getUpperBound(), parallelOp.getStep())) {
519 ForOp::create(rewriter, loc, lower, upper, step, iterArgs,
520 nullptr, parallelOp.getUnsignedCmp());
521 innermostForOp = forOp;
522 ivs.push_back(forOp.getInductionVar());
523 auto iterRange = forOp.getRegionIterArgs();
524 iterArgs.assign(iterRange.begin(), iterRange.end());
529 loopResults.assign(forOp.result_begin(), forOp.result_end());
531 }
else if (!forOp.getResults().empty()) {
535 scf::YieldOp::create(rewriter, loc, forOp.getResults());
548 SmallVector<Value> yieldOperands;
549 yieldOperands.reserve(parallelOp.getNumResults());
550 for (int64_t i = 0, e = parallelOp.getNumResults(); i < e; ++i) {
551 Block &reductionBody = reductionOp.getReductions()[i].front();
552 Value arg = iterArgs[yieldOperands.size()];
553 yieldOperands.push_back(
554 cast<ReduceReturnOp>(reductionBody.
getTerminator()).getResult());
557 {arg, reductionOp.getOperands()[i]});
563 if (newBody->
empty())
564 rewriter.
mergeBlocks(parallelOp.getBody(), newBody, ivs);
571 if (!yieldOperands.empty()) {
573 scf::YieldOp::create(rewriter, loc, yieldOperands);
576 rewriter.
replaceOp(parallelOp, loopResults);
582 PatternRewriter &rewriter)
const {
583 OpBuilder::InsertionGuard guard(rewriter);
584 Location loc = whileOp.getLoc();
588 Block *continuation =
592 Block *after = whileOp.getAfterBody();
593 Block *before = whileOp.getBeforeBody();
599 cf::BranchOp::create(rewriter, loc, before, whileOp.getInits());
606 SmallVector<Value> args = llvm::to_vector(condOp.getArgs());
608 after, condOp.getArgs(),
614 yieldOp.getResults());
625DoWhileLowering::matchAndRewrite(WhileOp whileOp,
626 PatternRewriter &rewriter)
const {
627 Block &afterBlock = *whileOp.getAfterBody();
628 if (!llvm::hasSingleElement(afterBlock))
630 "do-while simplification applicable "
631 "only if 'after' region has no payload");
633 auto yield = dyn_cast<scf::YieldOp>(&afterBlock.
front());
634 if (!yield || yield.getResults() != afterBlock.
getArguments())
636 "do-while simplification applicable "
637 "only to forwarding 'after' regions");
640 OpBuilder::InsertionGuard guard(rewriter);
642 Block *continuation =
646 Block *before = whileOp.getBeforeBody();
651 cf::BranchOp::create(rewriter, whileOp.getLoc(), before, whileOp.getInits());
656 auto latch = cf::CondBranchOp::create(
657 rewriter, condOp.getLoc(), condOp.getCondition(), before,
658 condOp.getArgs(), continuation,
ValueRange());
663 rewriter.
replaceOp(whileOp, condOp.getArgs());
671IndexSwitchLowering::matchAndRewrite(IndexSwitchOp op,
672 PatternRewriter &rewriter)
const {
679 SmallVector<Value> results;
680 results.reserve(op.getNumResults());
681 for (Type resultType : op.getResultTypes())
682 results.push_back(continueBlock->
addArgument(resultType, op.getLoc()));
685 auto convertRegion = [&](Region ®ion) -> FailureOr<Block *> {
686 Block *block = ®ion.front();
692 yield.getOperands());
700 SmallVector<Block *> caseSuccessors;
701 SmallVector<APInt> caseValues;
702 caseSuccessors.reserve(op.getCases().size());
703 caseValues.reserve(op.getCases().size());
704 for (
auto [region, value] : llvm::zip(op.getCaseRegions(), op.getCases())) {
705 FailureOr<Block *> block = convertRegion(region);
708 caseSuccessors.push_back(*block);
709 caseValues.push_back(APInt(64, value));
713 FailureOr<Block *> defaultBlock = convertRegion(op.getDefaultRegion());
719 SmallVector<ValueRange> caseOperands(caseSuccessors.size(), {});
722 Value caseValue = arith::IndexCastOp::create(
723 rewriter, op.getLoc(), rewriter.
getI64Type(), op.getArg());
725 cf::SwitchOp::create(rewriter, op.getLoc(), caseValue, *defaultBlock,
726 ValueRange(), caseValues, caseSuccessors, caseOperands);
727 rewriter.
replaceOp(op, continueBlock->getArguments());
731LogicalResult ForallLowering::matchAndRewrite(ForallOp forallOp,
732 PatternRewriter &rewriter)
const {
738 patterns.
add<ForallLowering, ForLowering, IfLowering, ParallelLowering,
744void SCFToControlFlowPass::runOnOperation() {
750 target.addIllegalOp<scf::ForallOp, scf::ForOp, scf::IfOp, scf::IndexSwitchOp,
751 scf::ParallelOp, scf::WhileOp, scf::ExecuteRegionOp>();
752 target.markUnknownOpDynamicallyLegal([](
Operation *) {
return true; });
753 ConversionConfig config;
754 config.allowPatternRollback = allowPatternRollback;
755 if (
failed(applyPartialConversion(getOperation(),
target, std::move(patterns),
static void propagateLoopAttrs(Operation *scfOp, Operation *brOp)
static void copyLLVMDialectAttrs(Operation *from, Operation *to)
OpListType::iterator iterator
Operation * getTerminator()
Get the terminator operation of this block.
BlockArgument addArgument(Type type, Location loc)
Add one value to the argument list.
BlockArgListType getArguments()
Block::iterator getInsertionPoint() const
Returns the current insertion point of the builder.
Block * createBlock(Region *parent, Region::iterator insertPt={}, TypeRange argTypes={}, ArrayRef< Location > locs={})
Add new block with 'argTypes' arguments and set the insertion point to the end of it.
void setInsertionPointToStart(Block *block)
Sets the insertion point to the start of the specified block.
void setInsertionPoint(Block *block, Block::iterator insertPoint)
Set the insertion point to the specified location.
void setInsertionPointToEnd(Block *block)
Sets the insertion point to the end of the specified block.
Block * getInsertionBlock() const
Return the block the current insertion point belongs to.
Operation is the basic unit of execution within MLIR.
operand_iterator operand_begin()
operand_iterator operand_end()
auto getDiscardableAttrs()
Return a range of all of discardable attributes on this operation.
operand_range getOperands()
Returns an iterator on the underlying Value's.
void setDiscardableAttrs(DictionaryAttr newAttrs)
Set the discardable attribute dictionary on this operation.
A special type of RewriterBase that coordinates the application of a rewrite pattern on the current I...
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.
Block * splitBlock(Block *block, Block::iterator before)
Split the operations starting at "before" (inclusive) out of the given block into a new block,...
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.
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.
void mergeBlocks(Block *source, Block *dest, ValueRange argValues={})
Inline the operations of block 'source' into the end of block 'dest'.
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 inlineRegionBefore(Region ®ion, Region &parent, Region::iterator before)
Move the blocks that belong to "region" before the given position in another region "parent".
OpTy replaceOpWithNewOp(Operation *op, Args &&...args)
Replace the results of the given (original) op with a new op that is created without verification (re...
LogicalResult forallToParallelLoop(RewriterBase &rewriter, ForallOp forallOp, ParallelOp *result=nullptr)
Try converting scf.forall into an scf.parallel loop.
Include the generated interface declarations.
void populateSCFToControlFlowConversionPatterns(RewritePatternSet &patterns)
Collect a set of patterns to convert SCF operations to CFG branch-based operations within the Control...
LogicalResult matchAndRewrite(WhileOp whileOp, OpAdaptor adaptor, ConversionPatternRewriter &rewriter) const override
OpRewritePattern is a wrapper around RewritePattern that allows for matching and rewriting against an...
OpRewritePattern(MLIRContext *context, PatternBenefit benefit=1, ArrayRef< StringRef > generatedNames={})