[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