[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