[llvm] [Transforms][Utils] Add LoopSplitUtils for iteration-space loop splitting (PR #205995)
Florian Hahn via llvm-commits
llvm-commits at lists.llvm.org
Mon Aug 3 12:00:44 PDT 2026
================
@@ -0,0 +1,617 @@
+//===- LoopSplitUtils.cpp - Split a loop's iteration space ----------------===//
+//
+// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
+// See https://llvm.org/LICENSE.txt for license information.
+// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
+//
+//===----------------------------------------------------------------------===//
+//
+// Splits a counted loop's iteration space into a chain of per-partition
+// sub-loops. See LoopSplitUtils.h for the high-level usage guidelines.
+//
+// Structure produced for partitions [S0,E0], [S1,E1], ... where E is the loop's
+// last iteration and each clamped end sel_i = min(E_i, E):
+//
+// guard0: ; every S_i and sel_i is computed here
+// if (S0 <= sel0) goto preheader0 else goto guard1 ; default guard check
+// loop0: ... ; latch stops at sel0
+// exit0 -> guard1
+// guard1:
+// if (S1 <= sel1) goto preheader1 else goto guard2 ; default guard check
+// loop1: ... ; latch stops at sel1
+// exit1 -> guard2
+// ...
+// final.exit: ; merges every partition's live-outs
+//
+// Each guard holds the "S_i <= sel_i" check and skips an empty partition by
+// falling through to the next guard. The check is replaced by an unconditional
+// branch when a partition is proven empty (to the next guard) or the caller
+// exempts it via avoidPartitionGuard() (to its preheader). All S_i/sel_i are
+// materialized once in guard0; the end clamp keeps the "runs at least once"
+// iteration in the right partition; live-outs are rebuilt one SSAUpdater each.
+//
+// A descending (step -1) loop uses the same structure mirrored: partitions run
+// high-to-low and the empty test, clamp, and predicates flip (>=/>).
+//
+// Usage guidelines:
+// - Caller bounds must not wrap the induction type. The clamp absorbs a bound
+// past the runtime trip count, but a Start +/- offset that overshoots the
+// type extreme wraps in the bound arithmetic and cannot be repaired here.
+// - Bounds must be loop-invariant: they are expanded in guard0 (the
+// preheader),
+// so a bound depending on a value defined inside the loop cannot be placed.
+// - The partitions must tile the original iteration space exactly -- same
+// iterations, same order -- so the split preserves program behaviour.
+// - A caller that drops a guard via avoidPartitionGuard() must itself ensure
+// that partition runs at least once, or the result is a spurious iteration.
+//
+//===----------------------------------------------------------------------===//
+
+#include "llvm/Transforms/Utils/LoopSplitUtils.h"
+#include "llvm/ADT/DenseMap.h"
+#include "llvm/Analysis/LoopInfo.h"
+#include "llvm/Analysis/ScalarEvolution.h"
+#include "llvm/Analysis/ScalarEvolutionExpressions.h"
+#include "llvm/Analysis/ScalarEvolutionPatternMatch.h"
+#include "llvm/IR/BasicBlock.h"
+#include "llvm/IR/CFG.h"
+#include "llvm/IR/Constants.h"
+#include "llvm/IR/Dominators.h"
+#include "llvm/IR/Function.h"
+#include "llvm/IR/IRBuilder.h"
+#include "llvm/IR/Instructions.h"
+#include "llvm/Support/Debug.h"
+#include "llvm/Transforms/Utils/BasicBlockUtils.h"
+#include "llvm/Transforms/Utils/Cloning.h"
+#include "llvm/Transforms/Utils/LoopUtils.h"
+#include "llvm/Transforms/Utils/SSAUpdater.h"
+#include "llvm/Transforms/Utils/ScalarEvolutionExpander.h"
+#include "llvm/Transforms/Utils/ValueMapper.h"
+#include <optional>
+
+using namespace llvm;
+using namespace llvm::SCEVPatternMatch;
+
+#define DEBUG_TYPE "loop-split-utils"
+
+//===----------------------------------------------------------------------===//
+// LoopSplitUtils - construction, partition list, induction analysis
+//===----------------------------------------------------------------------===//
+
+/// Per-split() scratch shared by the phase helpers; lives for one split() call.
+struct LoopSplitUtils::SplitState {
+ // Partition 0 reuses the original loop's preheader, exit, and entry guard;
+ // those blocks live in Partitions[0] rather than being duplicated here.
+ BasicBlock *FinalExit = nullptr; // where live-outs merge.
+ Loop *OuterLoop = nullptr; // parent of the new blocks, if any.
+ PHINode *Induction = nullptr; // the loop's induction variable.
+ bool Descending = false; // step is negative (loop counts down).
+ bool LatchComparesPHI = false; // latch compares the PHI, not the step.
+
+ /// A value that must be reconstructed after cloning because it is
+ /// loop-carried (feeds a later partition), live-out (used after the loop), or
+ /// both.
+ struct EscapingValue {
+ EscapingValue() = default;
+ EscapingValue(Value *Def) : Def(Def) {}
+
+ /// The value as it exists in partition 0 (the original).
+ Value *Def = nullptr;
+ /// The carried header PHI in partition 0, or null if \c Def needs no
+ /// per-partition start value seeded.
+ PHINode *CarriedHeaderPHI = nullptr;
+ /// True if \c Def is used outside the loop and must be merged at the final
+ /// exit.
+ bool EscapesOutside = false;
+ /// \c Def and \c CarriedHeaderPHI cloned into each partition (index 0 is
+ /// the original; \c PerPartitionPHI[0] is unused).
+ SmallVector<Value *, 4> PerPartitionDef;
+ SmallVector<PHINode *, 4> PerPartitionPHI;
+ };
+
+ /// Values that must survive across partitions (carried and/or live-out).
+ SmallVector<EscapingValue, 8> Escaping;
+
+ EscapingValue &addEscaping(Value *Def) { return Escaping.emplace_back(Def); }
+};
+
+// Record a new partition with the given inclusive iteration range.
+void LoopSplitUtils::addPartition(const SCEV *Start, const SCEV *End) {
+ Partitions.emplace_back(Start, End);
+}
+
+// Mark a partition so split() emits no entry guard for it.
+void LoopSplitUtils::avoidPartitionGuard(unsigned PartitionIndex) {
+ assert(PartitionIndex < Partitions.size() &&
+ "avoidPartitionGuard() called for an unknown partition");
+ Partitions[PartitionIndex].Guarded = false;
+}
+
+// Return a partition's original-to-clone map, or null if it has none.
+const ValueToValueMapTy *
+LoopSplitUtils::getPartitionValueMap(unsigned PartitionIndex) const {
+ if (PartitionIndex >= Partitions.size())
+ return nullptr;
+ return Partitions[PartitionIndex].VMap.get();
+}
+
+// Look up the counterpart of an original value in a given partition.
+Value *LoopSplitUtils::getPartitionValue(Value *V,
+ unsigned PartitionIndex) const {
+ assert(PartitionIndex < getNumPartitions() && "partition index out of range");
+ // Partition 0 reuses the original loop: every value maps to itself.
+ if (PartitionIndex == 0)
+ return V;
+ const ValueToValueMapTy *VMap = getPartitionValueMap(PartitionIndex);
+ if (!VMap)
+ return nullptr;
+ return VMap->lookup(V);
+}
+
+// Find the induction variable and the latch operand it is compared against;
+// returns the induction's add-recurrence, or null if the loop is unsuitable.
+// On success \p LatchIndOperand is set to the compared induction operand.
+static const SCEVAddRecExpr *analyzeInduction(Loop *L, ScalarEvolution *SE,
+ Value *&LatchIndOperand) {
+ ICmpInst *LatchCmp = L->getLatchCmpInst();
+
+ // SCEV's induction variable, restricted to a unit-step affine recurrence.
+ PHINode *Induction = L->getInductionVariable(*SE);
+ if (!Induction)
+ return nullptr;
+ const SCEV *IndSCEV = SE->getSCEV(Induction);
+ // Match an affine add-recurrence and capture its constant step; accept a unit
+ // step in either direction: +1 (ascending) or -1 (descending).
+ const APInt *Step;
+ if (!match(IndSCEV, m_scev_AffineAddRec(m_SCEV(), m_scev_APInt(Step))))
+ return nullptr;
+ if (!Step->isOne() && !Step->isAllOnes())
+ return nullptr;
+ const auto *AR = cast<SCEVAddRecExpr>(IndSCEV);
+
+ // The induction's "next" value (i + 1), produced in the latch.
+ auto *StepInst = dyn_cast<Instruction>(
+ Induction->getIncomingValueForBlock(L->getLoopLatch()));
+ if (!StepInst)
+ return nullptr;
+
+ // Select the compare operand that is the induction (PHI or its step).
+ if (LatchCmp->getOperand(0) == Induction ||
+ LatchCmp->getOperand(0) == StepInst)
+ LatchIndOperand = LatchCmp->getOperand(0);
+ else if (LatchCmp->getOperand(1) == Induction ||
+ LatchCmp->getOperand(1) == StepInst)
+ LatchIndOperand = LatchCmp->getOperand(1);
+ else
+ return nullptr;
+ return AR;
+}
+
+// Decide whether the iteration ordering is signed or unsigned; returns the
+// signedness, or nullopt if it cannot be proven.
+static std::optional<bool> computeSignedness(Loop *L,
+ const SCEVAddRecExpr *IndAR) {
+ ICmpInst::Predicate P = L->getLatchCmpInst()->getPredicate();
+ // A relational predicate gives the ordering directly; for eq/ne fall back to
+ // the recurrence's no-wrap flags.
+ if (ICmpInst::isRelational(P))
+ return ICmpInst::isSigned(P);
+ if (IndAR->hasNoSignedWrap())
+ return true;
+ if (IndAR->hasNoUnsignedWrap())
+ return false;
+ LLVM_DEBUG(dbgs() << DEBUG_TYPE
+ ": cannot prove iteration ordering signedness\n");
+ return std::nullopt;
+}
+
+// Check every structural precondition and record the induction analysis.
+bool LoopSplitUtils::isLegal() {
+ // Require a bottom-tested single-exit loop in LCSSA form with a preheader.
+ if (!L->getLoopPreheader() || !L->getLoopLatch() || !L->getExitingBlock() ||
+ !L->getExitBlock() || L->getExitingBlock() != L->getLoopLatch() ||
+ !L->isLCSSAForm(*DT)) {
+ LLVM_DEBUG(dbgs() << DEBUG_TYPE ": loop not in expected form\n");
+ return false;
+ }
+
+ // The latch compare must exist and reside in the latch.
+ ICmpInst *LatchCmp = L->getLatchCmpInst();
+ if (!LatchCmp || LatchCmp->getParent() != L->getLoopLatch()) {
+ LLVM_DEBUG(dbgs() << DEBUG_TYPE ": latch compare not in the loop latch\n");
+ return false;
+ }
+
+ // A computable backedge-taken count fixes the iteration space we rebuild.
+ const SCEV *BTC = SE->getBackedgeTakenCount(L);
+ if (isa<SCEVCouldNotCompute>(BTC)) {
+ LLVM_DEBUG(dbgs() << DEBUG_TYPE ": loop trip count uncomputable\n");
+ return false;
+ }
+
+ const SCEVAddRecExpr *IndAR = analyzeInduction(L, SE, LatchIndOperand);
+ if (!IndAR) {
+ LLVM_DEBUG(dbgs() << DEBUG_TYPE
+ ": no unique unit-step integer induction\n");
+ return false;
+ }
+
+ std::optional<bool> Signed = computeSignedness(L, IndAR);
+ if (!Signed)
+ return false;
+ InductionIsSigned = *Signed;
+
+ InductionEnd = IndAR->evaluateAtIteration(BTC, *SE);
+ // Start and end must share the induction type; reject any width mismatch.
+ if (InductionEnd->getType() != IndAR->getStart()->getType()) {
+ LLVM_DEBUG(dbgs() << DEBUG_TYPE ": induction end/start type mismatch\n");
+ return false;
+ }
+ return true;
+}
+
+//===----------------------------------------------------------------------===//
+// Transform
+//===----------------------------------------------------------------------===//
+
+// Latch "keep iterating" predicate (ascending </<=, descending >/>=); inclusive
+// when the latch compares the step value, strict when it compares the PHI.
+static ICmpInst::Predicate continuePredicate(bool Signed, bool Descending,
+ bool Inclusive) {
+ if (Descending)
+ return Inclusive ? (Signed ? ICmpInst::ICMP_SGE : ICmpInst::ICMP_UGE)
+ : (Signed ? ICmpInst::ICMP_SGT : ICmpInst::ICMP_UGT);
+ return Inclusive ? (Signed ? ICmpInst::ICMP_SLE : ICmpInst::ICMP_ULE)
+ : (Signed ? ICmpInst::ICMP_SLT : ICmpInst::ICMP_ULT);
+}
+
+// Guard "enter this partition" predicate: Start <= sel ascending, Start >= sel
+// descending.
+static ICmpInst::Predicate guardPredicate(bool Signed, bool Descending) {
+ if (Descending)
+ return Signed ? ICmpInst::ICMP_SGE : ICmpInst::ICMP_UGE;
+ return Signed ? ICmpInst::ICMP_SLE : ICmpInst::ICMP_ULE;
+}
+
+static void buildEntryGuard(BasicBlock *&Preheader, BasicBlock *&EntryGuard,
+ DominatorTree *DT, LoopInfo *LI);
+
+// Drive the whole transform: set up scratch state and run each phase in order.
+bool LoopSplitUtils::split() {
+ PHINode *Induction = L->getInductionVariable(*SE);
+ assert(Induction && "split() requires a successful isLegal()");
+ if (getNumPartitions() < 2)
+ return false;
+
+ if (!L->hasDedicatedExits() &&
+ !formDedicatedExitBlocks(L, DT, LI, /*MSSAU=*/nullptr,
+ /*PreserveLCSSA=*/true))
+ return false;
----------------
fhahn wrote:
untested?
https://github.com/llvm/llvm-project/pull/205995
More information about the llvm-commits
mailing list