[llvm] [SDPatternMatch] Add m_Node and m_SpecificOpc that take opcode from template argument (PR #228626)

Min-Yih Hsu via llvm-commits llvm-commits at lists.llvm.org
Fri Oct 2 17:07:25 PDT 2026


https://github.com/mshockwave created https://github.com/llvm/llvm-project/pull/228626

In most cases, `m_Node` and `m_SpecificOpc` is checking against a fixed opcode. By taking the opcode it is comparing against from a template argument, we can guarantee that opcode comparison will be trivial, even in debug builds. This could potentially improve the compilation time in the long run.

>From 75a2bf4557a75b81fade063b4e1f558039f27ad4 Mon Sep 17 00:00:00 2001
From: Min-Yih Hsu <min.hsu at sifive.com>
Date: Fri, 2 Oct 2026 16:57:49 -0700
Subject: [PATCH] [SDPatternMatch] Add m_Node and m_SpecificOpc that take
 Opcode from template argument

---
 llvm/include/llvm/CodeGen/SDPatternMatch.h    | 14 +++++++++++
 llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp | 24 +++++++++----------
 .../Target/AArch64/AArch64ISelLowering.cpp    | 20 ++++++++--------
 llvm/lib/Target/RISCV/RISCVISelLowering.cpp   |  6 ++---
 llvm/lib/Target/X86/X86ISelLowering.cpp       | 18 +++++++-------
 .../CodeGen/SelectionDAGPatternMatchTest.cpp  |  6 +++++
 6 files changed, 54 insertions(+), 34 deletions(-)

diff --git a/llvm/include/llvm/CodeGen/SDPatternMatch.h b/llvm/include/llvm/CodeGen/SDPatternMatch.h
index ad29772ce7063..6e9f20080ed7e 100644
--- a/llvm/include/llvm/CodeGen/SDPatternMatch.h
+++ b/llvm/include/llvm/CodeGen/SDPatternMatch.h
@@ -100,6 +100,10 @@ struct Opcode_match {
   bool match(SDValue N) { return N->getOpcode() == Opcode; }
 };
 
+template <unsigned Opcode> struct FixedOpcode_match {
+  bool match(SDValue N) { return N->getOpcode() == Opcode; }
+};
+
 // === Patterns combinators ===
 template <typename... Preds> struct And {
   bool match(SDValue N) { return true; }
@@ -152,6 +156,10 @@ template <typename... Preds> auto m_NoneOf(const Preds &...preds) {
   return m_Unless(m_AnyOf(preds...));
 }
 
+template <unsigned Opcode> inline auto m_SpecificOpc() {
+  return FixedOpcode_match<Opcode>();
+}
+
 inline Opcode_match m_SpecificOpc(unsigned Opcode) {
   return Opcode_match(Opcode);
 }
@@ -392,6 +400,12 @@ struct Operands_match<OpIdx, OpndPred, OpndPreds...>
   }
 };
 
+template <unsigned Opcode, typename... OpndPreds>
+auto m_Node(const OpndPreds &...Preds) {
+  return m_AllOf(m_SpecificOpc<Opcode>(),
+                 Operands_match<0, OpndPreds...>(Preds...));
+}
+
 template <typename... OpndPreds>
 auto m_Node(unsigned Opcode, const OpndPreds &...preds) {
   return m_AllOf(m_SpecificOpc(Opcode),
diff --git a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
index ea2955050332e..68221440a9467 100644
--- a/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
+++ b/llvm/lib/CodeGen/SelectionDAG/DAGCombiner.cpp
@@ -5076,7 +5076,8 @@ SDValue DAGCombiner::visitMUL(SDNode *N) {
   }
 
   // fold (mul (add x, c1), c2) -> (add (mul x, c2), c1*c2)
-  if (sd_match(N0, m_SpecificOpc(ISD::ADD)) && isConstantOrConstantVector(N1) &&
+  if (sd_match(N0, m_SpecificOpc<ISD::ADD>()) &&
+      isConstantOrConstantVector(N1) &&
       isConstantOrConstantVector(N0.getOperand(1)) &&
       isMulAddWithConstProfitable(N, N0, N1))
     return DAG.getNode(
@@ -11940,7 +11941,7 @@ SDValue DAGCombiner::visitSRL(SDNode *N) {
           N0,
           m_OneUse(m_BitwiseLogic(
               m_Value(X),
-              m_OneUse(m_Shl(m_Value(ZExtY, m_SpecificOpc(ISD::ZERO_EXTEND)),
+              m_OneUse(m_Shl(m_Value(ZExtY, m_SpecificOpc<ISD::ZERO_EXTEND>()),
                              m_Specific(N1))))))) {
     unsigned NumLeadingZeros = ZExtY.getScalarValueSizeInBits() -
                                ZExtY.getOperand(0).getScalarValueSizeInBits();
@@ -12148,15 +12149,15 @@ SDValue DAGCombiner::visitFunnelShift(SDNode *N) {
     unsigned C1Expected = IsFSHL ? BitWidth - ShAmt : ShAmt;
 
     if ((sd_match(N0, m_Srl(m_Value(Val), m_SpecificInt(C0Expected))) ||
-         sd_match(N0, m_Node(ISD::FSHR, m_Value(), m_Value(Val),
-                             m_SpecificInt(C0Expected))) ||
-         sd_match(N0, m_Node(ISD::FSHL, m_Value(), m_Value(Val),
-                             m_SpecificInt(C1Expected)))) &&
+         sd_match(N0, m_Node<ISD::FSHR>(m_Value(), m_Value(Val),
+                                        m_SpecificInt(C0Expected))) ||
+         sd_match(N0, m_Node<ISD::FSHL>(m_Value(), m_Value(Val),
+                                        m_SpecificInt(C1Expected)))) &&
         (sd_match(N1, m_Shl(m_Specific(Val), m_SpecificInt(C1Expected))) ||
-         sd_match(N1, m_Node(ISD::FSHL, m_Specific(Val), m_Value(),
-                             m_SpecificInt(C1Expected))) ||
-         sd_match(N1, m_Node(ISD::FSHR, m_Specific(Val), m_Value(),
-                             m_SpecificInt(C0Expected)))))
+         sd_match(N1, m_Node<ISD::FSHL>(m_Specific(Val), m_Value(),
+                                        m_SpecificInt(C1Expected))) ||
+         sd_match(N1, m_Node<ISD::FSHR>(m_Specific(Val), m_Value(),
+                                        m_SpecificInt(C0Expected)))))
       return Val;
 
     // fold (fshl ld1, ld0, c) -> (ld0[ofs]) iff ld0 and ld1 are consecutive.
@@ -27606,8 +27607,7 @@ static SDValue combineConcatVectorOfShuffles(SDNode *N, SelectionDAG &DAG,
                                              bool LegalOperations) {
   SDValue A, B;
   ArrayRef<int> M0, M1;
-  if (!sd_match(N,
-                m_Node(ISD::CONCAT_VECTORS,
+  if (!sd_match(N, m_Node<ISD::CONCAT_VECTORS>(
                        m_OneUse(m_Shuffle(m_NUses<2>(m_Value(A)),
                                           m_NUses<2>(m_Value(B)), m_Mask(M0))),
                        m_OneUse(m_Shuffle(m_Deferred(A), m_Deferred(B),
diff --git a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
index 9885dfc854d11..ca1e84a2ff58b 100644
--- a/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
+++ b/llvm/lib/Target/AArch64/AArch64ISelLowering.cpp
@@ -12239,8 +12239,8 @@ SDValue AArch64TargetLowering::LowerBR_CC(SDValue Op, SelectionDAG &DAG) const {
     SDValue Flags;
     uint64_t InverseCC;
     // `CSET <Wd>, <cond>` is an alias of `CSINC <Wd>, WZR, WZR, invert(<cond>)`
-    auto m_CSET = m_Node(AArch64ISD::CSINC, m_Zero(), m_Zero(),
-                         m_ConstInt(InverseCC), m_Value(Flags));
+    auto m_CSET = m_Node<AArch64ISD::CSINC>(
+        m_Zero(), m_Zero(), m_ConstInt(InverseCC), m_Value(Flags));
     // Note: We look through `& 1` as the result of CSET is known to be 0 or 1.
     if ((CC == ISD::SETEQ || CC == ISD::SETNE) && isNullConstant(RHS) &&
         sd_match(LHS, m_AnyOf(m_CSET, m_And(m_CSET, m_One())))) {
@@ -22597,14 +22597,14 @@ static bool hasSVEMultiVectorOps(const AArch64Subtarget *Subtarget) {
 
 static auto m_PredicateAsCounterWhile() {
   using namespace llvm::SDPatternMatch;
-  return m_AnyOf(m_SpecificOpc(AArch64ISD::WHILEGE_PRED_COUNTER),
-                 m_SpecificOpc(AArch64ISD::WHILEGT_PRED_COUNTER),
-                 m_SpecificOpc(AArch64ISD::WHILELT_PRED_COUNTER),
-                 m_SpecificOpc(AArch64ISD::WHILELE_PRED_COUNTER),
-                 m_SpecificOpc(AArch64ISD::WHILEHS_PRED_COUNTER),
-                 m_SpecificOpc(AArch64ISD::WHILEHI_PRED_COUNTER),
-                 m_SpecificOpc(AArch64ISD::WHILELO_PRED_COUNTER),
-                 m_SpecificOpc(AArch64ISD::WHILELS_PRED_COUNTER));
+  return m_AnyOf(m_SpecificOpc<AArch64ISD::WHILEGE_PRED_COUNTER>(),
+                 m_SpecificOpc<AArch64ISD::WHILEGT_PRED_COUNTER>(),
+                 m_SpecificOpc<AArch64ISD::WHILELT_PRED_COUNTER>(),
+                 m_SpecificOpc<AArch64ISD::WHILELE_PRED_COUNTER>(),
+                 m_SpecificOpc<AArch64ISD::WHILEHS_PRED_COUNTER>(),
+                 m_SpecificOpc<AArch64ISD::WHILEHI_PRED_COUNTER>(),
+                 m_SpecificOpc<AArch64ISD::WHILELO_PRED_COUNTER>(),
+                 m_SpecificOpc<AArch64ISD::WHILELS_PRED_COUNTER>());
 }
 
 /// Folds extracting the first lane from the first segment of a
diff --git a/llvm/lib/Target/RISCV/RISCVISelLowering.cpp b/llvm/lib/Target/RISCV/RISCVISelLowering.cpp
index c96a481705cec..c50b5987bc3fd 100644
--- a/llvm/lib/Target/RISCV/RISCVISelLowering.cpp
+++ b/llvm/lib/Target/RISCV/RISCVISelLowering.cpp
@@ -19757,7 +19757,7 @@ static SDValue combineNarrowableShiftedLoad(SDNode *N, SelectionDAG &DAG) {
   APInt MaskVal, ShiftVal;
   // (and (shl (load ...), ShiftAmt), Mask)
   if (!sd_match(
-          N, m_And(m_OneUse(m_Shl(m_Value(LoadNode, m_SpecificOpc(ISD::LOAD)),
+          N, m_And(m_OneUse(m_Shl(m_Value(LoadNode, m_SpecificOpc<ISD::LOAD>()),
                                   m_ConstInt(ShiftVal))),
                    m_ConstInt(MaskVal)))) {
     return SDValue();
@@ -22434,7 +22434,7 @@ static SDValue performBITREVERSECombine(SDNode *N, SelectionDAG &DAG,
 static auto m_ReverseEVL = [](auto X, auto EVL) {
   using namespace SDPatternMatch;
   return m_AnyOf(m_SpliceRight(m_OneUse(m_VectorReverse(X)), m_Poison(), EVL),
-                 m_Node(ISD::EXPERIMENTAL_VP_REVERSE, X, m_Value(), EVL));
+                 m_Node<ISD::EXPERIMENTAL_VP_REVERSE>(X, m_Value(), EVL));
 };
 
 // TODO: A vlse.v is not necessarily faster than a vrgather.vv on all uarchs.
@@ -26145,7 +26145,7 @@ SDValue RISCVTargetLowering::PerformDAGCombine(SDNode *N,
     if (!N->getOperand(0).isUndef() ||
         !sd_match(N->getOperand(2),
                   m_AnyOf(m_ExtractElt(m_Value(SrcVec), m_Zero()),
-                          m_Node(RISCVISD::VMV_X_S, m_Value(SrcVec)))))
+                          m_Node<RISCVISD::VMV_X_S>(m_Value(SrcVec)))))
       break;
 
     MVT SrcVecVT = SrcVec.getSimpleValueType();
diff --git a/llvm/lib/Target/X86/X86ISelLowering.cpp b/llvm/lib/Target/X86/X86ISelLowering.cpp
index a84d017d9c6cd..292a849b155cb 100644
--- a/llvm/lib/Target/X86/X86ISelLowering.cpp
+++ b/llvm/lib/Target/X86/X86ISelLowering.cpp
@@ -53461,11 +53461,11 @@ 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_Value(
-                      Y, m_SpecificVT(CondVT, m_SpecificOpc(ISD::SETCC)))))))) {
+        sd_match(N,
+                 m_And(m_Value(X),
+                       m_OneUse(m_SExt(m_Value(
+                           Y, m_SpecificVT(CondVT,
+                                           m_SpecificOpc<ISD::SETCC>()))))))) {
       return DAG.getSelect(dl, VT, Y, X,
                            getZeroVector(VT.getSimpleVT(), Subtarget, DAG, dl));
     }
@@ -60350,12 +60350,12 @@ static SDValue matchPMADDWD(SelectionDAG &DAG, SDNode *N,
     return SDValue();
 
   SDValue Op0, Op1, Accum;
-  if (!sd_match(N, m_Add(m_Value(Op0, m_SpecificOpc(ISD::BUILD_VECTOR)),
-                         m_Value(Op1, m_SpecificOpc(ISD::BUILD_VECTOR)))) &&
+  if (!sd_match(N, m_Add(m_Value(Op0, m_SpecificOpc<ISD::BUILD_VECTOR>()),
+                         m_Value(Op1, m_SpecificOpc<ISD::BUILD_VECTOR>()))) &&
       !sd_match(N,
-                m_Add(m_Value(Op0, m_SpecificOpc(ISD::BUILD_VECTOR)),
+                m_Add(m_Value(Op0, m_SpecificOpc<ISD::BUILD_VECTOR>()),
                       m_Add(m_Value(Accum),
-                            m_Value(Op1, m_SpecificOpc(ISD::BUILD_VECTOR))))))
+                            m_Value(Op1, m_SpecificOpc<ISD::BUILD_VECTOR>())))))
     return SDValue();
 
   // Check if one of Op0,Op1 is of the form:
diff --git a/llvm/unittests/CodeGen/SelectionDAGPatternMatchTest.cpp b/llvm/unittests/CodeGen/SelectionDAGPatternMatchTest.cpp
index 96b733ca359e1..b8e5996fcd9e7 100644
--- a/llvm/unittests/CodeGen/SelectionDAGPatternMatchTest.cpp
+++ b/llvm/unittests/CodeGen/SelectionDAGPatternMatchTest.cpp
@@ -302,6 +302,10 @@ TEST_F(SelectionDAGPatternMatchTest, matchBinaryOp) {
   EXPECT_FALSE(sd_match(Add, m_NSWAddLike(m_Value(), m_Value())));
   EXPECT_TRUE(sd_match(Mul, m_Mul(m_OneUse(m_SpecificOpc(ISD::SUB)),
                                   m_NUses<2>(m_Specific(Add)))));
+  EXPECT_TRUE(sd_match(Mul, m_Mul(m_OneUse(m_SpecificOpc<ISD::SUB>()),
+                                  m_NUses<2>(m_Specific(Add)))));
+  EXPECT_FALSE(sd_match(Mul, m_Mul(m_OneUse(m_SpecificOpc<ISD::MUL>()),
+                                   m_NUses<2>(m_Specific(Add)))));
   EXPECT_TRUE(
       sd_match(SFAdd, m_ChainedBinOp(ISD::STRICT_FADD, m_SpecificVT(Float32VT),
                                      m_SpecificVT(Float32VT))));
@@ -861,7 +865,9 @@ TEST_F(SelectionDAGPatternMatchTest, matchNode) {
 
   using namespace SDPatternMatch;
   EXPECT_TRUE(sd_match(Add, m_Node(ISD::ADD, m_Value(), m_Value())));
+  EXPECT_TRUE(sd_match(Add, m_Node<ISD::ADD>(m_Value(), m_Value())));
   EXPECT_FALSE(sd_match(Add, m_Node(ISD::SUB, m_Value(), m_Value())));
+  EXPECT_FALSE(sd_match(Add, m_Node<ISD::SUB>(m_Value(), m_Value())));
   EXPECT_FALSE(sd_match(Add, m_Node(ISD::ADD, m_Value())));
   EXPECT_FALSE(
       sd_match(Add, m_Node(ISD::ADD, m_Value(), m_Value(), m_Value())));



More information about the llvm-commits mailing list