[llvm] [SCEV] Batch common-factor folding in getAddExpr (NFC) (PR #184258)
Kevin McAfee via llvm-commits
llvm-commits at lists.llvm.org
Thu Apr 16 12:14:04 PDT 2026
================
@@ -2828,76 +2828,79 @@ const SCEV *ScalarEvolution::getAddExpr(SmallVectorImpl<SCEVUse> &Ops,
}
}
+ // Given a SCEVMulExpr and an operand index, return the product of all
+ // operands except the one at OpIdx.
+ auto StripFactor = [&](const SCEVMulExpr *M, unsigned OpIdx) -> SCEVUse {
+ if (M->getNumOperands() == 2)
+ return M->getOperand(OpIdx == 0);
+ SmallVector<SCEVUse, 4> Remaining(M->operands().take_front(OpIdx));
+ append_range(Remaining, M->operands().drop_front(OpIdx + 1));
+ return getMulExpr(Remaining, SCEV::FlagAnyWrap, Depth + 1);
+ };
+
// If we are adding something to a multiply expression, make sure the
// something is not already an operand of the multiply. If so, merge it into
// the multiply.
for (; Idx < Ops.size() && isa<SCEVMulExpr>(Ops[Idx]); ++Idx) {
const SCEVMulExpr *Mul = cast<SCEVMulExpr>(Ops[Idx]);
for (unsigned MulOp = 0, e = Mul->getNumOperands(); MulOp != e; ++MulOp) {
+ // Scan all terms to find every occurrence of common factor MulOpSCEV
+ // and fold them in one shot:
+ // A1*X + A2*X + ... + An*X --> X * (A1 + A2 + ... + An)
const SCEV *MulOpSCEV = Mul->getOperand(MulOp);
if (isa<SCEVConstant>(MulOpSCEV))
continue;
- for (unsigned AddOp = 0, e = Ops.size(); AddOp != e; ++AddOp)
+
+ // Cofactors: 1 for bare addends matching MulOpSCEV, or the
+ // remaining product for multiply terms containing MulOpSCEV.
+ SmallVector<SCEVUse, 4> Cofactors;
+ SmallVector<unsigned, 4> DeadIndices;
+ for (unsigned AddOp = 0, e = Ops.size(); AddOp != e; ++AddOp) {
if (MulOpSCEV == Ops[AddOp]) {
- // Fold W + X + (X * Y * Z) --> W + (X * ((Y*Z)+1))
- const SCEV *InnerMul = Mul->getOperand(MulOp == 0);
- if (Mul->getNumOperands() != 2) {
- // If the multiply has more than two operands, we must get the
- // Y*Z term.
- SmallVector<SCEVUse, 4> MulOps(Mul->operands().take_front(MulOp));
- append_range(MulOps, Mul->operands().drop_front(MulOp + 1));
- InnerMul = getMulExpr(MulOps, SCEV::FlagAnyWrap, Depth + 1);
- }
- const SCEV *AddOne =
- getAddExpr(getOne(Ty), InnerMul, SCEV::FlagAnyWrap, Depth + 1);
- const SCEV *OuterMul = getMulExpr(AddOne, MulOpSCEV,
- SCEV::FlagAnyWrap, Depth + 1);
- if (Ops.size() == 2) return OuterMul;
- if (AddOp < Idx) {
- Ops.erase(Ops.begin()+AddOp);
- Ops.erase(Ops.begin()+Idx-1);
- } else {
- Ops.erase(Ops.begin()+Idx);
- Ops.erase(Ops.begin()+AddOp-1);
- }
- Ops.push_back(OuterMul);
- return getAddExpr(Ops, SCEV::FlagAnyWrap, Depth + 1);
+ // W + X + (X * Y * Z) --> W + (X * ((Y*Z)+1))
+ Cofactors.push_back(getOne(Ty));
+ DeadIndices.push_back(AddOp);
+ continue;
}
- // Check this multiply against other multiplies being added together.
- for (unsigned OtherMulIdx = Idx+1;
- OtherMulIdx < Ops.size() && isa<SCEVMulExpr>(Ops[OtherMulIdx]);
- ++OtherMulIdx) {
- const SCEVMulExpr *OtherMul = cast<SCEVMulExpr>(Ops[OtherMulIdx]);
- // If MulOp occurs in OtherMul, we can fold the two multiplies
- // together.
- for (unsigned OMulOp = 0, e = OtherMul->getNumOperands();
- OMulOp != e; ++OMulOp)
+ if (AddOp == Idx || !isa<SCEVMulExpr>(Ops[AddOp]))
----------------
kalxr wrote:
I believe you're correct, nice catch
https://github.com/llvm/llvm-project/pull/184258
More information about the llvm-commits
mailing list