[llvm] [SDPatternMatch] Make m_SetCC work like m_ICmp from IR PatternMatch. (PR #226623)
Craig Topper via llvm-commits
llvm-commits at lists.llvm.org
Fri Sep 25 19:07:48 PDT 2026
https://github.com/topperc created https://github.com/llvm/llvm-project/pull/226623
The condition code is stored an operand, but we don't need to expose that to the interface.
This adds 2 signatures of m_Setcc, one that takes 2 operands and matches any condition code and one that takes the matched condition code by reference. For m_SpecificCondCode cases, I've added m_SpecificSetCC.
Similar changes have been applied to m_SelectCC and m_SelectCCLike.
Out of tree targets will need to update to the new interface. Alternatively, we could keep the current 3 operand form.
Assisted-by: Claude
>From 3fdd2aaaa81a8327bd46524850b3a78eb4f3492c Mon Sep 17 00:00:00 2001
From: Craig Topper <craig.topper at sifive.com>
Date: Fri, 25 Sep 2026 18:45:40 -0700
Subject: [PATCH] [SDPatternMatch] Make m_SetCC work like m_ICmp from IR
PatternMatch.
The condition code is stored an operand, but we don't need to expose
that to the interface.
This adds 2 signatures of m_Setcc, one that takes 2 operands and
matches any condition code and one that takes the matched condition
code by reference. For m_SpecificCondCode cases, I've added
m_SpecificSetCC.
Similar changes have been applied to m_SelectCC and m_SelectCCLike.
Out of tree targets will need to update to the new interface.
Alternatively, we could keep the current 3 operand form.
Assisted-by: Claude
---
llvm/include/llvm/CodeGen/SDPatternMatch.h | 168 +++++++++++++-----
llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp | 46 +++--
.../Target/AArch64/AArch64ISelLowering.cpp | 6 +-
llvm/lib/Target/NVPTX/NVPTXISelLowering.cpp | 17 +-
llvm/lib/Target/RISCV/RISCVISelLowering.cpp | 15 +-
.../WebAssembly/WebAssemblyISelLowering.cpp | 15 +-
llvm/lib/Target/X86/X86ISelLowering.cpp | 42 +++--
.../CodeGen/SelectionDAGPatternMatchTest.cpp | 97 +++++++---
8 files changed, 259 insertions(+), 147 deletions(-)
diff --git a/llvm/include/llvm/CodeGen/SDPatternMatch.h b/llvm/include/llvm/CodeGen/SDPatternMatch.h
index 2e82693708fca4..e653c3cc508d30 100644
--- a/llvm/include/llvm/CodeGen/SDPatternMatch.h
+++ b/llvm/include/llvm/CodeGen/SDPatternMatch.h
@@ -455,17 +455,88 @@ struct TernaryOpc_match {
}
};
-template <typename T0_P, typename T1_P, typename T2_P>
-inline TernaryOpc_match<T0_P, T1_P, T2_P>
-m_SetCC(const T0_P &LHS, const T1_P &RHS, const T2_P &CC) {
- return TernaryOpc_match<T0_P, T1_P, T2_P>(ISD::SETCC, LHS, RHS, CC);
+struct CondCode_match {
+ std::optional<ISD::CondCode> CCToMatch;
+ ISD::CondCode *BindCC = nullptr;
+
+ explicit CondCode_match(ISD::CondCode CC) : CCToMatch(CC) {}
+
+ explicit CondCode_match(ISD::CondCode *CC) : BindCC(CC) {}
+
+ bool match(SDValue N) {
+ if (auto *CC = dyn_cast<CondCodeSDNode>(N.getNode())) {
+ if (CCToMatch && *CCToMatch != CC->get())
+ return false;
+
+ if (BindCC)
+ *BindCC = CC->get();
+ return true;
+ }
+
+ return false;
+ }
+};
+
+/// Match any conditional code SDNode.
+inline CondCode_match m_CondCode() { return CondCode_match(nullptr); }
+/// Match any conditional code SDNode and return its ISD::CondCode value.
+inline CondCode_match m_CondCode(ISD::CondCode &CC) {
+ return CondCode_match(&CC);
+}
+/// Match a conditional code SDNode with a specific ISD::CondCode.
+inline CondCode_match m_SpecificCondCode(ISD::CondCode CC) {
+ return CondCode_match(CC);
}
-template <typename T0_P, typename T1_P, typename T2_P>
-inline TernaryOpc_match<T0_P, T1_P, T2_P, true, false>
-m_c_SetCC(const T0_P &LHS, const T1_P &RHS, const T2_P &CC) {
- return TernaryOpc_match<T0_P, T1_P, T2_P, true, false>(ISD::SETCC, LHS, RHS,
- CC);
+/// Match a SETCC with any condition code.
+template <typename T0_P, typename T1_P>
+inline TernaryOpc_match<T0_P, T1_P, CondCode_match> m_SetCC(const T0_P &LHS,
+ const T1_P &RHS) {
+ return TernaryOpc_match<T0_P, T1_P, CondCode_match>(ISD::SETCC, LHS, RHS,
+ m_CondCode());
+}
+
+/// Match a SETCC with any condition code and bind the condition code to CC.
+template <typename T0_P, typename T1_P>
+inline TernaryOpc_match<T0_P, T1_P, CondCode_match>
+m_SetCC(ISD::CondCode &CC, const T0_P &LHS, const T1_P &RHS) {
+ return TernaryOpc_match<T0_P, T1_P, CondCode_match>(ISD::SETCC, LHS, RHS,
+ m_CondCode(CC));
+}
+
+/// Match a SETCC with a specific condition code.
+template <typename T0_P, typename T1_P>
+inline TernaryOpc_match<T0_P, T1_P, CondCode_match>
+m_SpecificSetCC(ISD::CondCode CC, const T0_P &LHS, const T1_P &RHS) {
+ return TernaryOpc_match<T0_P, T1_P, CondCode_match>(ISD::SETCC, LHS, RHS,
+ m_SpecificCondCode(CC));
+}
+
+/// Match a SETCC with any condition code, allowing the operands to be
+/// commuted.
+template <typename T0_P, typename T1_P>
+inline TernaryOpc_match<T0_P, T1_P, CondCode_match, true, false>
+m_c_SetCC(const T0_P &LHS, const T1_P &RHS) {
+ return TernaryOpc_match<T0_P, T1_P, CondCode_match, true, false>(
+ ISD::SETCC, LHS, RHS, m_CondCode());
+}
+
+/// Match a SETCC with any condition code, allowing the operands to be
+/// commuted, and bind the condition code to CC.
+template <typename T0_P, typename T1_P>
+inline TernaryOpc_match<T0_P, T1_P, CondCode_match, true, false>
+m_c_SetCC(ISD::CondCode &CC, const T0_P &LHS, const T1_P &RHS) {
+ return TernaryOpc_match<T0_P, T1_P, CondCode_match, true, false>(
+ ISD::SETCC, LHS, RHS, m_CondCode(CC));
+}
+
+/// Match a SETCC with a specific condition code, allowing the operands to be
+/// commuted.
+template <typename T0_P, typename T1_P>
+inline TernaryOpc_match<T0_P, T1_P, CondCode_match, true, false>
+m_c_SpecificSetCC(ISD::CondCode CC, const T0_P &LHS, const T1_P &RHS) {
+ return TernaryOpc_match<T0_P, T1_P, CondCode_match, true, false>(
+ ISD::SETCC, LHS, RHS, m_SpecificCondCode(CC));
}
template <typename T0_P, typename T1_P, typename T2_P>
@@ -524,16 +595,48 @@ m_c_TernaryOp(unsigned Opc, const T0_P &Op0, const T1_P &Op1, const T2_P &Op2) {
return TernaryOpc_match<T0_P, T1_P, T2_P, true>(Opc, Op0, Op1, Op2);
}
-template <typename LTy, typename RTy, typename TTy, typename FTy, typename CCTy>
-inline auto m_SelectCC(const LTy &L, const RTy &R, const TTy &T, const FTy &F,
- const CCTy &CC) {
- return m_Node(ISD::SELECT_CC, L, R, T, F, CC);
+/// Match a SELECT_CC with any condition code.
+template <typename LTy, typename RTy, typename TTy, typename FTy>
+inline auto m_SelectCC(const LTy &L, const RTy &R, const TTy &T, const FTy &F) {
+ return m_Node(ISD::SELECT_CC, L, R, T, F, m_CondCode());
+}
+
+/// Match a SELECT_CC with any condition code and bind the condition code to
+/// CC.
+template <typename LTy, typename RTy, typename TTy, typename FTy>
+inline auto m_SelectCC(ISD::CondCode &CC, const LTy &L, const RTy &R,
+ const TTy &T, const FTy &F) {
+ return m_Node(ISD::SELECT_CC, L, R, T, F, m_CondCode(CC));
}
-template <typename LTy, typename RTy, typename TTy, typename FTy, typename CCTy>
+/// Match a SELECT_CC with a specific condition code.
+template <typename LTy, typename RTy, typename TTy, typename FTy>
+inline auto m_SpecificSelectCC(ISD::CondCode CC, const LTy &L, const RTy &R,
+ const TTy &T, const FTy &F) {
+ return m_Node(ISD::SELECT_CC, L, R, T, F, m_SpecificCondCode(CC));
+}
+
+/// Match a SELECT of a SETCC or a SELECT_CC with any condition code.
+template <typename LTy, typename RTy, typename TTy, typename FTy>
inline auto m_SelectCCLike(const LTy &L, const RTy &R, const TTy &T,
- const FTy &F, const CCTy &CC) {
- return m_AnyOf(m_Select(m_SetCC(L, R, CC), T, F), m_SelectCC(L, R, T, F, CC));
+ const FTy &F) {
+ return m_AnyOf(m_Select(m_SetCC(L, R), T, F), m_SelectCC(L, R, T, F));
+}
+
+/// Match a SELECT of a SETCC or a SELECT_CC with any condition code and bind
+/// the condition code to CC.
+template <typename LTy, typename RTy, typename TTy, typename FTy>
+inline auto m_SelectCCLike(ISD::CondCode &CC, const LTy &L, const RTy &R,
+ const TTy &T, const FTy &F) {
+ return m_AnyOf(m_Select(m_SetCC(CC, L, R), T, F), m_SelectCC(CC, L, R, T, F));
+}
+
+/// Match a SELECT of a SETCC or a SELECT_CC with a specific condition code.
+template <typename LTy, typename RTy, typename TTy, typename FTy>
+inline auto m_SpecificSelectCCLike(ISD::CondCode CC, const LTy &L, const RTy &R,
+ const TTy &T, const FTy &F) {
+ return m_AnyOf(m_Select(m_SpecificSetCC(CC, L, R), T, F),
+ m_SpecificSelectCC(CC, L, R, T, F));
}
// === Binary operations ===
@@ -1318,39 +1421,6 @@ inline auto m_True(const SelectionDAG &DAG) { return Bool_match<true>(DAG); }
/// TargetLowering.
inline auto m_False(const SelectionDAG &DAG) { return Bool_match<false>(DAG); }
-struct CondCode_match {
- std::optional<ISD::CondCode> CCToMatch;
- ISD::CondCode *BindCC = nullptr;
-
- explicit CondCode_match(ISD::CondCode CC) : CCToMatch(CC) {}
-
- explicit CondCode_match(ISD::CondCode *CC) : BindCC(CC) {}
-
- bool match(SDValue N) {
- if (auto *CC = dyn_cast<CondCodeSDNode>(N.getNode())) {
- if (CCToMatch && *CCToMatch != CC->get())
- return false;
-
- if (BindCC)
- *BindCC = CC->get();
- return true;
- }
-
- return false;
- }
-};
-
-/// Match any conditional code SDNode.
-inline CondCode_match m_CondCode() { return CondCode_match(nullptr); }
-/// Match any conditional code SDNode and return its ISD::CondCode value.
-inline CondCode_match m_CondCode(ISD::CondCode &CC) {
- return CondCode_match(&CC);
-}
-/// Match a conditional code SDNode with a specific ISD::CondCode.
-inline CondCode_match m_SpecificCondCode(ISD::CondCode CC) {
- return CondCode_match(CC);
-}
-
/// Match a negate as a sub(0, v)
template <typename ValTy>
inline BinaryOpc_match<Zero_match, ValTy, false> m_Neg(const ValTy &V) {
diff --git a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
index 4f201982448837..53498de41c7288 100644
--- a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
@@ -2482,8 +2482,7 @@ static bool isTruncateOf(SelectionDAG &DAG, SDValue N, SDValue &Op,
}
if (N.getValueType().getScalarType() != MVT::i1 ||
- !sd_match(
- N, m_c_SetCC(m_Value(Op), m_Zero(), m_SpecificCondCode(ISD::SETNE))))
+ !sd_match(N, m_c_SpecificSetCC(ISD::SETNE, m_Value(Op), m_Zero())))
return false;
Known = DAG.computeKnownBits(Op);
@@ -2718,8 +2717,9 @@ static SDValue foldAddSubBoolOfMaskedVal(SDNode *N, const SDLoc &DL,
return SDValue();
// Match the compare as: setcc (X & 1), 0, eq.
- if (!sd_match(Z.getOperand(0), m_SetCC(m_And(m_Value(), m_One()), m_Zero(),
- m_SpecificCondCode(ISD::SETEQ))))
+ if (!sd_match(
+ Z.getOperand(0),
+ m_SpecificSetCC(ISD::SETEQ, m_And(m_Value(), m_One()), m_Zero())))
return SDValue();
// We are adding/subtracting a constant and an inverted low bit. Turn that
@@ -3957,10 +3957,9 @@ static SDValue combineCarryDiamond(SelectionDAG &DAG, const TargetLowering &TLI,
static SDValue combineOrOfSetCCToUSUBOCarry(SDNode *N, SelectionDAG &DAG,
const TargetLowering &TLI) {
SDValue A, B, CarryIn;
- if (!sd_match(N, m_Or(m_SetCC(m_Value(A), m_Value(B),
- m_SpecificCondCode(ISD::SETULT)),
- m_And(m_c_SetCC(m_Deferred(A), m_Deferred(B),
- m_SpecificCondCode(ISD::SETEQ)),
+ if (!sd_match(N, m_Or(m_SpecificSetCC(ISD::SETULT, m_Value(A), m_Value(B)),
+ m_And(m_c_SpecificSetCC(ISD::SETEQ, m_Deferred(A),
+ m_Deferred(B)),
m_Value(CarryIn)))))
return SDValue();
@@ -4311,13 +4310,15 @@ SDValue DAGCombiner::visitSUB(SDNode *N) {
auto MS0 = m_Specific(N0);
auto MVY = m_Value(Y);
auto MZ = m_Zero();
- auto MCC1 = m_SpecificCondCode(ISD::SETULT);
- auto MCC2 = m_SpecificCondCode(ISD::SETUGE);
- if (sd_match(N1, m_SelectCCLike(MS0, MVY, MZ, m_Deferred(Y), MCC1)) ||
- sd_match(N1, m_SelectCCLike(MS0, MVY, m_Deferred(Y), MZ, MCC2)) ||
- sd_match(N1, m_VSelect(m_SetCC(MS0, MVY, MCC1), MZ, m_Deferred(Y))) ||
- sd_match(N1, m_VSelect(m_SetCC(MS0, MVY, MCC2), m_Deferred(Y), MZ)))
+ if (sd_match(N1, m_SpecificSelectCCLike(ISD::SETULT, MS0, MVY, MZ,
+ m_Deferred(Y))) ||
+ sd_match(N1, m_SpecificSelectCCLike(ISD::SETUGE, MS0, MVY,
+ m_Deferred(Y), MZ)) ||
+ sd_match(N1, m_VSelect(m_SpecificSetCC(ISD::SETULT, MS0, MVY), MZ,
+ m_Deferred(Y))) ||
+ sd_match(N1, m_VSelect(m_SpecificSetCC(ISD::SETUGE, MS0, MVY),
+ m_Deferred(Y), MZ)))
return DAG.getNode(ISD::UMIN, DL, VT, N0,
DAG.getNode(ISD::SUB, DL, VT, N0, Y));
@@ -6392,14 +6393,12 @@ static SDValue performNanGuardFpToSatCombine(SDNode *N, SelectionDAG &DAG) {
// select (setcc X, 0.0, uno), 0, (and (fp_to_sint/uint X), M)
// select (setcc X, 0.0, ord), (and (fp_to_sint/uint X), M), 0
SDValue X, GuardedVal;
- if (!sd_match(N,
- m_SelectLike(m_OneUse(m_SetCC(m_Value(X), m_AnyZeroFP(),
- m_SpecificCondCode(ISD::SETUO))),
- m_Zero(), m_Value(GuardedVal))) &&
- !sd_match(N,
- m_SelectLike(m_OneUse(m_SetCC(m_Value(X), m_AnyZeroFP(),
- m_SpecificCondCode(ISD::SETO))),
- m_Value(GuardedVal), m_Zero())))
+ if (!sd_match(N, m_SelectLike(m_OneUse(m_SpecificSetCC(ISD::SETUO, m_Value(X),
+ m_AnyZeroFP())),
+ m_Zero(), m_Value(GuardedVal))) &&
+ !sd_match(N, m_SelectLike(m_OneUse(m_SpecificSetCC(ISD::SETO, m_Value(X),
+ m_AnyZeroFP())),
+ m_Value(GuardedVal), m_Zero())))
return SDValue();
// The guarded value must be fp_to_sint/fp_to_uint of the same X, optionally
@@ -13179,8 +13178,7 @@ static SDValue foldVSelectToSignBitSplatMask(SDNode *N, SelectionDAG &DAG) {
SDValue Cond0, Cond1;
ISD::CondCode CC;
- if (!sd_match(N0, m_OneUse(m_SetCC(m_Value(Cond0), m_Value(Cond1),
- m_CondCode(CC)))) ||
+ if (!sd_match(N0, m_OneUse(m_SetCC(CC, m_Value(Cond0), m_Value(Cond1)))) ||
VT != Cond0.getValueType())
return SDValue();
diff --git a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
index 101a3961779fe9..48459f14345f38 100644
--- a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
+++ b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
@@ -30136,8 +30136,8 @@ static SDValue foldMaskedShiftToUSHL(SelectionDAG &DAG,
return SDValue();
unsigned EltSize = VT.getScalarSizeInBits();
- if (!sd_match(Cond, m_SetCC(m_Specific(Amt), m_SpecificInt(EltSize),
- m_SpecificCondCode(RequiredCC))))
+ if (!sd_match(Cond, m_SpecificSetCC(RequiredCC, m_Specific(Amt),
+ m_SpecificInt(EltSize))))
return SDValue();
SDLoc DL(N);
@@ -31715,7 +31715,7 @@ static SDValue performCTPOPCombine(SDNode *N,
EVT CmpVT;
// Use the same VT as the SETcc if -CTPOP would not overflow.
- if (sd_match(Mask, m_SetCC(m_VT(CmpVT), m_Value(), m_Value()))) {
+ if (sd_match(Mask, m_SetCC(m_VT(CmpVT), m_Value()))) {
CmpVT = CmpVT.changeVectorElementTypeToInteger();
if (Log2_64_Ceil(MaskVT.getSizeInBits()) <= CmpVT.getScalarSizeInBits() - 1)
ReduceInVT = CmpVT;
diff --git a/llvm/lib/Target/NVPTX/NVPTXISelLowering.cpp b/llvm/lib/Target/NVPTX/NVPTXISelLowering.cpp
index 5327d39c34d03a..64db745571ac6e 100644
--- a/llvm/lib/Target/NVPTX/NVPTXISelLowering.cpp
+++ b/llvm/lib/Target/NVPTX/NVPTXISelLowering.cpp
@@ -6913,18 +6913,17 @@ static SDValue PerformSELECTShiftCombine(SDNode *N,
m_Shl(m_Value(), m_TruncOrSelf(m_Deferred(ShiftAmt)))));
// shift_amt > BitWidth-1 ? 0 : shift_op
- bool MatchedUGT =
- sd_match(N, m_Select(m_SetCC(m_Value(ShiftAmt),
- m_SpecificInt(APInt(BitWidth, BitWidth - 1)),
- m_SpecificCondCode(ISD::SETUGT)),
- m_Zero(), LogicalShift));
+ bool MatchedUGT = sd_match(
+ N, m_Select(m_SpecificSetCC(ISD::SETUGT, m_Value(ShiftAmt),
+ m_SpecificInt(APInt(BitWidth, BitWidth - 1))),
+ m_Zero(), LogicalShift));
// shift_amt < BitWidth ? shift_op : 0
bool MatchedULT =
!MatchedUGT &&
- sd_match(N, m_Select(m_SetCC(m_Value(ShiftAmt),
- m_SpecificInt(APInt(BitWidth, BitWidth)),
- m_SpecificCondCode(ISD::SETULT)),
- LogicalShift, m_Zero()));
+ sd_match(
+ N, m_Select(m_SpecificSetCC(ISD::SETULT, m_Value(ShiftAmt),
+ m_SpecificInt(APInt(BitWidth, BitWidth))),
+ LogicalShift, m_Zero()));
if (!MatchedUGT && !MatchedULT)
return SDValue();
diff --git a/llvm/lib/Target/RISCV/RISCVISelLowering.cpp b/llvm/lib/Target/RISCV/RISCVISelLowering.cpp
index 3e442bd08d7f85..8952328d1c3e0a 100644
--- a/llvm/lib/Target/RISCV/RISCVISelLowering.cpp
+++ b/llvm/lib/Target/RISCV/RISCVISelLowering.cpp
@@ -22477,8 +22477,8 @@ static SDValue performMaskedLoadToVPLoadCombine(MaskedLoadSDNode *MLoad,
SDValue SetCCLHS, SetCCRHS;
ISD::CondCode CC;
- if (!sd_match(MLoad->getMask(), m_SetCC(m_Value(SetCCLHS), m_Value(SetCCRHS),
- m_CondCode(CC))) ||
+ if (!sd_match(MLoad->getMask(),
+ m_SetCC(CC, m_Value(SetCCLHS), m_Value(SetCCRHS))) ||
SetCCLHS->getOpcode() != ISD::BUILD_VECTOR ||
!(CC == ISD::SETULT || CC == ISD::SETLT) ||
!SetCCLHS.getValueType().isInteger())
@@ -23371,8 +23371,8 @@ static SDValue foldSelectToUSATI(SDNode *N, SelectionDAG &DAG,
using namespace SDPatternMatch;
SDValue Src, InnerSetCC, FalseSrc;
- if (!sd_match(N, m_Select(m_SetCC(m_Value(Src), m_SpecificInt(MaxVal),
- m_SpecificCondCode(ISD::SETUGT)),
+ if (!sd_match(N, m_Select(m_SpecificSetCC(ISD::SETUGT, m_Value(Src),
+ m_SpecificInt(MaxVal)),
m_SExt(m_Value(InnerSetCC)),
m_Trunc(m_Value(FalseSrc)))))
return SDValue();
@@ -23382,9 +23382,10 @@ static SDValue foldSelectToUSATI(SDNode *N, SelectionDAG &DAG,
return SDValue();
// Check inner setcc: src > -1 (signed comparison)
- if (!sd_match(InnerSetCC,
- m_SpecificVT(MVT::i1, m_SetCC(m_Specific(Src), m_AllOnes(),
- m_SpecificCondCode(ISD::SETGT)))))
+ if (!sd_match(
+ InnerSetCC,
+ m_SpecificVT(MVT::i1, m_SpecificSetCC(ISD::SETGT, m_Specific(Src),
+ m_AllOnes()))))
return SDValue();
// It's possible that the input to the setccs is also a truncate, in that
diff --git a/llvm/lib/Target/WebAssembly/WebAssemblyISelLowering.cpp b/llvm/lib/Target/WebAssembly/WebAssemblyISelLowering.cpp
index 1ecd232782679c..3664700009f69d 100644
--- a/llvm/lib/Target/WebAssembly/WebAssemblyISelLowering.cpp
+++ b/llvm/lib/Target/WebAssembly/WebAssemblyISelLowering.cpp
@@ -3411,8 +3411,8 @@ static SDValue performBitcastCombine(SDNode *N,
SDValue Concat, SetCCVector;
ISD::CondCode SetCond;
- if (!sd_match(N, m_BitCast(m_c_SetCC(m_Value(Concat), m_Value(SetCCVector),
- m_CondCode(SetCond)))))
+ if (!sd_match(N, m_BitCast(m_c_SetCC(SetCond, m_Value(Concat),
+ m_Value(SetCCVector)))))
return SDValue();
if (Concat.getOpcode() != ISD::CONCAT_VECTORS)
return SDValue();
@@ -3486,8 +3486,8 @@ static SDValue performBitmaskCombine(SDNode *N, SelectionDAG &DAG) {
return SDValue();
SDValue LHS;
- if (!sd_match(N->getOperand(1), m_c_SetCC(m_Value(LHS), m_Zero(),
- m_SpecificCondCode(ISD::SETLT))))
+ if (!sd_match(N->getOperand(1),
+ m_c_SpecificSetCC(ISD::SETLT, m_Value(LHS), m_Zero())))
return SDValue();
SDLoc DL(N);
@@ -3506,8 +3506,7 @@ static SDValue performAnyAllCombine(SDNode *N, SelectionDAG &DAG) {
SDValue LHS;
if (N->getNumOperands() < 2 ||
- !sd_match(N->getOperand(1),
- m_c_SetCC(m_Value(LHS), m_Zero(), m_CondCode())))
+ !sd_match(N->getOperand(1), m_c_SetCC(m_Value(LHS), m_Zero())))
return SDValue();
EVT LT = LHS.getValueType();
if (LT.getScalarSizeInBits() > 128 / LT.getVectorNumElements())
@@ -3520,8 +3519,8 @@ static SDValue performAnyAllCombine(SDNode *N, SelectionDAG &DAG) {
return SDValue();
SDValue LHS;
- if (!sd_match(N->getOperand(1), m_c_SetCC(m_Value(LHS), m_Zero(),
- m_SpecificCondCode(SetType))))
+ if (!sd_match(N->getOperand(1),
+ m_c_SpecificSetCC(SetType, m_Value(LHS), m_Zero())))
return SDValue();
SDLoc DL(N);
diff --git a/llvm/lib/Target/X86/X86ISelLowering.cpp b/llvm/lib/Target/X86/X86ISelLowering.cpp
index e2f3b5f3cd3d5e..6746e9b380356e 100644
--- a/llvm/lib/Target/X86/X86ISelLowering.cpp
+++ b/llvm/lib/Target/X86/X86ISelLowering.cpp
@@ -48902,10 +48902,9 @@ static SDValue commuteSelect(SDNode *N, SelectionDAG &DAG, const SDLoc &DL,
ISD::CondCode CC;
SDValue Cond, X, Y, LHS, RHS;
- if (!sd_match(
- N, m_VSelect(m_AllOf(m_Value(Cond),
- m_SetCC(m_Value(X), m_Value(Y), m_CondCode(CC))),
- m_Value(LHS), m_Value(RHS))))
+ if (!sd_match(N, m_VSelect(m_AllOf(m_Value(Cond),
+ m_SetCC(CC, m_Value(X), m_Value(Y))),
+ m_Value(LHS), m_Value(RHS))))
return SDValue();
if (canCombineAsMaskOperation(LHS, Subtarget) ||
@@ -49235,9 +49234,9 @@ static SDValue combineSelect(SDNode *N, SelectionDAG &DAG,
if ((LHS.getOpcode() == ISD::SRL || LHS.getOpcode() == ISD::SHL) &&
supportedVectorVarShift(VT, Subtarget, LHS.getOpcode()) &&
ISD::isConstantSplatVectorAllZeros(RHS.getNode()) &&
- sd_match(Cond, m_SetCC(m_Specific(LHS.getOperand(1)),
- m_SpecificInt(VT.getScalarSizeInBits()),
- m_SpecificCondCode(ISD::SETULT)))) {
+ sd_match(Cond,
+ m_SpecificSetCC(ISD::SETULT, m_Specific(LHS.getOperand(1)),
+ m_SpecificInt(VT.getScalarSizeInBits())))) {
return DAG.getNode(LHS.getOpcode() == ISD::SRL ? X86ISD::VSRLV
: X86ISD::VSHLV,
DL, VT, LHS.getOperand(0), LHS.getOperand(1));
@@ -49247,9 +49246,9 @@ static SDValue combineSelect(SDNode *N, SelectionDAG &DAG,
if ((RHS.getOpcode() == ISD::SRL || RHS.getOpcode() == ISD::SHL) &&
supportedVectorVarShift(VT, Subtarget, RHS.getOpcode()) &&
ISD::isConstantSplatVectorAllZeros(LHS.getNode()) &&
- sd_match(Cond, m_SetCC(m_Specific(RHS.getOperand(1)),
- m_SpecificInt(VT.getScalarSizeInBits()),
- m_SpecificCondCode(ISD::SETUGE)))) {
+ sd_match(Cond,
+ m_SpecificSetCC(ISD::SETUGE, m_Specific(RHS.getOperand(1)),
+ m_SpecificInt(VT.getScalarSizeInBits())))) {
return DAG.getNode(RHS.getOpcode() == ISD::SRL ? X86ISD::VSRLV
: X86ISD::VSHLV,
DL, VT, RHS.getOperand(0), RHS.getOperand(1));
@@ -51431,14 +51430,14 @@ static SDValue combineShiftLeft(SDNode *N, SelectionDAG &DAG,
SDValue N01 = N0.getOperand(2);
// fold shl(select(icmp_ult(amt,BW),x,0),amt) -> avx2 psllv(x,amt)
if (ISD::isConstantSplatVectorAllZeros(N01.getNode()) &&
- sd_match(Cond, m_SetCC(m_Specific(N1), m_SpecificInt(EltSizeInBits),
- m_SpecificCondCode(ISD::SETULT)))) {
+ sd_match(Cond, m_SpecificSetCC(ISD::SETULT, m_Specific(N1),
+ m_SpecificInt(EltSizeInBits)))) {
return DAG.getNode(X86ISD::VSHLV, DL, VT, N00, N1);
}
// fold shl(select(icmp_uge(amt,BW),0,x),amt) -> avx2 psllv(x,amt)
if (ISD::isConstantSplatVectorAllZeros(N00.getNode()) &&
- sd_match(Cond, m_SetCC(m_Specific(N1), m_SpecificInt(EltSizeInBits),
- m_SpecificCondCode(ISD::SETUGE)))) {
+ sd_match(Cond, m_SpecificSetCC(ISD::SETUGE, m_Specific(N1),
+ m_SpecificInt(EltSizeInBits)))) {
return DAG.getNode(X86ISD::VSHLV, DL, VT, N01, N1);
}
}
@@ -51570,14 +51569,14 @@ static SDValue combineShiftRightLogical(SDNode *N, SelectionDAG &DAG,
SDValue N01 = N0.getOperand(2);
// fold srl(select(icmp_ult(amt,BW),x,0),amt) -> avx2 psrlv(x,amt)
if (ISD::isConstantSplatVectorAllZeros(N01.getNode()) &&
- sd_match(Cond, m_SetCC(m_Specific(N1), m_SpecificInt(EltSizeInBits),
- m_SpecificCondCode(ISD::SETULT)))) {
+ sd_match(Cond, m_SpecificSetCC(ISD::SETULT, m_Specific(N1),
+ m_SpecificInt(EltSizeInBits)))) {
return DAG.getNode(X86ISD::VSRLV, DL, VT, N00, N1);
}
// fold srl(select(icmp_uge(amt,BW),0,x),amt) -> avx2 psrlv(x,amt)
if (ISD::isConstantSplatVectorAllZeros(N00.getNode()) &&
- sd_match(Cond, m_SetCC(m_Specific(N1), m_SpecificInt(EltSizeInBits),
- m_SpecificCondCode(ISD::SETUGE)))) {
+ sd_match(Cond, m_SpecificSetCC(ISD::SETUGE, m_Specific(N1),
+ m_SpecificInt(EltSizeInBits)))) {
return DAG.getNode(X86ISD::VSRLV, DL, VT, N01, N1);
}
}
@@ -53392,10 +53391,9 @@ static SDValue combineAnd(SDNode *N, SelectionDAG &DAG,
if (TLI.isTypeLegal(VT) && TLI.isTypeLegal(CondVT) &&
(VT.is512BitVector() || Subtarget.hasVLX()) &&
(VT.getScalarSizeInBits() >= 32 || Subtarget.hasBWI()) &&
- sd_match(N, m_And(m_Value(X),
- m_OneUse(m_SExt(m_AllOf(
- m_Value(Y), m_SpecificVT(CondVT),
- m_SetCC(m_Value(), m_Value(), m_Value()))))))) {
+ sd_match(N, m_And(m_Value(X), m_OneUse(m_SExt(m_AllOf(
+ m_Value(Y), m_SpecificVT(CondVT),
+ m_SetCC(m_Value(), m_Value()))))))) {
return DAG.getSelect(dl, VT, Y, X,
getZeroVector(VT.getSimpleVT(), Subtarget, DAG, dl));
}
diff --git a/llvm/unittests/CodeGen/SelectionDAGPatternMatchTest.cpp b/llvm/unittests/CodeGen/SelectionDAGPatternMatchTest.cpp
index 02aab153191b1d..96b733ca359e1a 100644
--- a/llvm/unittests/CodeGen/SelectionDAGPatternMatchTest.cpp
+++ b/llvm/unittests/CodeGen/SelectionDAGPatternMatchTest.cpp
@@ -124,26 +124,32 @@ TEST_F(SelectionDAGPatternMatchTest, matchTernaryOp) {
using namespace SDPatternMatch;
ISD::CondCode CC;
- EXPECT_TRUE(sd_match(ICMP_UGT, m_SetCC(m_Value(), m_Value(),
- m_SpecificCondCode(ISD::SETUGT))));
EXPECT_TRUE(
- sd_match(ICMP_UGT, m_SetCC(m_Value(), m_Value(), m_CondCode(CC))));
+ sd_match(ICMP_UGT, m_SpecificSetCC(ISD::SETUGT, m_Value(), m_Value())));
+ EXPECT_TRUE(sd_match(ICMP_UGT, m_SetCC(CC, m_Value(), m_Value())));
EXPECT_TRUE(CC == ISD::SETUGT);
- EXPECT_FALSE(sd_match(
- ICMP_UGT, m_SetCC(m_Value(), m_Value(), m_SpecificCondCode(ISD::SETLE))));
-
- EXPECT_TRUE(sd_match(ICMP_EQ01, m_SetCC(m_Specific(Op0), m_Specific(Op1),
- m_SpecificCondCode(ISD::SETEQ))));
- EXPECT_TRUE(sd_match(ICMP_EQ10, m_SetCC(m_Specific(Op1), m_Specific(Op0),
- m_SpecificCondCode(ISD::SETEQ))));
- EXPECT_FALSE(sd_match(ICMP_EQ01, m_SetCC(m_Specific(Op1), m_Specific(Op0),
- m_SpecificCondCode(ISD::SETEQ))));
- EXPECT_FALSE(sd_match(ICMP_EQ10, m_SetCC(m_Specific(Op0), m_Specific(Op1),
- m_SpecificCondCode(ISD::SETEQ))));
- EXPECT_TRUE(sd_match(ICMP_EQ01, m_c_SetCC(m_Specific(Op1), m_Specific(Op0),
- m_SpecificCondCode(ISD::SETEQ))));
- EXPECT_TRUE(sd_match(ICMP_EQ10, m_c_SetCC(m_Specific(Op0), m_Specific(Op1),
- m_SpecificCondCode(ISD::SETEQ))));
+ EXPECT_FALSE(
+ sd_match(ICMP_UGT, m_SpecificSetCC(ISD::SETLE, m_Value(), m_Value())));
+
+ EXPECT_TRUE(sd_match(ICMP_EQ01, m_SpecificSetCC(ISD::SETEQ, m_Specific(Op0),
+ m_Specific(Op1))));
+ EXPECT_TRUE(sd_match(ICMP_EQ10, m_SpecificSetCC(ISD::SETEQ, m_Specific(Op1),
+ m_Specific(Op0))));
+ EXPECT_FALSE(sd_match(ICMP_EQ01, m_SpecificSetCC(ISD::SETEQ, m_Specific(Op1),
+ m_Specific(Op0))));
+ EXPECT_FALSE(sd_match(ICMP_EQ10, m_SpecificSetCC(ISD::SETEQ, m_Specific(Op0),
+ m_Specific(Op1))));
+ EXPECT_TRUE(sd_match(ICMP_EQ01, m_c_SpecificSetCC(ISD::SETEQ, m_Specific(Op1),
+ m_Specific(Op0))));
+ EXPECT_TRUE(sd_match(ICMP_EQ10, m_c_SpecificSetCC(ISD::SETEQ, m_Specific(Op0),
+ m_Specific(Op1))));
+ EXPECT_TRUE(sd_match(ICMP_UGT, m_SetCC(m_Value(), m_Value())));
+ EXPECT_FALSE(sd_match(Select, m_SetCC(m_Value(), m_Value())));
+ EXPECT_TRUE(sd_match(ICMP_EQ01, m_c_SetCC(m_Specific(Op1), m_Specific(Op0))));
+ CC = ISD::SETCC_INVALID;
+ EXPECT_TRUE(
+ sd_match(ICMP_EQ10, m_c_SetCC(CC, m_Specific(Op0), m_Specific(Op1))));
+ EXPECT_TRUE(CC == ISD::SETEQ);
EXPECT_TRUE(sd_match(
Select, m_Select(m_Specific(Cond), m_Specific(T), m_Specific(F))));
@@ -1218,10 +1224,43 @@ TEST_F(SelectionDAGPatternMatchTest, MatchSelectCCLike) {
SDValue Select = DAG->getNode(ISD::SELECT_CC, SDLoc(), MVT::i32, LHS, RHS,
TVal, FVal, DAG->getCondCode(ISD::SETLT));
- ISD::CondCode CC = ISD::SETLT;
- EXPECT_TRUE(sd_match(
- Select, m_SelectCCLike(m_Specific(LHS), m_Specific(RHS), m_Specific(TVal),
- m_Specific(FVal), m_CondCode(CC))));
+ ISD::CondCode CC = ISD::SETCC_INVALID;
+ EXPECT_TRUE(
+ sd_match(Select, m_SelectCCLike(CC, m_Specific(LHS), m_Specific(RHS),
+ m_Specific(TVal), m_Specific(FVal))));
+ EXPECT_TRUE(CC == ISD::SETLT);
+ EXPECT_TRUE(
+ sd_match(Select, m_SelectCCLike(m_Specific(LHS), m_Specific(RHS),
+ m_Specific(TVal), m_Specific(FVal))));
+ EXPECT_TRUE(sd_match(Select, m_SpecificSelectCCLike(
+ ISD::SETLT, m_Specific(LHS), m_Specific(RHS),
+ m_Specific(TVal), m_Specific(FVal))));
+ EXPECT_FALSE(
+ sd_match(Select, m_SpecificSelectCCLike(ISD::SETGT, m_Specific(LHS),
+ m_Specific(RHS), m_Specific(TVal),
+ m_Specific(FVal))));
+
+ // Use non-constant operands so the SETCC isn't constant folded.
+ SDValue X = DAG->getCopyFromReg(DAG->getEntryNode(), SDLoc(),
+ Register::index2VirtReg(1), MVT::i32);
+ SDValue Y = DAG->getCopyFromReg(DAG->getEntryNode(), SDLoc(),
+ Register::index2VirtReg(2), MVT::i32);
+ SDValue Cond = DAG->getSetCC(SDLoc(), MVT::i1, X, Y, ISD::SETULT);
+ SDValue SelectOfSetCC =
+ DAG->getNode(ISD::SELECT, SDLoc(), MVT::i32, Cond, TVal, FVal);
+ CC = ISD::SETCC_INVALID;
+ EXPECT_TRUE(sd_match(SelectOfSetCC,
+ m_SelectCCLike(CC, m_Specific(X), m_Specific(Y),
+ m_Specific(TVal), m_Specific(FVal))));
+ EXPECT_TRUE(CC == ISD::SETULT);
+ EXPECT_TRUE(
+ sd_match(SelectOfSetCC,
+ m_SpecificSelectCCLike(ISD::SETULT, m_Specific(X), m_Specific(Y),
+ m_Specific(TVal), m_Specific(FVal))));
+ EXPECT_FALSE(
+ sd_match(SelectOfSetCC,
+ m_SpecificSelectCCLike(ISD::SETLT, m_Specific(X), m_Specific(Y),
+ m_Specific(TVal), m_Specific(FVal))));
}
TEST_F(SelectionDAGPatternMatchTest, MatchSelectCC) {
@@ -1234,10 +1273,18 @@ TEST_F(SelectionDAGPatternMatchTest, MatchSelectCC) {
SDValue Select = DAG->getNode(ISD::SELECT_CC, SDLoc(), MVT::i32, LHS, RHS,
TVal, FVal, DAG->getCondCode(ISD::SETLT));
- ISD::CondCode CC = ISD::SETLT;
+ ISD::CondCode CC = ISD::SETCC_INVALID;
+ EXPECT_TRUE(sd_match(Select, m_SelectCC(CC, m_Specific(LHS), m_Specific(RHS),
+ m_Specific(TVal), m_Specific(FVal))));
+ EXPECT_TRUE(CC == ISD::SETLT);
EXPECT_TRUE(sd_match(Select, m_SelectCC(m_Specific(LHS), m_Specific(RHS),
- m_Specific(TVal), m_Specific(FVal),
- m_CondCode(CC))));
+ m_Specific(TVal), m_Specific(FVal))));
+ EXPECT_TRUE(sd_match(
+ Select, m_SpecificSelectCC(ISD::SETLT, m_Specific(LHS), m_Specific(RHS),
+ m_Specific(TVal), m_Specific(FVal))));
+ EXPECT_FALSE(sd_match(
+ Select, m_SpecificSelectCC(ISD::SETGE, m_Specific(LHS), m_Specific(RHS),
+ m_Specific(TVal), m_Specific(FVal))));
}
TEST_F(SelectionDAGPatternMatchTest, MatchSpecificNeg) {
More information about the llvm-commits
mailing list