[llvm] [VPlan] Implement VPlan-based unit-strideness speculation (PR #182595)
Florian Hahn via llvm-commits
llvm-commits at lists.llvm.org
Mon Aug 24 04:44:49 PDT 2026
================
@@ -5520,6 +5529,187 @@ void VPlanTransforms::makeMemOpWideningDecisions(VPlan &Plan, VFRange &Range,
});
}
+void VPlanTransforms::multiversionForUnitStridedMemOps(
+ VPlan &Plan, VPCostContext &CostCtx, VFRange &Range,
+ ArrayRef<VPInstruction *> MemOps) {
+ ScalarEvolution *SE = CostCtx.PSE.getSE();
+ SCEVUnionPredicate StridePredicates({}, *SE);
+
+ for (VPInstruction *VPI : MemOps) {
+ bool IsLoad = VPI->getOpcode() == Instruction::Load;
+ VPValue *PtrOp = IsLoad ? VPI->getOperand(0) : VPI->getOperand(1);
+
+ const SCEV *PtrSCEV =
+ vputils::getSCEVExprForVPValue(PtrOp, CostCtx.PSE, CostCtx.L);
+ const SCEV *Start, *Stride;
+
+ if (!match(PtrSCEV, m_scev_AffineAddRec(m_SCEV(Start), m_SCEV(Stride),
+ m_SpecificLoop(CostCtx.L))))
+ continue;
+
+ Type *ScalarTy =
+ IsLoad ? VPI->getScalarType() : VPI->getOperand(0)->getScalarType();
+
+ if (VPI->getMask()) {
+ Instruction *I = VPI->getUnderlyingInstr();
+ // Don't speculate unit-strideness if it won't result in any unit-strided
+ // loads, as we'd pay the price of not taking vector loop if the runtime
+ // condition is false for no benefits.
+ if (!CostCtx.Config.isLegalMaskedLoadOrStore(IsLoad, ScalarTy,
+ getLoadStoreAlignment(I),
+ getLoadStoreAddressSpace(I)))
+ continue;
+ }
+
+ if (isa<SCEVConstant>(Stride))
+ continue;
+
+ const auto *TypeSize = cast<SCEVConstant>(SE->getSizeOfExpr(
+ Stride->getType(), SE->getDataLayout().getTypeAllocSize(ScalarTy)));
+
+ const SCEVConstant *StrideConstantMultiplier;
+ const SCEV *StrideNonConstantMultiplier;
+
+ const SCEV *ToMultiVersion = Stride;
+ const SCEV *MVConst = TypeSize;
+ if (match(Stride, m_scev_c_Mul(m_SCEVConstant(StrideConstantMultiplier),
+ m_SCEV(StrideNonConstantMultiplier)))) {
+ if (TypeSize != StrideConstantMultiplier) {
+ // TODO: Support `TypeSize = N * StrideConstantMultiplier`,
+ // including negative `N`. For now, only process when they're equal,
+ // which matches the useful part of the legacy behavior that
+ // multiversiones GEP index for stride one.
+ continue;
+ }
+ ToMultiVersion = StrideNonConstantMultiplier;
+ MVConst = SE->getOne(ToMultiVersion->getType());
+ } else if (!TypeSize->isOne()) {
+ // Likewise - try to match legacy behavior.
+ continue;
+ }
+
+ while (auto *C = dyn_cast<SCEVIntegralCastExpr>(ToMultiVersion)) {
+ ToMultiVersion = C->getOperand();
+ MVConst = SE->getTruncateOrSignExtend(MVConst, ToMultiVersion->getType());
+ }
+
+ if (match(ToMultiVersion, m_scev_UndefOrPoison()))
+ continue;
+
+ if (!isa<SCEVUnknown>(ToMultiVersion)) {
+ // Match legacy behavior.
+ // If/when changed, make sure that explicit poison/undef in the defining
+ // expression doesn't cause any issues.
+ continue;
+ }
+
+ Value *StrideVal = cast<SCEVUnknown>(ToMultiVersion)->getValue();
+
+ const SCEVPredicate *NewPred =
+ SE->getComparePredicate(CmpInst::ICMP_EQ, ToMultiVersion, MVConst);
+
+ auto *PredicatedMaxBTC = SE->rewriteUsingPredicate(
+ CostCtx.PSE.getSymbolicMaxBackedgeTakenCount(), CostCtx.L,
+ StridePredicates.getUnionWith(NewPred, *SE)
+ .getUnionWith(&CostCtx.PSE.getPredicate(), *SE));
+ Type *BTCTy = PredicatedMaxBTC->getType();
+
+ // If predicate implies scalar loop never takes the backedge, don't perform
+ // multiversioning.
+ if (SE->isKnownPredicate(ICmpInst::ICMP_ULT, PredicatedMaxBTC,
+ SE->getOne(BTCTy)))
+ continue;
+
+ // If we don't fold the tail, we need enough scalar iterations to fill the
+ // full vector.
+ if (!Plan.hasTailFolded() &&
+ LoopVectorizationPlanner::getDecisionAndClampRange(
+ [&](ElementCount VF) {
+ return SE->isKnownPredicate(
+ ICmpInst::ICMP_ULT, PredicatedMaxBTC,
+ SE->getAddExpr(SE->getElementCount(BTCTy, VF),
+ SE->getMinusOne(BTCTy)));
+ },
+ Range))
+ continue;
+
+ StridePredicates = StridePredicates.getUnionWith(NewPred, *SE);
+
+ auto ReplaceUsesInVectorLoop = [&](Value *V, const SCEV *ToSCEV) {
+ VPValue *From = Plan.getLiveIn(V);
+ if (!From)
+ return;
+
+ assert(From->getScalarType() == ToSCEV->getType() &&
+ "Wrong type for ToSCEV!");
+ VPValue *To = Plan.getConstantInt(cast<SCEVConstant>(ToSCEV)->getAPInt());
+
+ // Original scalar loop can still use `From`, make sure to only rewrite
+ // uses inside the vector loop that we guard with the checks.
+ From->replaceUsesWithIf(To, [&](VPUser &U, unsigned) {
+ auto *R = cast<VPRecipeBase>(&U);
+ return R->getRegion() || R->getParent() == Plan.getVectorPreheader();
+ });
+ };
+
+ ReplaceUsesInVectorLoop(StrideVal, MVConst);
+ // If `StrideVal` has casts defined outside VPlan that are live-ins, replace
+ // them too.
+ for (auto *U : StrideVal->users())
+ if (isa<SExtInst>(U))
+ ReplaceUsesInVectorLoop(U,
+ SE->getSignExtendExpr(MVConst, U->getType()));
+ else if (isa<ZExtInst, TruncInst>(U))
+ ReplaceUsesInVectorLoop(
+ U, SE->getTruncateOrZeroExtend(MVConst, U->getType()));
+ }
+
+ if (StridePredicates.isAlwaysTrue())
+ return;
+
+ VPBasicBlock *StridesCheckVPBB = Plan.createVPBasicBlock("strides.check");
+ // We will replace the condition once we expand the predicate.
+ attachVPCheckBlock(Plan, Plan.getTrue(), StridesCheckVPBB,
----------------
fhahn wrote:
Do we have a test case where the stride is provably != 1? I think this may currently crash
Something like
```
define void @range_attr(ptr noalias %out, ptr %p, i64 range(i64 4, 8) %stride) {
entry:
br label %loop
loop:
%i = phi i64 [ 0, %entry ], [ %i.next, %loop ]
%idx = mul i64 %i, %stride
%gep = getelementptr inbounds i32, ptr %p, i64 %idx
%l = load i32, ptr %gep, align 4
%o = getelementptr inbounds i32, ptr %out, i64 %i
store i32 %l, ptr %o, align 4
%i.next = add i64 %i, 1
%ec = icmp eq i64 %i.next, 1024
br i1 %ec, label %exit, label %loop
exit:
ret void
}
```
https://github.com/llvm/llvm-project/pull/182595
More information about the llvm-commits
mailing list