1 //===- TosaToSCF.cpp - Lowering Tosa to SCF Dialect -----------------------===// 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 // These rewriters lower from the Tosa to the SCF dialect. 10 // 11 //===----------------------------------------------------------------------===// 12 13 #include "mlir/Conversion/TosaToSCF/TosaToSCF.h" 14 #include "mlir/Dialect/SCF/SCF.h" 15 #include "mlir/Dialect/Tensor/IR/Tensor.h" 16 #include "mlir/Dialect/Tosa/IR/TosaOps.h" 17 #include "mlir/IR/BlockAndValueMapping.h" 18 #include "mlir/IR/PatternMatch.h" 19 #include "mlir/Transforms/GreedyPatternRewriteDriver.h" 20 21 using namespace mlir; 22 using namespace tosa; 23 24 static void inlineIfCase(Region &srcRegion, Region &dstRegion, 25 OperandRange operands, PatternRewriter &rewriter) { 26 rewriter.cloneRegionBefore(srcRegion, &dstRegion.front()); 27 rewriter.eraseBlock(&dstRegion.back()); 28 29 Block *headBlock = &dstRegion.front(); 30 for (auto it : llvm::zip(headBlock->getArguments(), operands)) 31 std::get<0>(it).replaceAllUsesWith(std::get<1>(it)); 32 33 auto yield = cast<YieldOp>(headBlock->getTerminator()); 34 rewriter.setInsertionPoint(yield); 35 rewriter.create<scf::YieldOp>(yield.getLoc(), yield.inputs()); 36 rewriter.eraseOp(yield); 37 38 headBlock->eraseArguments( 39 llvm::to_vector<4>(llvm::seq<unsigned>(0, headBlock->getNumArguments()))); 40 } 41 42 static void inlineWhileCase(Region &srcRegion, Region &dstRegion, 43 PatternRewriter &rewriter, bool isCond) { 44 rewriter.cloneRegionBefore(srcRegion, &dstRegion.back()); 45 rewriter.eraseBlock(&dstRegion.back()); 46 47 Block *headBlock = &dstRegion.front(); 48 49 auto yield = cast<YieldOp>(headBlock->getTerminator()); 50 rewriter.setInsertionPoint(yield); 51 if (isCond) { 52 auto condition = 53 rewriter.create<tensor::ExtractOp>(yield.getLoc(), yield.getOperand(0)); 54 rewriter.create<scf::ConditionOp>(yield.getLoc(), condition, 55 headBlock->getArguments()); 56 } else { 57 rewriter.setInsertionPoint(yield); 58 rewriter.create<scf::YieldOp>(yield.getLoc(), yield.inputs()); 59 } 60 rewriter.eraseOp(yield); 61 } 62 63 namespace { 64 65 class IfOpConverter : public OpRewritePattern<tosa::IfOp> { 66 public: 67 using OpRewritePattern<tosa::IfOp>::OpRewritePattern; 68 69 LogicalResult matchAndRewrite(tosa::IfOp op, 70 PatternRewriter &rewriter) const final { 71 auto condition = rewriter.create<tensor::ExtractOp>(op.getLoc(), op.cond()); 72 auto newIf = rewriter.create<scf::IfOp>(op.getLoc(), op.getResultTypes(), 73 condition, true); 74 75 inlineIfCase(op.then_branch(), newIf.thenRegion(), op.inputs(), rewriter); 76 inlineIfCase(op.else_branch(), newIf.elseRegion(), op.inputs(), rewriter); 77 78 rewriter.replaceOp(op, newIf.getResults()); 79 return success(); 80 } 81 }; 82 83 class WhileOpConverter : public OpRewritePattern<tosa::WhileOp> { 84 public: 85 using OpRewritePattern<tosa::WhileOp>::OpRewritePattern; 86 87 LogicalResult matchAndRewrite(tosa::WhileOp op, 88 PatternRewriter &rewriter) const final { 89 auto newWhile = rewriter.create<scf::WhileOp>( 90 op.getLoc(), op.getResultTypes(), op.inputs()); 91 rewriter.createBlock(&newWhile.before()); 92 rewriter.createBlock(&newWhile.after()); 93 94 inlineWhileCase(op.cond(), newWhile.before(), rewriter, true); 95 inlineWhileCase(op.body(), newWhile.after(), rewriter, false); 96 97 rewriter.replaceOp(op, newWhile.getResults()); 98 99 return success(); 100 } 101 }; 102 103 } // namespace 104 105 void mlir::tosa::populateTosaToSCFConversionPatterns( 106 RewritePatternSet *patterns) { 107 patterns->add<IfOpConverter>(patterns->getContext()); 108 patterns->add<WhileOpConverter>(patterns->getContext()); 109 } 110