[llvm] 8a0c5d3 - [DAG] Replace llvm::isNeutralConstant with SelectionDAG::isIdentityElement (#195827)
via llvm-commits
llvm-commits at lists.llvm.org
Tue May 5 05:00:43 PDT 2026
Author: Simon Pilgrim
Date: 2026-05-05T12:00:37Z
New Revision: 8a0c5d3f43b27c8e2895c6106d6b551d626979fd
URL: https://github.com/llvm/llvm-project/commit/8a0c5d3f43b27c8e2895c6106d6b551d626979fd
DIFF: https://github.com/llvm/llvm-project/commit/8a0c5d3f43b27c8e2895c6106d6b551d626979fd.diff
LOG: [DAG] Replace llvm::isNeutralConstant with SelectionDAG::isIdentityElement (#195827)
Initial step towards generalising this - move to SelectionDAG like other
valuetracker helpers, add DemandedElts/Depth controls, etc.
We can add target node handling when it becomes necessary.
Renamed to "IdentityElement" to match llvm naming conventions.
Added:
Modified:
llvm/include/llvm/CodeGen/SelectionDAG.h
llvm/include/llvm/CodeGen/SelectionDAGNodes.h
llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp
llvm/lib/CodeGen/SelectionDAG/SelectionDAG.cpp
llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
llvm/lib/Target/RISCV/RISCVISelLowering.cpp
Removed:
################################################################################
diff --git a/llvm/include/llvm/CodeGen/SelectionDAG.h b/llvm/include/llvm/CodeGen/SelectionDAG.h
index b68d1fc7ea6ca..2512bce02dbeb 100644
--- a/llvm/include/llvm/CodeGen/SelectionDAG.h
+++ b/llvm/include/llvm/CodeGen/SelectionDAG.h
@@ -2277,6 +2277,19 @@ class SelectionDAG {
return computeOverflowForMul(IsSigned, N0, N1) == OFK_Never;
}
+ /// Returns true if \p V is an identity element of Opc with Flags.
+ /// When OperandNo is 0, it checks that V is a left identity. Otherwise, it
+ /// checks that V is a right identity.
+ LLVM_ABI bool isIdentityElement(unsigned Opc, SDNodeFlags Flags, SDValue V,
+ unsigned OperandNo, unsigned Depth = 0) const;
+
+ /// Returns true if the demanded vector elements of \p V is an identity
+ /// element of Opc with Flags. When OperandNo is 0, it checks that V is a left
+ /// identity. Otherwise, it checks that V is a right identity.
+ LLVM_ABI bool isIdentityElement(unsigned Opc, SDNodeFlags Flags, SDValue V,
+ const APInt &DemandedElts, unsigned OperandNo,
+ unsigned Depth = 0) const;
+
/// Test if the given value is known to have exactly one bit set. This
diff ers
/// from computeKnownBits in that it doesn't necessarily determine which bit
/// is set. If 'OrZero' is set, then return true if the given value is either
@@ -2726,9 +2739,9 @@ class SelectionDAG {
LLVM_ABI bool shouldOptForSize() const;
- /// Get the (commutative) neutral element for the given opcode, if it exists.
- LLVM_ABI SDValue getNeutralElement(unsigned Opcode, const SDLoc &DL, EVT VT,
- SDNodeFlags Flags);
+ /// Get the (commutative) identity element for the given opcode, if it exists.
+ LLVM_ABI SDValue getIdentityElement(unsigned Opcode, const SDLoc &DL, EVT VT,
+ SDNodeFlags Flags);
/// Get an expression that implements a partial multiply-subtract reduction.
/// In practice this means that parts of the expression are negated, e.g.
diff --git a/llvm/include/llvm/CodeGen/SelectionDAGNodes.h b/llvm/include/llvm/CodeGen/SelectionDAGNodes.h
index 7bc07f257de4d..aad16a16b6a8c 100644
--- a/llvm/include/llvm/CodeGen/SelectionDAGNodes.h
+++ b/llvm/include/llvm/CodeGen/SelectionDAGNodes.h
@@ -1949,12 +1949,6 @@ LLVM_ABI bool isOneConstant(SDValue V);
/// Returns true if \p V is a constant min signed integer value.
LLVM_ABI bool isMinSignedConstant(SDValue V);
-/// Returns true if \p V is a neutral element of Opc with Flags.
-/// When OperandNo is 0, it checks that V is a left identity. Otherwise, it
-/// checks that V is a right identity.
-LLVM_ABI bool isNeutralConstant(unsigned Opc, SDNodeFlags Flags, SDValue V,
- unsigned OperandNo);
-
/// Return the non-bitcasted source operand of \p V if it exists.
/// If \p V is not a bitcasted value, it is returned as-is.
LLVM_ABI SDValue peekThroughBitcasts(SDValue V);
diff --git a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
index 2e079ce0a4576..2debe9bacf40f 100644
--- a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
@@ -2538,7 +2538,7 @@ static SDValue foldSelectWithIdentityConstant(SDNode *N, SelectionDAG &DAG,
// This transform increases uses of N0, so freeze it to be safe.
// binop N0, (vselect Cond, IDC, FVal) --> vselect Cond, N0, (binop N0, FVal)
unsigned OpNo = ShouldCommuteOperands ? 0 : 1;
- if (isNeutralConstant(Opcode, N->getFlags(), TVal, OpNo) &&
+ if (DAG.isIdentityElement(Opcode, N->getFlags(), TVal, OpNo) &&
TLI.shouldFoldSelectWithIdentityConstant(Opcode, VT, SelOpcode, N0,
FVal)) {
SDValue F0 = DAG.getFreeze(N0);
@@ -2546,7 +2546,7 @@ static SDValue foldSelectWithIdentityConstant(SDNode *N, SelectionDAG &DAG,
return DAG.getSelect(SDLoc(N), VT, Cond, F0, NewBO);
}
// binop N0, (vselect Cond, TVal, IDC) --> vselect Cond, (binop N0, TVal), N0
- if (isNeutralConstant(Opcode, N->getFlags(), FVal, OpNo) &&
+ if (DAG.isIdentityElement(Opcode, N->getFlags(), FVal, OpNo) &&
TLI.shouldFoldSelectWithIdentityConstant(Opcode, VT, SelOpcode, N0,
TVal)) {
SDValue F0 = DAG.getFreeze(N0);
diff --git a/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp b/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp
index 856590ed2624f..cbd67675aab96 100644
--- a/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/LegalizeVectorTypes.cpp
@@ -8328,7 +8328,7 @@ SDValue DAGTypeLegalizer::WidenVecOp_VECREDUCE(SDNode *N) {
unsigned Opc = N->getOpcode();
unsigned BaseOpc = ISD::getVecReduceBaseOpcode(Opc);
- SDValue NeutralElem = DAG.getNeutralElement(BaseOpc, dl, ElemVT, Flags);
+ SDValue NeutralElem = DAG.getIdentityElement(BaseOpc, dl, ElemVT, Flags);
assert(NeutralElem && "Neutral element must exist");
// Pad the vector with the neutral element.
@@ -8382,7 +8382,7 @@ SDValue DAGTypeLegalizer::WidenVecOp_VECREDUCE_SEQ(SDNode *N) {
unsigned Opc = N->getOpcode();
unsigned BaseOpc = ISD::getVecReduceBaseOpcode(Opc);
- SDValue NeutralElem = DAG.getNeutralElement(BaseOpc, dl, ElemVT, Flags);
+ SDValue NeutralElem = DAG.getIdentityElement(BaseOpc, dl, ElemVT, Flags);
// Pad the vector with the neutral element.
unsigned OrigElts = OrigVT.getVectorMinNumElements();
diff --git a/llvm/lib/CodeGen/SelectionDAG/SelectionDAG.cpp b/llvm/lib/CodeGen/SelectionDAG/SelectionDAG.cpp
index 80d52e34bffde..857ff98f84b32 100644
--- a/llvm/lib/CodeGen/SelectionDAG/SelectionDAG.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/SelectionDAG.cpp
@@ -13544,11 +13544,19 @@ bool llvm::isMinSignedConstant(SDValue V) {
return Const != nullptr && Const->isMinSignedValue();
}
-bool llvm::isNeutralConstant(unsigned Opcode, SDNodeFlags Flags, SDValue V,
- unsigned OperandNo) {
+bool SelectionDAG::isIdentityElement(unsigned Opcode, SDNodeFlags Flags,
+ SDValue V, unsigned OperandNo,
+ unsigned Depth) const {
+ APInt DemandedElts = getDemandAllEltsMask(V);
+ return isIdentityElement(Opcode, Flags, V, DemandedElts, OperandNo, Depth);
+}
+
+bool SelectionDAG::isIdentityElement(unsigned Opcode, SDNodeFlags Flags,
+ SDValue V, const APInt &DemandedElts,
+ unsigned OperandNo, unsigned Depth) const {
// NOTE: The cases should match with IR's ConstantExpr::getBinOpIdentity().
// TODO: Target-specific opcodes could be added.
- if (auto *ConstV = isConstOrConstSplat(V, /*AllowUndefs*/ false,
+ if (auto *ConstV = isConstOrConstSplat(V, DemandedElts, /*AllowUndefs*/ false,
/*AllowTruncation*/ true)) {
APInt Const = ConstV->getAPIntValue().trunc(V.getScalarValueSizeInBits());
switch (Opcode) {
@@ -13575,7 +13583,7 @@ bool llvm::isNeutralConstant(unsigned Opcode, SDNodeFlags Flags, SDValue V,
case ISD::SDIV:
return OperandNo == 1 && Const.isOne();
}
- } else if (auto *ConstFP = isConstOrConstSplatFP(V)) {
+ } else if (auto *ConstFP = isConstOrConstSplatFP(V, DemandedElts)) {
switch (Opcode) {
case ISD::FADD:
return ConstFP->isZero() &&
@@ -13592,11 +13600,9 @@ bool llvm::isNeutralConstant(unsigned Opcode, SDNodeFlags Flags, SDValue V,
// Neutral element for fminnum is NaN, Inf or FLT_MAX, depending on FMF.
EVT VT = V.getValueType();
const fltSemantics &Semantics = VT.getFltSemantics();
- APFloat NeutralAF = !Flags.hasNoNaNs()
- ? APFloat::getQNaN(Semantics)
- : !Flags.hasNoInfs()
- ? APFloat::getInf(Semantics)
- : APFloat::getLargest(Semantics);
+ APFloat NeutralAF = !Flags.hasNoNaNs() ? APFloat::getQNaN(Semantics)
+ : !Flags.hasNoInfs() ? APFloat::getInf(Semantics)
+ : APFloat::getLargest(Semantics);
if (Opcode == ISD::FMAXNUM)
NeutralAF.changeSign();
@@ -14906,8 +14912,8 @@ SDValue SelectionDAG::getTokenFactor(const SDLoc &DL,
return getNode(ISD::TokenFactor, DL, MVT::Other, Vals);
}
-SDValue SelectionDAG::getNeutralElement(unsigned Opcode, const SDLoc &DL,
- EVT VT, SDNodeFlags Flags) {
+SDValue SelectionDAG::getIdentityElement(unsigned Opcode, const SDLoc &DL,
+ EVT VT, SDNodeFlags Flags) {
switch (Opcode) {
default:
return SDValue();
diff --git a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
index f5082b779d1db..356ebdf407cff 100644
--- a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
+++ b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
@@ -17368,7 +17368,7 @@ SDValue AArch64TargetLowering::LowerVECREDUCE_MUL(SDValue Op,
SDVTList SrcVTs = DAG.getVTList(SrcVT, SrcVT);
unsigned BaseOpc = ISD::getVecReduceBaseOpcode(Op.getOpcode());
- SDValue Identity = DAG.getNeutralElement(BaseOpc, DL, SrcVT, Op->getFlags());
+ SDValue Identity = DAG.getIdentityElement(BaseOpc, DL, SrcVT, Op->getFlags());
// Whilst we don't know the size of the vector we do know the maximum size so
// can perform a tree reduction with an identity vector, which means once we
diff --git a/llvm/lib/Target/RISCV/RISCVISelLowering.cpp b/llvm/lib/Target/RISCV/RISCVISelLowering.cpp
index 52249c3d258d8..d126616312748 100644
--- a/llvm/lib/Target/RISCV/RISCVISelLowering.cpp
+++ b/llvm/lib/Target/RISCV/RISCVISelLowering.cpp
@@ -12111,7 +12111,7 @@ SDValue RISCVTargetLowering::lowerVECREDUCE(SDValue Op,
SDValue StartV;
switch (BaseOpc) {
default:
- StartV = DAG.getNeutralElement(BaseOpc, DL, VecEltVT, SDNodeFlags());
+ StartV = DAG.getIdentityElement(BaseOpc, DL, VecEltVT, SDNodeFlags());
break;
case ISD::AND:
case ISD::OR:
@@ -15746,8 +15746,8 @@ static SDValue combineBinOpToReduce(SDNode *N, SelectionDAG &DAG,
// Check the scalar of ScalarV is neutral element
// TODO: Deal with value other than neutral element.
- if (!isNeutralConstant(N->getOpcode(), N->getFlags(), ScalarV.getOperand(1),
- 0))
+ if (!DAG.isIdentityElement(N->getOpcode(), N->getFlags(),
+ ScalarV.getOperand(1), 0))
return SDValue();
// If the AVL is zero, operand 0 will be returned. So it's not safe to fold.
@@ -19651,7 +19651,7 @@ static SDValue tryFoldSelectIntoOp(SDNode *N, SelectionDAG &DAG,
SDValue OtherOp = TrueVal.getOperand(1 - OpToFold);
EVT OtherOpVT = OtherOp.getValueType();
SDValue IdentityOperand =
- DAG.getNeutralElement(Opc, DL, OtherOpVT, N->getFlags());
+ DAG.getIdentityElement(Opc, DL, OtherOpVT, N->getFlags());
if (!Commutative)
IdentityOperand = DAG.getConstant(0, DL, OtherOpVT);
assert(IdentityOperand && "No identity operand!");
More information about the llvm-commits
mailing list