[llvm] [DAG] ISD::matchUnaryPredicate / matchUnaryFpPredicate / matchBinaryPredicate - add DemandedElts variant (PR #183013)
via llvm-commits
llvm-commits at lists.llvm.org
Sun Aug 2 02:02:59 PDT 2026
https://github.com/VachanVY updated https://github.com/llvm/llvm-project/pull/183013
>From 14018a0012a78e707d000304a92f450c6db20930 Mon Sep 17 00:00:00 2001
From: Vachan V Y <vachanvy05 at gmail.com>
Date: Sun, 2 Aug 2026 14:32:38 +0530
Subject: [PATCH] [DAG] ISD::matchUnaryPredicate / matchUnaryFpPredicate /
matchBinaryPredicate - add DemandedElts variant
---
llvm/include/llvm/CodeGen/SelectionDAGNodes.h | 52 +++++++++++++++--
.../lib/CodeGen/SelectionDAG/SelectionDAG.cpp | 56 ++++++-------------
2 files changed, 64 insertions(+), 44 deletions(-)
diff --git a/llvm/include/llvm/CodeGen/SelectionDAGNodes.h b/llvm/include/llvm/CodeGen/SelectionDAGNodes.h
index 5654817a92930..6a44bd65a045a 100644
--- a/llvm/include/llvm/CodeGen/SelectionDAGNodes.h
+++ b/llvm/include/llvm/CodeGen/SelectionDAGNodes.h
@@ -3485,40 +3485,80 @@ namespace ISD {
}
/// Attempt to match a unary predicate against a scalar/splat constant or
- /// every element of a constant BUILD_VECTOR.
+ /// every element of a constant BUILD_VECTOR. The DemandedElts argument
+ /// allows us to only collect the known bits that are shared by the requested
+ /// vector elements.
/// If AllowUndef is true, then UNDEF elements will pass nullptr to Match.
template <typename ConstNodeType>
- bool matchUnaryPredicateImpl(SDValue Op,
+ bool matchUnaryPredicateImpl(SDValue Op, const APInt &DemandedElts,
std::function<bool(ConstNodeType *)> Match,
bool AllowUndefs = false,
bool AllowTruncation = false);
/// Hook for matching ConstantSDNode predicate
+ inline bool matchUnaryPredicate(SDValue Op, const APInt &DemandedElts,
+ std::function<bool(ConstantSDNode *)> Match,
+ bool AllowUndefs = false,
+ bool AllowTruncation = false) {
+ return matchUnaryPredicateImpl<ConstantSDNode>(
+ Op, DemandedElts, Match, AllowUndefs, AllowTruncation);
+ }
+
inline bool matchUnaryPredicate(SDValue Op,
std::function<bool(ConstantSDNode *)> Match,
bool AllowUndefs = false,
bool AllowTruncation = false) {
- return matchUnaryPredicateImpl<ConstantSDNode>(Op, Match, AllowUndefs,
- AllowTruncation);
+ EVT VT = Op.getValueType();
+ APInt DemandedElts = VT.isFixedLengthVector()
+ ? APInt::getAllOnes(VT.getVectorNumElements())
+ : APInt(1, 1);
+ return matchUnaryPredicate(Op, DemandedElts, Match, AllowUndefs,
+ AllowTruncation);
}
/// Hook for matching ConstantFPSDNode predicate
+ inline bool
+ matchUnaryFpPredicate(SDValue Op, const APInt &DemandedElts,
+ std::function<bool(ConstantFPSDNode *)> Match,
+ bool AllowUndefs = false) {
+ return matchUnaryPredicateImpl<ConstantFPSDNode>(Op, DemandedElts, Match,
+ AllowUndefs);
+ }
+
inline bool
matchUnaryFpPredicate(SDValue Op,
std::function<bool(ConstantFPSDNode *)> Match,
bool AllowUndefs = false) {
- return matchUnaryPredicateImpl<ConstantFPSDNode>(Op, Match, AllowUndefs);
+ EVT VT = Op.getValueType();
+ APInt DemandedElts = VT.isFixedLengthVector()
+ ? APInt::getAllOnes(VT.getVectorNumElements())
+ : APInt(1, 1);
+ return matchUnaryFpPredicate(Op, DemandedElts, Match, AllowUndefs);
}
/// Attempt to match a binary predicate against a pair of scalar/splat
/// constants or every element of a pair of constant BUILD_VECTORs.
+ /// The DemandedElts argument allows us to only collect the
+ /// known bits that are shared by the requested vector elements.
/// If AllowUndef is true, then UNDEF elements will pass nullptr to Match.
/// If AllowTypeMismatch is true then RetType + ArgTypes don't need to match.
LLVM_ABI bool matchBinaryPredicate(
- SDValue LHS, SDValue RHS,
+ SDValue LHS, SDValue RHS, const APInt &DemandedElts,
std::function<bool(ConstantSDNode *, ConstantSDNode *)> Match,
bool AllowUndefs = false, bool AllowTypeMismatch = false);
+ inline bool matchBinaryPredicate(
+ SDValue LHS, SDValue RHS,
+ std::function<bool(ConstantSDNode *, ConstantSDNode *)> Match,
+ bool AllowUndefs = false, bool AllowTypeMismatch = false) {
+ EVT VT = LHS.getValueType();
+ APInt DemandedElts = VT.isFixedLengthVector()
+ ? APInt::getAllOnes(VT.getVectorNumElements())
+ : APInt(1, 1);
+ return matchBinaryPredicate(LHS, RHS, DemandedElts, Match, AllowUndefs,
+ AllowTypeMismatch);
+ }
+
/// Returns true if the specified value is the overflow result from one
/// of the overflow intrinsic nodes.
inline bool isOverflowIntrOpRes(SDValue Op) {
diff --git a/llvm/lib/CodeGen/SelectionDAG/SelectionDAG.cpp b/llvm/lib/CodeGen/SelectionDAG/SelectionDAG.cpp
index 5c80c4c1b5bff..55aec5c0c8cf9 100644
--- a/llvm/lib/CodeGen/SelectionDAG/SelectionDAG.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/SelectionDAG.cpp
@@ -349,7 +349,7 @@ bool ISD::isFreezeUndef(const SDNode *N) {
}
template <typename ConstNodeType>
-bool ISD::matchUnaryPredicateImpl(SDValue Op,
+bool ISD::matchUnaryPredicateImpl(SDValue Op, const APInt &DemandedElts,
std::function<bool(ConstNodeType *)> Match,
bool AllowUndefs, bool AllowTruncation) {
// FIXME: Add support for scalar UNDEF cases?
@@ -361,8 +361,14 @@ bool ISD::matchUnaryPredicateImpl(SDValue Op,
ISD::SPLAT_VECTOR != Op.getOpcode())
return false;
+ if (ISD::SPLAT_VECTOR == Op.getOpcode() && !DemandedElts)
+ return true;
+
EVT SVT = Op.getValueType().getScalarType();
for (unsigned i = 0, e = Op.getNumOperands(); i != e; ++i) {
+ if (ISD::SPLAT_VECTOR != Op.getOpcode() && !DemandedElts[i])
+ continue;
+
if (AllowUndefs && Op.getOperand(i).isUndef()) {
if (!Match(nullptr))
return false;
@@ -378,12 +384,13 @@ bool ISD::matchUnaryPredicateImpl(SDValue Op,
}
// Build used template types.
template bool ISD::matchUnaryPredicateImpl<ConstantSDNode>(
- SDValue, std::function<bool(ConstantSDNode *)>, bool, bool);
+ SDValue, const APInt &, std::function<bool(ConstantSDNode *)>, bool, bool);
template bool ISD::matchUnaryPredicateImpl<ConstantFPSDNode>(
- SDValue, std::function<bool(ConstantFPSDNode *)>, bool, bool);
+ SDValue, const APInt &, std::function<bool(ConstantFPSDNode *)>, bool,
+ bool);
bool ISD::matchBinaryPredicate(
- SDValue LHS, SDValue RHS,
+ SDValue LHS, SDValue RHS, const APInt &DemandedElts,
std::function<bool(ConstantSDNode *, ConstantSDNode *)> Match,
bool AllowUndefs, bool AllowTypeMismatch) {
if (!AllowTypeMismatch && LHS.getValueType() != RHS.getValueType())
@@ -400,8 +407,13 @@ bool ISD::matchBinaryPredicate(
LHS.getOpcode() != ISD::SPLAT_VECTOR))
return false;
+ if (ISD::SPLAT_VECTOR == LHS.getOpcode() && !DemandedElts)
+ return true;
+
EVT SVT = LHS.getValueType().getScalarType();
for (unsigned i = 0, e = LHS.getNumOperands(); i != e; ++i) {
+ if (ISD::SPLAT_VECTOR != LHS.getOpcode() && !DemandedElts[i])
+ continue;
SDValue LHSOp = LHS.getOperand(i);
SDValue RHSOp = RHS.getOperand(i);
bool LHSUndef = AllowUndefs && LHSOp.isUndef();
@@ -4743,26 +4755,10 @@ bool SelectionDAG::isKnownToBeAPowerOfTwo(SDValue Val,
};
// Is the constant a known power of 2 or zero?
- if (ISD::matchUnaryPredicate(Val, IsPowerOfTwoOrZero))
+ if (ISD::matchUnaryPredicate(Val, DemandedElts, IsPowerOfTwoOrZero))
return true;
switch (Val.getOpcode()) {
- case ISD::BUILD_VECTOR:
- // Are all operands of a build vector constant powers of two or zero?
- if (all_of(enumerate(Val->ops()), [&](auto P) {
- auto *C = dyn_cast<ConstantSDNode>(P.value());
- return !DemandedElts[P.index()] || (C && IsPowerOfTwoOrZero(C));
- }))
- return true;
- break;
-
- case ISD::SPLAT_VECTOR:
- // Is the operand of a splat vector a constant power of two?
- if (auto *C = dyn_cast<ConstantSDNode>(Val->getOperand(0)))
- if (IsPowerOfTwoOrZero(C))
- return true;
- break;
-
case ISD::EXTRACT_VECTOR_ELT: {
SDValue InVec = Val.getOperand(0);
SDValue EltNo = Val.getOperand(1);
@@ -6489,7 +6485,7 @@ bool SelectionDAG::isKnownNeverZero(SDValue Op, const APInt &DemandedElts,
return !V.isZero();
};
- if (ISD::matchUnaryPredicate(Op, IsNeverZero))
+ if (ISD::matchUnaryPredicate(Op, DemandedElts, IsNeverZero))
return true;
// TODO: Recognize more cases here. Most of the cases are also incomplete to
@@ -6498,22 +6494,6 @@ bool SelectionDAG::isKnownNeverZero(SDValue Op, const APInt &DemandedElts,
default:
break;
- case ISD::BUILD_VECTOR:
- // Are all operands of a build vector constant non-zero?
- if (all_of(enumerate(Op->ops()), [&](auto P) {
- auto *C = dyn_cast<ConstantSDNode>(P.value());
- return !DemandedElts[P.index()] || (C && IsNeverZero(C));
- }))
- return true;
- break;
-
- case ISD::SPLAT_VECTOR:
- // Is the operand of a splat vector a constant non-zero?
- if (auto *C = dyn_cast<ConstantSDNode>(Op->getOperand(0)))
- if (IsNeverZero(C))
- return true;
- break;
-
case ISD::EXTRACT_VECTOR_ELT: {
SDValue InVec = Op.getOperand(0);
SDValue EltNo = Op.getOperand(1);
More information about the llvm-commits
mailing list