14#include "mlir/Analysis/CFGLoopInfo.h"
15#include "mlir/Dialect/ControlFlow/IR/ControlFlowOps.h"
16#include "mlir/IR/Dominance.h"
17#include "mlir/IR/IRMapping.h"
18#include "mlir/IR/Matchers.h"
19#include "mlir/Pass/Pass.h"
20#include "llvm/ADT/PostOrderIterator.h"
21#include "llvm/Support/Debug.h"
23#define DEBUG_TYPE "llhd-unroll-loops"
27#define GEN_PASS_DEF_UNROLLLOOPSPASS
28#include "circt/Dialect/LLHD/LLHDPasses.h.inc"
35using llvm::SmallDenseSet;
45static void cloneBlocks(ArrayRef<Block *> blocks, Region ®ion,
46 Region::iterator before, IRMapping &mapper) {
54 SmallVector<Block *> newBlocks;
55 newBlocks.reserve(blocks.size());
56 for (
auto *block : blocks) {
57 auto *newBlock =
new Block();
58 mapper.map(block, newBlock);
59 for (
auto arg : block->getArguments())
60 mapper.map(arg, newBlock->addArgument(arg.getType(), arg.getLoc()));
61 region.getBlocks().insert(before, newBlock);
62 newBlocks.push_back(newBlock);
71 Operation::CloneOptions::all().cloneRegions(
false).cloneOperands(
false);
72 for (
auto [oldBlock, newBlock] : llvm::zip(blocks, newBlocks))
73 for (
auto &op : *oldBlock)
74 newBlock->push_back(op.clone(mapper, cloneOptions));
78 SmallVector<Value> operands;
79 for (
auto [oldBlock, newBlock] : llvm::zip(blocks, newBlocks)) {
80 for (
auto [oldOp, newOp] : llvm::zip(*oldBlock, *newBlock)) {
81 operands.resize(oldOp.getNumOperands());
83 oldOp.getOperands(), operands.begin(),
84 [&](Value operand) { return mapper.lookupOrDefault(operand); });
85 newOp.setOperands(operands);
86 for (
auto [oldRegion, newRegion] :
87 llvm::zip(oldOp.getRegions(), newOp.getRegions()))
88 oldRegion.cloneInto(&newRegion, mapper);
100 Loop(
unsigned loopId, CFGLoop &cfgLoop) : loopId(loopId), cfgLoop(cfgLoop) {}
101 bool failMatch(
const Twine &msg)
const;
103 void unroll(CFGLoopInfo &cfgLoopInfo);
110 BlockOperand *exitEdge =
nullptr;
121 comb::ICmpPredicate predicate;
123 APInt indVarIncrement;
129 unsigned tripCount = 0;
133static llvm::raw_ostream &
operator<<(llvm::raw_ostream &os,
const Loop &loop) {
134 os <<
"#" << loop.loopId <<
" from ";
135 loop.cfgLoop.getHeader()->printAsOperand(os);
137 loop.cfgLoop.getLoopLatch()->printAsOperand(os);
142bool Loop::failMatch(
const Twine &msg)
const {
143 LLVM_DEBUG(llvm::dbgs() <<
"- Ignoring loop " << *
this <<
": " << msg
152 SmallVector<BlockOperand *> exits;
153 for (
auto *block : cfgLoop.getBlocks())
154 for (auto &edge : block->getTerminator()->getBlockOperands())
155 if (!cfgLoop.contains(edge.
get()))
156 exits.push_back(&edge);
157 if (exits.size() != 1)
158 return failMatch(
"multiple exits");
159 exitEdge = exits.back();
162 auto exitBranch = dyn_cast<cf::CondBranchOp>(exitEdge->getOwner());
164 return failMatch(
"unsupported exit branch");
165 exitCondition = exitBranch.getCondition();
166 exitInverted = exitEdge->getOperandNumber() == 1;
170 if (
auto icmpOp = exitCondition.getDefiningOp<comb::ICmpOp>()) {
171 IntegerAttr boundAttr;
172 if (!matchPattern(icmpOp.getRhs(), m_Constant(&boundAttr)))
173 return failMatch(
"non-constant loop bound");
174 indVar = icmpOp.getLhs();
175 predicate = icmpOp.getPredicate();
176 endBound = boundAttr.getValue();
178 return failMatch(
"unsupported exit condition");
184 predicate = comb::ICmpOp::getNegatedPredicate(predicate);
187 auto *header = cfgLoop.getHeader();
188 auto *latch = cfgLoop.getLoopLatch();
189 auto indVarArg = dyn_cast<BlockArgument>(indVar);
190 if (!indVarArg || indVarArg.getOwner() != header)
191 return failMatch(
"induction variable is not a header block argument");
192 IntegerAttr beginBoundAttr;
193 for (
auto &pred : header->getUses()) {
194 auto branchOp = dyn_cast<BranchOpInterface>(pred.getOwner());
196 return failMatch(
"header predecessor terminator is not a branch op");
197 auto indVarValue = branchOp.getSuccessorOperands(
198 pred.getOperandNumber())[indVarArg.getArgNumber()];
199 IntegerAttr boundAttr;
200 if (pred.getOwner()->getBlock() == latch) {
201 indVarNext = indVarValue;
202 }
else if (matchPattern(indVarValue, m_Constant(&boundAttr))) {
204 beginBoundAttr = boundAttr;
205 else if (boundAttr != beginBoundAttr)
206 return failMatch(
"multiple initial bounds");
208 return failMatch(
"unsupported induction variable value");
212 return failMatch(
"no initial bound");
213 beginBound = beginBoundAttr.getValue();
216 if (
auto addOp = indVarNext.getDefiningOp<
comb::AddOp>();
217 addOp && addOp.getNumOperands() == 2) {
218 if (addOp.getOperand(0) != indVarArg)
219 return failMatch(
"increment LHS not the induction variable");
221 if (!matchPattern(addOp.getOperand(1), m_Constant(&incAttr)))
222 return failMatch(
"increment RHS non-constant");
223 indVarIncrement = incAttr.getValue();
225 return failMatch(
"unsupported increment");
228 std::optional<unsigned> range;
231 if (predicate == comb::ICmpPredicate::ult && beginBound.ule(endBound) &&
232 indVarIncrement.sgt(0)) {
233 range = endBound.getZExtValue() - beginBound.getZExtValue();
236 if (predicate == comb::ICmpPredicate::slt && !endBound.isNegative() &&
237 beginBound.sle(endBound) && indVarIncrement.sgt(0)) {
238 range = endBound.getZExtValue() - beginBound.getZExtValue();
241 if (predicate == comb::ICmpPredicate::sgt && !beginBound.isNegative() &&
242 endBound.sle(beginBound) && indVarIncrement.isNegative()) {
243 if (!endBound.isNegative())
244 range = beginBound.getZExtValue() - endBound.getZExtValue();
246 else if (endBound.isAllOnes())
247 range = beginBound.getZExtValue() + 1;
250 if (predicate == comb::ICmpPredicate::eq && indVarIncrement != 0 &&
251 beginBound == endBound) {
256 if (!range.has_value())
257 return failMatch(
"unsupported loop bounds");
260 unsigned stride = indVarIncrement.abs().getZExtValue();
261 tripCount = (*range + stride - 1) / stride;
263 if (tripCount >= 1024)
264 return failMatch(
"unsupported loop bounds");
271void Loop::unroll(CFGLoopInfo &cfgLoopInfo) {
272 LLVM_DEBUG(llvm::dbgs() <<
"- Unrolling loop " << *
this <<
"\n");
277 auto *header = cfgLoop.getHeader();
278 SmallVector<Block *> orderedBody;
279 for (
auto &block : *header->getParent())
280 if (cfgLoop.contains(&block))
281 orderedBody.push_back(&block);
284 auto *latch = cfgLoop.getLoopLatch();
285 OpBuilder builder(indVar.getContext());
286 auto indValue = beginBound;
287 for (
unsigned trip = 0; trip < tripCount; ++trip) {
290 cloneBlocks(orderedBody, *header->getParent(), header->getIterator(),
292 auto *clonedHeader = mapper.lookup(header);
293 auto *clonedTail = mapper.lookup(latch);
296 auto iterIndVar = mapper.lookup(indVar);
298 builder.setInsertionPointAfterValue(iterIndVar);
299 iterIndVar.replaceAllUsesWith(
304 for (
auto &blockOperand :
llvm::make_early_inc_range(header->getUses()))
305 if (blockOperand.getOwner()->getBlock() != latch)
306 blockOperand.set(clonedHeader);
310 for (
auto &blockOperand : clonedTail->getTerminator()->getBlockOperands())
311 if (blockOperand.
get() == clonedHeader)
312 blockOperand.set(header);
317 cast<cf::CondBranchOp>(mapper.lookup(exitEdge->getOwner()));
318 Block *continueDest = exitBranchOp.getTrueDest();
319 ValueRange continueDestOperands = exitBranchOp.getTrueDestOperands();
320 if (exitEdge->getOperandNumber() == 0) {
321 continueDest = exitBranchOp.getFalseDest();
322 continueDestOperands = exitBranchOp.getFalseDestOperands();
324 builder.setInsertionPoint(exitBranchOp);
325 cf::BranchOp::create(builder, exitBranchOp.getLoc(), continueDest,
326 continueDestOperands);
328 exitBranchOp.erase();
331 for (
auto *block : orderedBody) {
332 auto *newBlock = mapper.lookup(block);
333 cfgLoop.addBasicBlockToLoop(newBlock, cfgLoopInfo);
337 indValue += indVarIncrement;
344 builder.setInsertionPointAfterValue(indVar);
345 indVar.replaceAllUsesWith(
351 auto exitBranchOp = cast<cf::CondBranchOp>(exitEdge->getOwner());
352 Block *exitDest = exitBranchOp.getTrueDest();
353 ValueRange exitDestOperands = exitBranchOp.getTrueDestOperands();
354 if (exitEdge->getOperandNumber() == 1) {
355 exitDest = exitBranchOp.getFalseDest();
356 exitDestOperands = exitBranchOp.getFalseDestOperands();
358 builder.setInsertionPoint(exitBranchOp);
359 cf::BranchOp::create(builder, exitBranchOp.getLoc(), exitDest,
362 exitBranchOp.erase();
366 SmallPtrSet<Block *, 8> blocksToPrune;
367 for (
auto *block : cfgLoop.getBlocks())
368 if (block->use_empty())
369 blocksToPrune.insert(block);
370 while (!blocksToPrune.empty()) {
371 auto *block = *blocksToPrune.begin();
372 blocksToPrune.erase(block);
373 if (!block->use_empty())
375 for (
auto *succ : block->getSuccessors())
376 if (cfgLoop.contains(succ))
377 blocksToPrune.insert(succ);
378 block->dropAllDefinedValueUses();
379 cfgLoopInfo.removeBlock(block);
388 for (
auto &block : *header->getParent()) {
389 if (!cfgLoop.contains(&block))
392 auto branchOp = dyn_cast<cf::BranchOp>(block.getTerminator());
395 auto *otherBlock = branchOp.getDest();
396 if (!cfgLoop.contains(otherBlock) || !otherBlock->getSinglePredecessor())
398 for (
auto [blockArg, branchArg] :
399 llvm::zip(otherBlock->getArguments(), branchOp.getDestOperands()))
400 blockArg.replaceAllUsesWith(branchArg);
401 block.getOperations().splice(branchOp->getIterator(),
402 otherBlock->getOperations());
404 cfgLoopInfo.removeBlock(otherBlock);
415struct UnrollLoopsPass
416 :
public llhd::impl::UnrollLoopsPassBase<UnrollLoopsPass> {
417 void runOnOperation()
override;
418 void runOnOperation(CombinationalOp op);
422void UnrollLoopsPass::runOnOperation() {
423 for (
auto op : getOperation().getOps<CombinationalOp>())
427void UnrollLoopsPass::runOnOperation(CombinationalOp op) {
430 if (op.getBody().hasOneBlock())
434 LLVM_DEBUG(llvm::dbgs() <<
"Unrolling loops in " << op.getLoc() <<
"\n");
435 DominanceInfo domInfo(op);
436 CFGLoopInfo cfgLoopInfo(domInfo.getDomTree(&op.getBody()));
442 SmallVector<Loop> loops;
443 for (
auto *cfgLoop : cfgLoopInfo.getLoopsInPreorder()) {
446 auto *header = cfgLoop->getHeader();
447 auto *latch = cfgLoop->getLoopLatch();
452 llvm::dbgs() <<
"- ";
453 cfgLoop->print(llvm::dbgs(),
false,
false);
454 llvm::dbgs() <<
"\n";
456 Loop loop(loops.size(), *cfgLoop);
460 auto *parent = cfgLoop->getParentLoop();
461 while (parent && parent->getHeader() != header)
462 parent = parent->getParentLoop();
464 loop.failMatch(
"header block shared across multiple loops");
470 parent = cfgLoop->getParentLoop();
471 while (parent && !parent->isLoopLatch(latch))
472 parent = parent->getParentLoop();
474 loop.failMatch(
"latch block shared across multiple loops");
480 loops.push_back(std::move(loop));
488 auto &os = llvm::dbgs();
489 for (
auto &loop : loops) {
490 os <<
"- Loop " << loop <<
":\n";
492 loop.cfgLoop.print(os,
false,
false);
495 loop.exitEdge->get()->printAsOperand(os);
497 if (loop.exitInverted)
499 os << loop.exitCondition;
501 os <<
" - Induction variable: ";
502 loop.indVar.printAsOperand(os, OpPrintingFlags().useLocalScope());
503 os <<
", from " << loop.beginBound <<
", while " << loop.predicate <<
" "
504 << loop.endBound <<
", increment " << loop.indVarIncrement <<
"\n";
505 os <<
" - Trip count: " << loop.tripCount <<
"\n";
511 for (
auto &loop :
llvm::reverse(loops))
512 loop.unroll(cfgLoopInfo);
static void cloneBlocks(ArrayRef< Block * > blocks, Region ®ion, Region::iterator before, IRMapping &mapper)
Clone a list of blocks into a region before the given block.
Direction get(bool isOutput)
Returns an output direction if isOutput is true, otherwise returns an input direction.
OS & operator<<(OS &os, const InnerSymTarget &target)
Printing InnerSymTarget's.
The InstanceGraph op interface, see InstanceGraphInterface.td for more details.
Utility that tracks operations that have potentially become unused and allows them to be cleaned up a...
void eraseLaterIfUnused(Operation *op)
Mark an op the be erased later if it is unused at that point.
void eraseNow(Operation *op)
Erase an operation immediately, and remove it from the set of ops to be removed later.