28#define GEN_PASS_DEF_SCFTOCONTROLFLOWPASS
29#include "mlir/Conversion/Passes.h.inc"
37struct SCFToControlFlowPass
38 :
public impl::SCFToControlFlowPassBase<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())) {
518 ForOp forOp = ForOp::create(rewriter, loc, lower, upper, step, iterArgs);
519 innermostForOp = forOp;
520 ivs.push_back(forOp.getInductionVar());
521 auto iterRange = forOp.getRegionIterArgs();
522 iterArgs.assign(iterRange.begin(), iterRange.end());
527 loopResults.assign(forOp.result_begin(), forOp.result_end());
529 }
else if (!forOp.getResults().empty()) {
533 scf::YieldOp::create(rewriter, loc, forOp.getResults());
546 SmallVector<Value> yieldOperands;
547 yieldOperands.reserve(parallelOp.getNumResults());
548 for (int64_t i = 0, e = parallelOp.getNumResults(); i < e; ++i) {
549 Block &reductionBody = reductionOp.getReductions()[i].front();
550 Value arg = iterArgs[yieldOperands.size()];
551 yieldOperands.push_back(
552 cast<ReduceReturnOp>(reductionBody.
getTerminator()).getResult());
555 {arg, reductionOp.getOperands()[i]});
561 if (newBody->
empty())
562 rewriter.
mergeBlocks(parallelOp.getBody(), newBody, ivs);
569 if (!yieldOperands.empty()) {
571 scf::YieldOp::create(rewriter, loc, yieldOperands);
574 rewriter.
replaceOp(parallelOp, loopResults);
580 PatternRewriter &rewriter)
const {
581 OpBuilder::InsertionGuard guard(rewriter);
582 Location loc = whileOp.getLoc();
586 Block *continuation =
590 Block *after = whileOp.getAfterBody();
591 Block *before = whileOp.getBeforeBody();
597 cf::BranchOp::create(rewriter, loc, before, whileOp.getInits());
604 SmallVector<Value> args = llvm::to_vector(condOp.getArgs());
606 after, condOp.getArgs(),
612 yieldOp.getResults());
623DoWhileLowering::matchAndRewrite(WhileOp whileOp,
624 PatternRewriter &rewriter)
const {
625 Block &afterBlock = *whileOp.getAfterBody();
626 if (!llvm::hasSingleElement(afterBlock))
628 "do-while simplification applicable "
629 "only if 'after' region has no payload");
631 auto yield = dyn_cast<scf::YieldOp>(&afterBlock.
front());
632 if (!yield || yield.getResults() != afterBlock.
getArguments())
634 "do-while simplification applicable "
635 "only to forwarding 'after' regions");
638 OpBuilder::InsertionGuard guard(rewriter);
640 Block *continuation =
644 Block *before = whileOp.getBeforeBody();
649 cf::BranchOp::create(rewriter, whileOp.getLoc(), before, whileOp.getInits());
654 auto latch = cf::CondBranchOp::create(
655 rewriter, condOp.getLoc(), condOp.getCondition(), before,
656 condOp.getArgs(), continuation,
ValueRange());
661 rewriter.
replaceOp(whileOp, condOp.getArgs());
669IndexSwitchLowering::matchAndRewrite(IndexSwitchOp op,
670 PatternRewriter &rewriter)
const {
677 SmallVector<Value> results;
678 results.reserve(op.getNumResults());
679 for (Type resultType : op.getResultTypes())
680 results.push_back(continueBlock->
addArgument(resultType, op.getLoc()));
683 auto convertRegion = [&](Region ®ion) -> FailureOr<Block *> {
684 Block *block = ®ion.front();
690 yield.getOperands());
698 SmallVector<Block *> caseSuccessors;
699 SmallVector<APInt> caseValues;
700 caseSuccessors.reserve(op.getCases().size());
701 caseValues.reserve(op.getCases().size());
702 for (
auto [region, value] : llvm::zip(op.getCaseRegions(), op.getCases())) {
703 FailureOr<Block *> block = convertRegion(region);
706 caseSuccessors.push_back(*block);
707 caseValues.push_back(APInt(64, value));
711 FailureOr<Block *> defaultBlock = convertRegion(op.getDefaultRegion());
717 SmallVector<ValueRange> caseOperands(caseSuccessors.size(), {});
720 Value caseValue = arith::IndexCastOp::create(
721 rewriter, op.getLoc(), rewriter.
getI64Type(), op.getArg());
723 cf::SwitchOp::create(rewriter, op.getLoc(), caseValue, *defaultBlock,
724 ValueRange(), caseValues, caseSuccessors, caseOperands);
725 rewriter.
replaceOp(op, continueBlock->getArguments());
729LogicalResult ForallLowering::matchAndRewrite(ForallOp forallOp,
730 PatternRewriter &rewriter)
const {
736 patterns.
add<ForallLowering, ForLowering, IfLowering, ParallelLowering,
742void SCFToControlFlowPass::runOnOperation() {
748 target.addIllegalOp<scf::ForallOp, scf::ForOp, scf::IfOp, scf::IndexSwitchOp,
749 scf::ParallelOp, scf::WhileOp, scf::ExecuteRegionOp>();
750 target.markUnknownOpDynamicallyLegal([](
Operation *) {
return true; });
751 ConversionConfig config;
752 config.allowPatternRollback = allowPatternRollback;
753 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={})