1//==--------------- llvm/CodeGen/SDPatternMatch.h ---------------*- C++ -*-===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8/// \file
9/// Contains matchers for matching SelectionDAG nodes and values.
10///
11//===----------------------------------------------------------------------===//
12
13#ifndef LLVM_CODEGEN_SDPATTERNMATCH_H
14#define LLVM_CODEGEN_SDPATTERNMATCH_H
15
16#include "llvm/ADT/APInt.h"
17#include "llvm/ADT/ArrayRef.h"
18#include "llvm/ADT/STLExtras.h"
19#include "llvm/ADT/SmallBitVector.h"
20#include "llvm/ADT/bit.h"
21#include "llvm/CodeGen/SelectionDAG.h"
22#include "llvm/CodeGen/SelectionDAGNodes.h"
23#include "llvm/CodeGen/TargetLowering.h"
24#include "llvm/Support/KnownBits.h"
25
26#include <type_traits>
27
28namespace llvm {
29namespace SDPatternMatch {
30
31template <typename Pattern>
32[[nodiscard]] bool sd_match(SDValue N, Pattern &&P) {
33 return P.match(N);
34}
35
36template <typename Pattern>
37[[nodiscard]] bool sd_match(SDNode *N, Pattern &&P) {
38 return sd_match(SDValue(N, 0), P);
39}
40
41// === Utilities ===
42struct Value_match {
43 SDValue MatchVal;
44
45 Value_match() = default;
46
47 explicit Value_match(SDValue Match) : MatchVal(Match) {}
48
49 bool match(SDValue N) {
50 if (MatchVal)
51 return MatchVal == N;
52 return N.getNode();
53 }
54};
55
56/// Match any valid SDValue.
57inline Value_match m_Value() { return Value_match(); }
58
59inline Value_match m_Specific(SDValue N) {
60 assert(N);
61 return Value_match(N);
62}
63
64template <unsigned ResNo, typename Pattern> struct Result_match {
65 Pattern P;
66
67 explicit Result_match(const Pattern &P) : P(P) {}
68
69 bool match(SDValue N) { return N.getResNo() == ResNo && P.match(N); }
70};
71
72/// Match only if the SDValue is a certain result at ResNo.
73template <unsigned ResNo, typename Pattern>
74inline Result_match<ResNo, Pattern> m_Result(const Pattern &P) {
75 return Result_match<ResNo, Pattern>(P);
76}
77
78struct DeferredValue_match {
79 SDValue &MatchVal;
80
81 explicit DeferredValue_match(SDValue &Match) : MatchVal(Match) {}
82
83 bool match(SDValue N) { return N == MatchVal; }
84};
85
86/// Similar to m_Specific, but the specific value to match is determined by
87/// another sub-pattern in the same sd_match() expression. For instance,
88/// We cannot match `(add V, V)` with `m_Add(m_Value(X), m_Specific(X))` since
89/// `X` is not initialized at the time it got copied into `m_Specific`. Instead,
90/// we should use `m_Add(m_Value(X), m_Deferred(X))`.
91inline DeferredValue_match m_Deferred(SDValue &V) {
92 return DeferredValue_match(V);
93}
94
95struct Opcode_match {
96 unsigned Opcode;
97
98 explicit Opcode_match(unsigned Opc) : Opcode(Opc) {}
99
100 bool match(SDValue N) { return N->getOpcode() == Opcode; }
101};
102
103template <unsigned Opcode> struct FixedOpcode_match {
104 bool match(SDValue N) { return N->getOpcode() == Opcode; }
105};
106
107// === Patterns combinators ===
108template <typename... Preds> struct And {
109 bool match(SDValue N) { return true; }
110};
111
112template <typename Pred, typename... Preds>
113struct And<Pred, Preds...> : And<Preds...> {
114 Pred P;
115 And(const Pred &p, const Preds &...preds) : And<Preds...>(preds...), P(p) {}
116
117 bool match(SDValue N) { return P.match(N) && And<Preds...>::match(N); }
118};
119
120template <typename... Preds> struct Or {
121 bool match(SDValue N) { return false; }
122};
123
124template <typename Pred, typename... Preds>
125struct Or<Pred, Preds...> : Or<Preds...> {
126 Pred P;
127 Or(const Pred &p, const Preds &...preds) : Or<Preds...>(preds...), P(p) {}
128
129 bool match(SDValue N) { return P.match(N) || Or<Preds...>::match(N); }
130};
131
132template <typename Pred> struct Not {
133 Pred P;
134
135 explicit Not(const Pred &P) : P(P) {}
136
137 bool match(SDValue N) { return !P.match(N); }
138};
139// Explicit deduction guide.
140template <typename Pred> Not(const Pred &P) -> Not<Pred>;
141
142/// Match if the inner pattern does NOT match.
143template <typename Pred> inline Not<Pred> m_Unless(const Pred &P) {
144 return Not{P};
145}
146
147template <typename... Preds> And<Preds...> m_AllOf(const Preds &...preds) {
148 return And<Preds...>(preds...);
149}
150
151template <typename... Preds> Or<Preds...> m_AnyOf(const Preds &...preds) {
152 return Or<Preds...>(preds...);
153}
154
155template <typename... Preds> auto m_NoneOf(const Preds &...preds) {
156 return m_Unless(m_AnyOf(preds...));
157}
158
159template <unsigned Opcode> inline auto m_SpecificOpc() {
160 return FixedOpcode_match<Opcode>();
161}
162
163inline Opcode_match m_SpecificOpc(unsigned Opcode) {
164 return Opcode_match(Opcode);
165}
166
167inline auto m_Undef() {
168 return m_AnyOf(preds: Opcode_match(ISD::UNDEF), preds: Opcode_match(ISD::POISON));
169}
170
171inline Opcode_match m_Poison() { return Opcode_match(ISD::POISON); }
172
173template <unsigned NumUses, typename Pattern> struct NUses_match {
174 Pattern P;
175
176 explicit NUses_match(const Pattern &P) : P(P) {}
177
178 bool match(SDValue N) {
179 // SDNode::hasNUsesOfValue is pretty expensive when the SDNode produces
180 // multiple results, hence we check the subsequent pattern here before
181 // checking the number of value users.
182 return P.match(N) && N->hasNUsesOfValue(NUses: NumUses, Value: N.getResNo());
183 }
184};
185
186template <typename Pattern>
187inline NUses_match<1, Pattern> m_OneUse(const Pattern &P) {
188 return NUses_match<1, Pattern>(P);
189}
190template <unsigned N, typename Pattern>
191inline NUses_match<N, Pattern> m_NUses(const Pattern &P) {
192 return NUses_match<N, Pattern>(P);
193}
194
195inline NUses_match<1, Value_match> m_OneUse() {
196 return NUses_match<1, Value_match>(m_Value());
197}
198template <unsigned N> inline NUses_match<N, Value_match> m_NUses() {
199 return NUses_match<N, Value_match>(m_Value());
200}
201
202struct Value_bind {
203 SDValue &BindVal;
204
205 Value_bind(SDValue &N) : BindVal(N) {}
206
207 bool match(SDValue N) {
208 BindVal = N;
209 return true;
210 }
211};
212
213inline auto m_Value(SDValue &N) { return Value_bind(N); }
214/// Conditionally bind an SDValue based on the predicate.
215template <typename PredPattern>
216inline auto m_Value(SDValue &N, const PredPattern &P) {
217 return m_AllOf(P, Value_bind(N));
218}
219
220template <typename Pattern, typename PredFuncT> struct TLI_pred_match {
221 Pattern P;
222 PredFuncT PredFunc;
223
224 TLI_pred_match(const PredFuncT &Pred, const Pattern &P)
225 : P(P), PredFunc(Pred) {}
226
227 bool match(SDValue N) { return PredFunc(N) && P.match(N); }
228};
229
230// Explicit deduction guide.
231template <typename PredFuncT, typename Pattern>
232TLI_pred_match(const PredFuncT &Pred, const Pattern &P)
233 -> TLI_pred_match<Pattern, PredFuncT>;
234
235/// Match legal SDNodes based on the information provided by TargetLowering.
236template <typename Pattern>
237inline auto m_LegalOp(const SelectionDAG &DAG, const Pattern &P) {
238 return TLI_pred_match{[&DAG](SDValue N) {
239 return DAG.getTargetLoweringInfo().isOperationLegal(
240 Op: N->getOpcode(), VT: N.getValueType());
241 },
242 P};
243}
244
245// === Value type ===
246
247template <typename Pattern> struct ValueType_bind {
248 EVT &BindVT;
249 Pattern P;
250
251 explicit ValueType_bind(EVT &Bind, const Pattern &P) : BindVT(Bind), P(P) {}
252
253 bool match(SDValue N) {
254 BindVT = N.getValueType();
255 return P.match(N);
256 }
257};
258
259template <typename Pattern>
260ValueType_bind(const Pattern &P) -> ValueType_bind<Pattern>;
261
262/// Retreive the ValueType of the current SDValue.
263inline auto m_VT(EVT &VT) { return ValueType_bind(VT, m_Value()); }
264
265template <typename Pattern> inline auto m_VT(EVT &VT, const Pattern &P) {
266 return ValueType_bind(VT, P);
267}
268
269template <typename Pattern, typename PredFuncT> struct ValueType_match {
270 PredFuncT PredFunc;
271 Pattern P;
272
273 ValueType_match(const PredFuncT &Pred, const Pattern &P)
274 : PredFunc(Pred), P(P) {}
275
276 bool match(SDValue N) { return PredFunc(N.getValueType()) && P.match(N); }
277};
278
279// Explicit deduction guide.
280template <typename PredFuncT, typename Pattern>
281ValueType_match(const PredFuncT &Pred, const Pattern &P)
282 -> ValueType_match<Pattern, PredFuncT>;
283
284/// Match a specific ValueType.
285template <typename Pattern>
286inline auto m_SpecificVT(EVT RefVT, const Pattern &P) {
287 return ValueType_match{[=](EVT VT) { return VT == RefVT; }, P};
288}
289inline auto m_SpecificVT(EVT RefVT) {
290 return ValueType_match{[=](EVT VT) { return VT == RefVT; }, m_Value()};
291}
292
293inline auto m_Glue() { return m_SpecificVT(RefVT: MVT::Glue); }
294inline auto m_OtherVT() { return m_SpecificVT(RefVT: MVT::Other); }
295
296/// Match a scalar ValueType.
297template <typename Pattern>
298inline auto m_SpecificScalarVT(EVT RefVT, const Pattern &P) {
299 return ValueType_match{[=](EVT VT) { return VT.getScalarType() == RefVT; },
300 P};
301}
302inline auto m_SpecificScalarVT(EVT RefVT) {
303 return ValueType_match{[=](EVT VT) { return VT.getScalarType() == RefVT; },
304 m_Value()};
305}
306
307/// Match a vector ValueType.
308template <typename Pattern>
309inline auto m_SpecificVectorElementVT(EVT RefVT, const Pattern &P) {
310 return ValueType_match{[=](EVT VT) {
311 return VT.isVector() &&
312 VT.getVectorElementType() == RefVT;
313 },
314 P};
315}
316inline auto m_SpecificVectorElementVT(EVT RefVT) {
317 return ValueType_match{[=](EVT VT) {
318 return VT.isVector() &&
319 VT.getVectorElementType() == RefVT;
320 },
321 m_Value()};
322}
323
324/// Match any integer ValueTypes.
325template <typename Pattern> inline auto m_IntegerVT(const Pattern &P) {
326 return ValueType_match{[](EVT VT) { return VT.isInteger(); }, P};
327}
328inline auto m_IntegerVT() {
329 return ValueType_match{[](EVT VT) { return VT.isInteger(); }, m_Value()};
330}
331
332/// Match any floating point ValueTypes.
333template <typename Pattern> inline auto m_FloatingPointVT(const Pattern &P) {
334 return ValueType_match{[](EVT VT) { return VT.isFloatingPoint(); }, P};
335}
336inline auto m_FloatingPointVT() {
337 return ValueType_match{[](EVT VT) { return VT.isFloatingPoint(); },
338 m_Value()};
339}
340
341/// Match any vector ValueTypes.
342template <typename Pattern> inline auto m_VectorVT(const Pattern &P) {
343 return ValueType_match{[](EVT VT) { return VT.isVector(); }, P};
344}
345inline auto m_VectorVT() {
346 return ValueType_match{[](EVT VT) { return VT.isVector(); }, m_Value()};
347}
348
349/// Match fixed-length vector ValueTypes.
350template <typename Pattern> inline auto m_FixedVectorVT(const Pattern &P) {
351 return ValueType_match{[](EVT VT) { return VT.isFixedLengthVector(); }, P};
352}
353inline auto m_FixedVectorVT() {
354 return ValueType_match{[](EVT VT) { return VT.isFixedLengthVector(); },
355 m_Value()};
356}
357
358/// Match scalable vector ValueTypes.
359template <typename Pattern> inline auto m_ScalableVectorVT(const Pattern &P) {
360 return ValueType_match{[](EVT VT) { return VT.isScalableVector(); }, P};
361}
362inline auto m_ScalableVectorVT() {
363 return ValueType_match{[](EVT VT) { return VT.isScalableVector(); },
364 m_Value()};
365}
366
367/// Match legal ValueTypes based on the information provided by TargetLowering.
368template <typename Pattern>
369inline auto m_LegalType(const SelectionDAG &DAG, const Pattern &P) {
370 return TLI_pred_match{[&DAG](SDValue N) {
371 return DAG.getTargetLoweringInfo().isTypeLegal(
372 VT: N.getValueType());
373 },
374 P};
375}
376
377// === Generic node matching ===
378template <unsigned OpIdx, typename... OpndPreds> struct Operands_match {
379 bool match(SDValue N) {
380 // Returns false if there are more operands than predicates;
381 return N->getNumOperands() == OpIdx;
382 }
383};
384
385template <unsigned OpIdx, typename OpndPred, typename... OpndPreds>
386struct Operands_match<OpIdx, OpndPred, OpndPreds...>
387 : Operands_match<OpIdx + 1, OpndPreds...> {
388 OpndPred P;
389
390 Operands_match(const OpndPred &p, const OpndPreds &...preds)
391 : Operands_match<OpIdx + 1, OpndPreds...>(preds...), P(p) {}
392
393 bool match(SDValue N) {
394 if (OpIdx < N->getNumOperands())
395 return P.match(N->getOperand(Num: OpIdx)) &&
396 Operands_match<OpIdx + 1, OpndPreds...>::match(N);
397
398 // This is the case where there are more predicates than operands.
399 return false;
400 }
401};
402
403template <unsigned Opcode, typename... OpndPreds>
404auto m_Node(const OpndPreds &...Preds) {
405 return m_AllOf(m_SpecificOpc<Opcode>(),
406 Operands_match<0, OpndPreds...>(Preds...));
407}
408
409template <typename... OpndPreds>
410auto m_Node(unsigned Opcode, const OpndPreds &...preds) {
411 return m_AllOf(m_SpecificOpc(Opcode),
412 Operands_match<0, OpndPreds...>(preds...));
413}
414
415/// Provide number of operands that are not chain or glue, as well as the first
416/// index of such operand.
417template <bool ExcludeChain> struct EffectiveOperands {
418 unsigned Size = 0;
419 unsigned FirstIndex = 0;
420
421 explicit EffectiveOperands(SDValue N) : Size(N->getNumOperands()) {
422 if (ExcludeChain) {
423 // Glue if present, is the last operand.
424 if (Size != 0 && N->getOperand(Num: Size - 1).getValueType() == MVT::Glue)
425 --Size;
426 // Chain if present, is the first operand.
427 if (Size != 0 && N->getOperand(Num: 0).getValueType() == MVT::Other) {
428 ++FirstIndex;
429 --Size;
430 }
431 }
432 }
433};
434
435// === Ternary operations ===
436template <typename T0_P, typename T1_P, typename T2_P, bool Commutable = false,
437 bool ExcludeChain = false>
438struct TernaryOpc_match {
439 unsigned Opcode;
440 T0_P Op0;
441 T1_P Op1;
442 T2_P Op2;
443
444 TernaryOpc_match(unsigned Opc, const T0_P &Op0, const T1_P &Op1,
445 const T2_P &Op2)
446 : Opcode(Opc), Op0(Op0), Op1(Op1), Op2(Op2) {}
447
448 bool match(SDValue N) {
449 if (sd_match(N, P: m_SpecificOpc(Opcode))) {
450 EffectiveOperands<ExcludeChain> EO(N);
451 assert(EO.Size == 3);
452 return ((Op0.match(N->getOperand(Num: EO.FirstIndex)) &&
453 Op1.match(N->getOperand(Num: EO.FirstIndex + 1))) ||
454 (Commutable && Op0.match(N->getOperand(Num: EO.FirstIndex + 1)) &&
455 Op1.match(N->getOperand(Num: EO.FirstIndex)))) &&
456 Op2.match(N->getOperand(Num: EO.FirstIndex + 2));
457 }
458
459 return false;
460 }
461};
462
463struct CondCode_match {
464 std::optional<ISD::CondCode> CCToMatch;
465 ISD::CondCode *BindCC = nullptr;
466
467 explicit CondCode_match(ISD::CondCode CC) : CCToMatch(CC) {}
468
469 explicit CondCode_match(ISD::CondCode *CC) : BindCC(CC) {}
470
471 bool match(SDValue N) {
472 if (auto *CC = dyn_cast<CondCodeSDNode>(Val: N.getNode())) {
473 if (CCToMatch && *CCToMatch != CC->get())
474 return false;
475
476 if (BindCC)
477 *BindCC = CC->get();
478 return true;
479 }
480
481 return false;
482 }
483};
484
485/// Match any conditional code SDNode.
486inline CondCode_match m_CondCode() { return CondCode_match(nullptr); }
487/// Match any conditional code SDNode and return its ISD::CondCode value.
488inline CondCode_match m_CondCode(ISD::CondCode &CC) {
489 return CondCode_match(&CC);
490}
491/// Match a conditional code SDNode with a specific ISD::CondCode.
492inline CondCode_match m_SpecificCondCode(ISD::CondCode CC) {
493 return CondCode_match(CC);
494}
495
496/// Match a SETCC with any condition code.
497template <typename T0_P, typename T1_P>
498inline TernaryOpc_match<T0_P, T1_P, CondCode_match> m_SetCC(const T0_P &LHS,
499 const T1_P &RHS) {
500 return TernaryOpc_match<T0_P, T1_P, CondCode_match>(ISD::SETCC, LHS, RHS,
501 m_CondCode());
502}
503
504/// Match a SETCC with any condition code and bind the condition code to CC.
505template <typename T0_P, typename T1_P>
506inline TernaryOpc_match<T0_P, T1_P, CondCode_match>
507m_SetCC(ISD::CondCode &CC, const T0_P &LHS, const T1_P &RHS) {
508 return TernaryOpc_match<T0_P, T1_P, CondCode_match>(ISD::SETCC, LHS, RHS,
509 m_CondCode(CC));
510}
511
512/// Match a SETCC with a specific condition code.
513template <typename T0_P, typename T1_P>
514inline TernaryOpc_match<T0_P, T1_P, CondCode_match>
515m_SpecificSetCC(ISD::CondCode CC, const T0_P &LHS, const T1_P &RHS) {
516 return TernaryOpc_match<T0_P, T1_P, CondCode_match>(ISD::SETCC, LHS, RHS,
517 m_SpecificCondCode(CC));
518}
519
520/// Match a SETCC with any condition code, allowing the operands to be
521/// commuted.
522template <typename T0_P, typename T1_P>
523inline TernaryOpc_match<T0_P, T1_P, CondCode_match, true, false>
524m_c_SetCC(const T0_P &LHS, const T1_P &RHS) {
525 return TernaryOpc_match<T0_P, T1_P, CondCode_match, true, false>(
526 ISD::SETCC, LHS, RHS, m_CondCode());
527}
528
529/// Match a SETCC with any condition code, allowing the operands to be
530/// commuted, and bind the condition code to CC.
531template <typename T0_P, typename T1_P>
532inline TernaryOpc_match<T0_P, T1_P, CondCode_match, true, false>
533m_c_SetCC(ISD::CondCode &CC, const T0_P &LHS, const T1_P &RHS) {
534 return TernaryOpc_match<T0_P, T1_P, CondCode_match, true, false>(
535 ISD::SETCC, LHS, RHS, m_CondCode(CC));
536}
537
538/// Match a SETCC with a specific condition code, allowing the operands to be
539/// commuted.
540template <typename T0_P, typename T1_P>
541inline TernaryOpc_match<T0_P, T1_P, CondCode_match, true, false>
542m_c_SpecificSetCC(ISD::CondCode CC, const T0_P &LHS, const T1_P &RHS) {
543 return TernaryOpc_match<T0_P, T1_P, CondCode_match, true, false>(
544 ISD::SETCC, LHS, RHS, m_SpecificCondCode(CC));
545}
546
547template <typename T0_P, typename T1_P, typename T2_P>
548inline TernaryOpc_match<T0_P, T1_P, T2_P>
549m_Select(const T0_P &Cond, const T1_P &T, const T2_P &F) {
550 return TernaryOpc_match<T0_P, T1_P, T2_P>(ISD::SELECT, Cond, T, F);
551}
552
553template <typename T0_P, typename T1_P, typename T2_P>
554inline TernaryOpc_match<T0_P, T1_P, T2_P>
555m_VSelect(const T0_P &Cond, const T1_P &T, const T2_P &F) {
556 return TernaryOpc_match<T0_P, T1_P, T2_P>(ISD::VSELECT, Cond, T, F);
557}
558
559template <typename T0_P, typename T1_P, typename T2_P>
560inline auto m_SelectLike(const T0_P &Cond, const T1_P &T, const T2_P &F) {
561 return m_AnyOf(m_Select(Cond, T, F), m_VSelect(Cond, T, F));
562}
563
564template <typename T0_P, typename T1_P, typename T2_P>
565inline Result_match<0, TernaryOpc_match<T0_P, T1_P, T2_P>>
566m_Load(const T0_P &Ch, const T1_P &Ptr, const T2_P &Offset) {
567 return m_Result<0>(
568 TernaryOpc_match<T0_P, T1_P, T2_P>(ISD::LOAD, Ch, Ptr, Offset));
569}
570
571template <typename T0_P, typename T1_P, typename T2_P>
572inline TernaryOpc_match<T0_P, T1_P, T2_P>
573m_InsertElt(const T0_P &Vec, const T1_P &Val, const T2_P &Idx) {
574 return TernaryOpc_match<T0_P, T1_P, T2_P>(ISD::INSERT_VECTOR_ELT, Vec, Val,
575 Idx);
576}
577
578template <typename LHS, typename RHS, typename IDX>
579inline TernaryOpc_match<LHS, RHS, IDX>
580m_InsertSubvector(const LHS &Base, const RHS &Sub, const IDX &Idx) {
581 return TernaryOpc_match<LHS, RHS, IDX>(ISD::INSERT_SUBVECTOR, Base, Sub, Idx);
582}
583
584template <typename T0_P, typename T1_P, typename T2_P>
585inline TernaryOpc_match<T0_P, T1_P, T2_P>
586m_SpliceRight(const T0_P &V1, const T1_P &V2, const T2_P &Offset) {
587 return TernaryOpc_match<T0_P, T1_P, T2_P>(ISD::VECTOR_SPLICE_RIGHT, V1, V2,
588 Offset);
589}
590
591template <typename T0_P, typename T1_P, typename T2_P>
592inline TernaryOpc_match<T0_P, T1_P, T2_P>
593m_TernaryOp(unsigned Opc, const T0_P &Op0, const T1_P &Op1, const T2_P &Op2) {
594 return TernaryOpc_match<T0_P, T1_P, T2_P>(Opc, Op0, Op1, Op2);
595}
596
597template <typename T0_P, typename T1_P, typename T2_P>
598inline TernaryOpc_match<T0_P, T1_P, T2_P, true>
599m_c_TernaryOp(unsigned Opc, const T0_P &Op0, const T1_P &Op1, const T2_P &Op2) {
600 return TernaryOpc_match<T0_P, T1_P, T2_P, true>(Opc, Op0, Op1, Op2);
601}
602
603/// Match a SELECT_CC with any condition code.
604template <typename LTy, typename RTy, typename TTy, typename FTy>
605inline auto m_SelectCC(const LTy &L, const RTy &R, const TTy &T, const FTy &F) {
606 return m_Node(ISD::SELECT_CC, L, R, T, F, m_CondCode());
607}
608
609/// Match a SELECT_CC with any condition code and bind the condition code to
610/// CC.
611template <typename LTy, typename RTy, typename TTy, typename FTy>
612inline auto m_SelectCC(ISD::CondCode &CC, const LTy &L, const RTy &R,
613 const TTy &T, const FTy &F) {
614 return m_Node(ISD::SELECT_CC, L, R, T, F, m_CondCode(CC));
615}
616
617/// Match a SELECT_CC with a specific condition code.
618template <typename LTy, typename RTy, typename TTy, typename FTy>
619inline auto m_SpecificSelectCC(ISD::CondCode CC, const LTy &L, const RTy &R,
620 const TTy &T, const FTy &F) {
621 return m_Node(ISD::SELECT_CC, L, R, T, F, m_SpecificCondCode(CC));
622}
623
624/// Match a SELECT of a SETCC or a SELECT_CC with any condition code.
625template <typename LTy, typename RTy, typename TTy, typename FTy>
626inline auto m_SelectCCLike(const LTy &L, const RTy &R, const TTy &T,
627 const FTy &F) {
628 return m_AnyOf(m_Select(m_SetCC(L, R), T, F), m_SelectCC(L, R, T, F));
629}
630
631/// Match a SELECT of a SETCC or a SELECT_CC with any condition code and bind
632/// the condition code to CC.
633template <typename LTy, typename RTy, typename TTy, typename FTy>
634inline auto m_SelectCCLike(ISD::CondCode &CC, const LTy &L, const RTy &R,
635 const TTy &T, const FTy &F) {
636 return m_AnyOf(m_Select(m_SetCC(CC, L, R), T, F), m_SelectCC(CC, L, R, T, F));
637}
638
639/// Match a SELECT of a SETCC or a SELECT_CC with a specific condition code.
640template <typename LTy, typename RTy, typename TTy, typename FTy>
641inline auto m_SpecificSelectCCLike(ISD::CondCode CC, const LTy &L, const RTy &R,
642 const TTy &T, const FTy &F) {
643 return m_AnyOf(m_Select(m_SpecificSetCC(CC, L, R), T, F),
644 m_SpecificSelectCC(CC, L, R, T, F));
645}
646
647// === Binary operations ===
648template <typename LHS_P, typename RHS_P, bool Commutable = false,
649 bool ExcludeChain = false>
650struct BinaryOpc_match {
651 unsigned Opcode;
652 LHS_P LHS;
653 RHS_P RHS;
654 SDNodeFlags Flags;
655 BinaryOpc_match(unsigned Opc, const LHS_P &L, const RHS_P &R,
656 SDNodeFlags Flgs = SDNodeFlags())
657 : Opcode(Opc), LHS(L), RHS(R), Flags(Flgs) {}
658
659 bool match(SDValue N) {
660 if (sd_match(N, P: m_SpecificOpc(Opcode))) {
661 EffectiveOperands<ExcludeChain> EO(N);
662 assert(EO.Size == 2);
663 if (!((LHS.match(N->getOperand(Num: EO.FirstIndex)) &&
664 RHS.match(N->getOperand(Num: EO.FirstIndex + 1))) ||
665 (Commutable && LHS.match(N->getOperand(Num: EO.FirstIndex + 1)) &&
666 RHS.match(N->getOperand(Num: EO.FirstIndex)))))
667 return false;
668
669 return (Flags & N->getFlags()) == Flags;
670 }
671
672 return false;
673 }
674};
675
676/// Matching while capturing mask
677template <typename T0, typename T1, typename T2> struct SDShuffle_match {
678 T0 Op1;
679 T1 Op2;
680 T2 Mask;
681
682 SDShuffle_match(const T0 &Op1, const T1 &Op2, const T2 &Mask)
683 : Op1(Op1), Op2(Op2), Mask(Mask) {}
684
685 bool match(SDValue N) {
686 if (auto *I = dyn_cast<ShuffleVectorSDNode>(Val&: N)) {
687 return Op1.match(I->getOperand(Num: 0)) && Op2.match(I->getOperand(Num: 1)) &&
688 Mask.match(I->getMask());
689 }
690 return false;
691 }
692};
693struct m_Mask {
694 ArrayRef<int> &MaskRef;
695 m_Mask(ArrayRef<int> &MaskRef) : MaskRef(MaskRef) {}
696 bool match(ArrayRef<int> Mask) {
697 MaskRef = Mask;
698 return true;
699 }
700};
701
702struct m_SpecificMask {
703 ArrayRef<int> MaskRef;
704 m_SpecificMask(ArrayRef<int> MaskRef) : MaskRef(MaskRef) {}
705 bool match(ArrayRef<int> Mask) { return MaskRef == Mask; }
706};
707
708template <typename LHS_P, typename RHS_P, typename Pred_t,
709 bool Commutable = false>
710struct MaxMin_match {
711 using PredType = Pred_t;
712 LHS_P LHS;
713 RHS_P RHS;
714
715 MaxMin_match(const LHS_P &L, const RHS_P &R) : LHS(L), RHS(R) {}
716
717 bool match(SDValue N) {
718 auto MatchMinMax = [&](SDValue L, SDValue R, SDValue TrueValue,
719 SDValue FalseValue, ISD::CondCode CC) {
720 if ((TrueValue != L || FalseValue != R) &&
721 (TrueValue != R || FalseValue != L))
722 return false;
723
724 ISD::CondCode Cond =
725 TrueValue == L ? CC : getSetCCInverse(Operation: CC, Type: L.getValueType());
726 if (!Pred_t::match(Cond))
727 return false;
728
729 return (LHS.match(L) && RHS.match(R)) ||
730 (Commutable && LHS.match(R) && RHS.match(L));
731 };
732
733 if (N.getOpcode() == ISD::SELECT || N.getOpcode() == ISD::VSELECT) {
734 assert(N.getNumOperands() == 3);
735 SDValue Cond = N.getOperand(i: 0);
736 SDValue TrueValue = N.getOperand(i: 1);
737 SDValue FalseValue = N.getOperand(i: 2);
738
739 if (Cond.getOpcode() == ISD::SETCC) {
740 assert(Cond.getNumOperands() == 3);
741 SDValue L = Cond.getOperand(i: 0);
742 SDValue R = Cond.getOperand(i: 1);
743 ISD::CondCode CC = cast<CondCodeSDNode>(Val: Cond.getOperand(i: 2))->get();
744 return MatchMinMax(L, R, TrueValue, FalseValue, CC);
745 }
746 }
747
748 if (N.getOpcode() == ISD::SELECT_CC) {
749 assert(N.getNumOperands() == 5);
750 SDValue L = N.getOperand(i: 0);
751 SDValue R = N.getOperand(i: 1);
752 SDValue TrueValue = N.getOperand(i: 2);
753 SDValue FalseValue = N.getOperand(i: 3);
754 ISD::CondCode CC = cast<CondCodeSDNode>(Val: N->getOperand(Num: 4))->get();
755 return MatchMinMax(L, R, TrueValue, FalseValue, CC);
756 }
757
758 return false;
759 }
760};
761
762// Helper class for identifying signed max predicates.
763struct smax_pred_ty {
764 static bool match(ISD::CondCode Cond) {
765 return Cond == ISD::CondCode::SETGT || Cond == ISD::CondCode::SETGE;
766 }
767};
768
769// Helper class for identifying unsigned max predicates.
770struct umax_pred_ty {
771 static bool match(ISD::CondCode Cond) {
772 return Cond == ISD::CondCode::SETUGT || Cond == ISD::CondCode::SETUGE;
773 }
774};
775
776// Helper class for identifying signed min predicates.
777struct smin_pred_ty {
778 static bool match(ISD::CondCode Cond) {
779 return Cond == ISD::CondCode::SETLT || Cond == ISD::CondCode::SETLE;
780 }
781};
782
783// Helper class for identifying unsigned min predicates.
784struct umin_pred_ty {
785 static bool match(ISD::CondCode Cond) {
786 return Cond == ISD::CondCode::SETULT || Cond == ISD::CondCode::SETULE;
787 }
788};
789
790template <typename LHS, typename RHS>
791inline BinaryOpc_match<LHS, RHS> m_BinOp(unsigned Opc, const LHS &L,
792 const RHS &R,
793 SDNodeFlags Flgs = SDNodeFlags()) {
794 return BinaryOpc_match<LHS, RHS>(Opc, L, R, Flgs);
795}
796template <typename LHS, typename RHS>
797inline BinaryOpc_match<LHS, RHS, true>
798m_c_BinOp(unsigned Opc, const LHS &L, const RHS &R,
799 SDNodeFlags Flgs = SDNodeFlags()) {
800 return BinaryOpc_match<LHS, RHS, true>(Opc, L, R, Flgs);
801}
802
803template <typename LHS, typename RHS>
804inline BinaryOpc_match<LHS, RHS, false, true>
805m_ChainedBinOp(unsigned Opc, const LHS &L, const RHS &R) {
806 return BinaryOpc_match<LHS, RHS, false, true>(Opc, L, R);
807}
808template <typename LHS, typename RHS>
809inline BinaryOpc_match<LHS, RHS, true, true>
810m_c_ChainedBinOp(unsigned Opc, const LHS &L, const RHS &R) {
811 return BinaryOpc_match<LHS, RHS, true, true>(Opc, L, R);
812}
813
814// Common binary operations
815template <typename LHS, typename RHS>
816inline BinaryOpc_match<LHS, RHS, true> m_Add(const LHS &L, const RHS &R) {
817 return BinaryOpc_match<LHS, RHS, true>(ISD::ADD, L, R);
818}
819
820template <typename LHS, typename RHS>
821inline auto m_NUWAdd(const LHS &L, const RHS &R) {
822 return BinaryOpc_match<LHS, RHS, true>(ISD::ADD, L, R,
823 SDNodeFlags::NoUnsignedWrap);
824}
825
826template <typename LHS, typename RHS>
827inline auto m_NSWAdd(const LHS &L, const RHS &R) {
828 return BinaryOpc_match<LHS, RHS, true>(ISD::ADD, L, R,
829 SDNodeFlags::NoSignedWrap);
830}
831
832template <typename LHS, typename RHS>
833inline BinaryOpc_match<LHS, RHS> m_Sub(const LHS &L, const RHS &R) {
834 return BinaryOpc_match<LHS, RHS>(ISD::SUB, L, R);
835}
836
837template <typename LHS, typename RHS>
838inline BinaryOpc_match<LHS, RHS, true> m_Mul(const LHS &L, const RHS &R) {
839 return BinaryOpc_match<LHS, RHS, true>(ISD::MUL, L, R);
840}
841
842template <typename LHS, typename RHS>
843inline BinaryOpc_match<LHS, RHS, true> m_And(const LHS &L, const RHS &R) {
844 return BinaryOpc_match<LHS, RHS, true>(ISD::AND, L, R);
845}
846
847template <typename LHS, typename RHS>
848inline BinaryOpc_match<LHS, RHS, true> m_Or(const LHS &L, const RHS &R) {
849 return BinaryOpc_match<LHS, RHS, true>(ISD::OR, L, R);
850}
851
852template <typename LHS, typename RHS>
853inline BinaryOpc_match<LHS, RHS, true> m_DisjointOr(const LHS &L,
854 const RHS &R) {
855 return BinaryOpc_match<LHS, RHS, true>(ISD::OR, L, R, SDNodeFlags::Disjoint);
856}
857
858template <typename LHS, typename RHS>
859inline auto m_AddLike(const LHS &L, const RHS &R) {
860 return m_AnyOf(m_Add(L, R), m_DisjointOr(L, R));
861}
862
863template <typename LHS, typename RHS>
864inline auto m_NSWAddLike(const LHS &L, const RHS &R) {
865 return m_AnyOf(m_NSWAdd(L, R), m_DisjointOr(L, R));
866}
867
868template <typename LHS, typename RHS>
869inline auto m_NUWAddLike(const LHS &L, const RHS &R) {
870 return m_AnyOf(m_NUWAdd(L, R), m_DisjointOr(L, R));
871}
872
873template <typename LHS, typename RHS>
874inline BinaryOpc_match<LHS, RHS, true> m_Xor(const LHS &L, const RHS &R) {
875 return BinaryOpc_match<LHS, RHS, true>(ISD::XOR, L, R);
876}
877
878template <typename LHS, typename RHS>
879inline auto m_BitwiseLogic(const LHS &L, const RHS &R) {
880 return m_AnyOf(m_And(L, R), m_Or(L, R), m_Xor(L, R));
881}
882
883template <unsigned Opc, typename Pred, typename LHS, typename RHS>
884inline auto m_MaxMinLike(const LHS &L, const RHS &R) {
885 return m_AnyOf(BinaryOpc_match<LHS, RHS, true>(Opc, L, R),
886 MaxMin_match<LHS, RHS, Pred, true>(L, R));
887}
888
889template <typename LHS, typename RHS>
890inline BinaryOpc_match<LHS, RHS, true> m_SMin(const LHS &L, const RHS &R) {
891 return BinaryOpc_match<LHS, RHS, true>(ISD::SMIN, L, R);
892}
893
894template <typename LHS, typename RHS>
895inline auto m_SMinLike(const LHS &L, const RHS &R) {
896 return m_MaxMinLike<ISD::SMIN, smin_pred_ty>(L, R);
897}
898
899template <typename LHS, typename RHS>
900inline BinaryOpc_match<LHS, RHS, true> m_SMax(const LHS &L, const RHS &R) {
901 return BinaryOpc_match<LHS, RHS, true>(ISD::SMAX, L, R);
902}
903
904template <typename LHS, typename RHS>
905inline auto m_SMaxLike(const LHS &L, const RHS &R) {
906 return m_MaxMinLike<ISD::SMAX, smax_pred_ty>(L, R);
907}
908
909template <typename LHS, typename RHS>
910inline BinaryOpc_match<LHS, RHS, true> m_UMin(const LHS &L, const RHS &R) {
911 return BinaryOpc_match<LHS, RHS, true>(ISD::UMIN, L, R);
912}
913
914template <typename LHS, typename RHS>
915inline auto m_UMinLike(const LHS &L, const RHS &R) {
916 return m_MaxMinLike<ISD::UMIN, umin_pred_ty>(L, R);
917}
918
919template <typename LHS, typename RHS>
920inline BinaryOpc_match<LHS, RHS, true> m_UMax(const LHS &L, const RHS &R) {
921 return BinaryOpc_match<LHS, RHS, true>(ISD::UMAX, L, R);
922}
923
924template <typename LHS, typename RHS>
925inline auto m_UMaxLike(const LHS &L, const RHS &R) {
926 return m_MaxMinLike<ISD::UMAX, umax_pred_ty>(L, R);
927}
928
929template <typename LHS, typename RHS>
930inline BinaryOpc_match<LHS, RHS> m_UDiv(const LHS &L, const RHS &R) {
931 return BinaryOpc_match<LHS, RHS>(ISD::UDIV, L, R);
932}
933template <typename LHS, typename RHS>
934inline BinaryOpc_match<LHS, RHS> m_SDiv(const LHS &L, const RHS &R) {
935 return BinaryOpc_match<LHS, RHS>(ISD::SDIV, L, R);
936}
937
938template <typename LHS, typename RHS>
939inline BinaryOpc_match<LHS, RHS> m_URem(const LHS &L, const RHS &R) {
940 return BinaryOpc_match<LHS, RHS>(ISD::UREM, L, R);
941}
942template <typename LHS, typename RHS>
943inline BinaryOpc_match<LHS, RHS> m_SRem(const LHS &L, const RHS &R) {
944 return BinaryOpc_match<LHS, RHS>(ISD::SREM, L, R);
945}
946
947template <typename LHS, typename RHS>
948inline BinaryOpc_match<LHS, RHS> m_Shl(const LHS &L, const RHS &R) {
949 return BinaryOpc_match<LHS, RHS>(ISD::SHL, L, R);
950}
951
952template <typename LHS, typename RHS>
953inline BinaryOpc_match<LHS, RHS> m_Sra(const LHS &L, const RHS &R) {
954 return BinaryOpc_match<LHS, RHS>(ISD::SRA, L, R);
955}
956template <typename LHS, typename RHS>
957inline BinaryOpc_match<LHS, RHS> m_Srl(const LHS &L, const RHS &R) {
958 return BinaryOpc_match<LHS, RHS>(ISD::SRL, L, R);
959}
960template <typename LHS, typename RHS>
961inline auto m_ExactSr(const LHS &L, const RHS &R) {
962 return m_AnyOf(BinaryOpc_match<LHS, RHS>(ISD::SRA, L, R, SDNodeFlags::Exact),
963 BinaryOpc_match<LHS, RHS>(ISD::SRL, L, R, SDNodeFlags::Exact));
964}
965
966template <typename LHS, typename RHS>
967inline BinaryOpc_match<LHS, RHS> m_Rotl(const LHS &L, const RHS &R) {
968 return BinaryOpc_match<LHS, RHS>(ISD::ROTL, L, R);
969}
970
971template <typename LHS, typename RHS>
972inline BinaryOpc_match<LHS, RHS> m_Rotr(const LHS &L, const RHS &R) {
973 return BinaryOpc_match<LHS, RHS>(ISD::ROTR, L, R);
974}
975
976template <typename T0_P, typename T1_P, typename T2_P>
977inline TernaryOpc_match<T0_P, T1_P, T2_P>
978m_FShL(const T0_P &Op0, const T1_P &Op1, const T2_P &Op2) {
979 return m_TernaryOp(ISD::FSHL, Op0, Op1, Op2);
980}
981
982template <typename T0_P, typename T1_P, typename T2_P>
983inline TernaryOpc_match<T0_P, T1_P, T2_P>
984m_FShR(const T0_P &Op0, const T1_P &Op1, const T2_P &Op2) {
985 return m_TernaryOp(ISD::FSHR, Op0, Op1, Op2);
986}
987
988template <typename T0_P, typename T1_P, typename T2_P, bool Left>
989struct FunnelShiftLike_match {
990 T0_P Op0;
991 T1_P Op1;
992 T2_P Op2;
993
994 FunnelShiftLike_match(const T0_P &Op0, const T1_P &Op1, const T2_P &Op2)
995 : Op0(Op0), Op1(Op1), Op2(Op2) {}
996
997 static bool hasComplementaryConstantShifts(const APInt &ShlV,
998 const APInt &SrlV,
999 unsigned BitWidth) {
1000 unsigned SumWidth = std::max(a: ShlV.getBitWidth(), b: SrlV.getBitWidth()) + 1;
1001 unsigned BitWidthBits = llvm::bit_width(Value: BitWidth);
1002 if (BitWidthBits > SumWidth)
1003 return false;
1004
1005 return ShlV.zext(width: SumWidth) + SrlV.zext(width: SumWidth) ==
1006 APInt(SumWidth, BitWidth);
1007 }
1008
1009 bool matchOperands(SDValue X, SDValue Y, SDValue Z) {
1010 return Op0.match(X) && Op1.match(Y) && Op2.match(Z);
1011 }
1012
1013 bool matchShiftOr(SDValue N, unsigned BitWidth);
1014
1015 bool match(SDValue N) {
1016 if (sd_match(N, Left ? m_FShL(Op0, Op1, Op2) : m_FShR(Op0, Op1, Op2)))
1017 return true;
1018
1019 SDValue X, Z;
1020 if (sd_match(N, P: Left ? m_Rotl(L: m_Value(N&: X), R: m_Value(N&: Z))
1021 : m_Rotr(L: m_Value(N&: X), R: m_Value(N&: Z))))
1022 return matchOperands(X, Y: X, Z);
1023
1024 return matchShiftOr(N, BitWidth: N.getValueType().getScalarSizeInBits());
1025 }
1026};
1027
1028template <typename T0_P, typename T1_P, typename T2_P>
1029inline FunnelShiftLike_match<T0_P, T1_P, T2_P, true>
1030m_FShLLike(const T0_P &Op0, const T1_P &Op1, const T2_P &Op2) {
1031 return FunnelShiftLike_match<T0_P, T1_P, T2_P, true>(Op0, Op1, Op2);
1032}
1033
1034template <typename T0_P, typename T1_P, typename T2_P>
1035inline FunnelShiftLike_match<T0_P, T1_P, T2_P, false>
1036m_FShRLike(const T0_P &Op0, const T1_P &Op1, const T2_P &Op2) {
1037 return FunnelShiftLike_match<T0_P, T1_P, T2_P, false>(Op0, Op1, Op2);
1038}
1039
1040template <typename LHS, typename RHS>
1041inline BinaryOpc_match<LHS, RHS, true> m_Clmul(const LHS &L, const RHS &R) {
1042 return BinaryOpc_match<LHS, RHS, true>(ISD::CLMUL, L, R);
1043}
1044
1045template <typename LHS, typename RHS>
1046inline BinaryOpc_match<LHS, RHS, true> m_FAdd(const LHS &L, const RHS &R) {
1047 return BinaryOpc_match<LHS, RHS, true>(ISD::FADD, L, R);
1048}
1049
1050template <typename LHS, typename RHS>
1051inline BinaryOpc_match<LHS, RHS> m_FSub(const LHS &L, const RHS &R) {
1052 return BinaryOpc_match<LHS, RHS>(ISD::FSUB, L, R);
1053}
1054
1055template <typename LHS, typename RHS>
1056inline BinaryOpc_match<LHS, RHS, true> m_FMul(const LHS &L, const RHS &R) {
1057 return BinaryOpc_match<LHS, RHS, true>(ISD::FMUL, L, R);
1058}
1059
1060template <typename LHS, typename RHS>
1061inline BinaryOpc_match<LHS, RHS> m_FDiv(const LHS &L, const RHS &R) {
1062 return BinaryOpc_match<LHS, RHS>(ISD::FDIV, L, R);
1063}
1064
1065template <typename LHS, typename RHS>
1066inline BinaryOpc_match<LHS, RHS> m_FRem(const LHS &L, const RHS &R) {
1067 return BinaryOpc_match<LHS, RHS>(ISD::FREM, L, R);
1068}
1069
1070template <typename V1_t, typename V2_t>
1071inline BinaryOpc_match<V1_t, V2_t> m_Shuffle(const V1_t &v1, const V2_t &v2) {
1072 return BinaryOpc_match<V1_t, V2_t>(ISD::VECTOR_SHUFFLE, v1, v2);
1073}
1074
1075template <typename V1_t, typename V2_t, typename Mask_t>
1076inline SDShuffle_match<V1_t, V2_t, Mask_t>
1077m_Shuffle(const V1_t &v1, const V2_t &v2, const Mask_t &mask) {
1078 return SDShuffle_match<V1_t, V2_t, Mask_t>(v1, v2, mask);
1079}
1080
1081template <typename LHS, typename RHS>
1082inline BinaryOpc_match<LHS, RHS> m_ExtractElt(const LHS &Vec, const RHS &Idx) {
1083 return BinaryOpc_match<LHS, RHS>(ISD::EXTRACT_VECTOR_ELT, Vec, Idx);
1084}
1085
1086template <typename LHS, typename RHS>
1087inline BinaryOpc_match<LHS, RHS> m_ExtractSubvector(const LHS &Vec,
1088 const RHS &Idx) {
1089 return BinaryOpc_match<LHS, RHS>(ISD::EXTRACT_SUBVECTOR, Vec, Idx);
1090}
1091
1092// === Unary operations ===
1093template <typename Opnd_P, bool ExcludeChain = false> struct UnaryOpc_match {
1094 unsigned Opcode;
1095 Opnd_P Opnd;
1096 SDNodeFlags Flags;
1097 UnaryOpc_match(unsigned Opc, const Opnd_P &Op,
1098 SDNodeFlags Flgs = SDNodeFlags())
1099 : Opcode(Opc), Opnd(Op), Flags(Flgs) {}
1100
1101 bool match(SDValue N) {
1102 if (sd_match(N, P: m_SpecificOpc(Opcode))) {
1103 EffectiveOperands<ExcludeChain> EO(N);
1104 assert(EO.Size == 1);
1105 if (!Opnd.match(N->getOperand(Num: EO.FirstIndex)))
1106 return false;
1107
1108 return (Flags & N->getFlags()) == Flags;
1109 }
1110
1111 return false;
1112 }
1113};
1114
1115template <typename Opnd>
1116inline UnaryOpc_match<Opnd> m_UnaryOp(unsigned Opc, const Opnd &Op) {
1117 return UnaryOpc_match<Opnd>(Opc, Op);
1118}
1119template <typename Opnd>
1120inline UnaryOpc_match<Opnd, true> m_ChainedUnaryOp(unsigned Opc,
1121 const Opnd &Op) {
1122 return UnaryOpc_match<Opnd, true>(Opc, Op);
1123}
1124
1125template <typename Opnd> inline UnaryOpc_match<Opnd> m_BitCast(const Opnd &Op) {
1126 return UnaryOpc_match<Opnd>(ISD::BITCAST, Op);
1127}
1128
1129template <typename Opnd>
1130inline UnaryOpc_match<Opnd> m_BSwap(const Opnd &Op) {
1131 return UnaryOpc_match<Opnd>(ISD::BSWAP, Op);
1132}
1133
1134template <typename Opnd>
1135inline UnaryOpc_match<Opnd> m_BitReverse(const Opnd &Op) {
1136 return UnaryOpc_match<Opnd>(ISD::BITREVERSE, Op);
1137}
1138
1139template <typename Opnd> inline UnaryOpc_match<Opnd> m_ZExt(const Opnd &Op) {
1140 return UnaryOpc_match<Opnd>(ISD::ZERO_EXTEND, Op);
1141}
1142
1143template <typename Opnd>
1144inline UnaryOpc_match<Opnd> m_NNegZExt(const Opnd &Op) {
1145 return UnaryOpc_match<Opnd>(ISD::ZERO_EXTEND, Op, SDNodeFlags::NonNeg);
1146}
1147
1148template <typename Opnd> inline auto m_SExt(const Opnd &Op) {
1149 return UnaryOpc_match<Opnd>(ISD::SIGN_EXTEND, Op);
1150}
1151
1152template <typename Opnd> inline UnaryOpc_match<Opnd> m_AnyExt(const Opnd &Op) {
1153 return UnaryOpc_match<Opnd>(ISD::ANY_EXTEND, Op);
1154}
1155
1156template <typename Opnd> inline UnaryOpc_match<Opnd> m_Trunc(const Opnd &Op) {
1157 return UnaryOpc_match<Opnd>(ISD::TRUNCATE, Op);
1158}
1159
1160template <typename Opnd> inline auto m_Abs(const Opnd &Op) {
1161 return m_AnyOf(UnaryOpc_match<Opnd>(ISD::ABS, Op),
1162 UnaryOpc_match<Opnd>(ISD::ABS_MIN_POISON, Op));
1163}
1164
1165template <typename Opnd> inline UnaryOpc_match<Opnd> m_FAbs(const Opnd &Op) {
1166 return UnaryOpc_match<Opnd>(ISD::FABS, Op);
1167}
1168
1169/// Match a zext or identity
1170/// Allows to peek through optional extensions
1171template <typename Opnd> inline auto m_ZExtOrSelf(const Opnd &Op) {
1172 return m_AnyOf(m_ZExt(Op), Op);
1173}
1174
1175/// Match a sext or identity
1176/// Allows to peek through optional extensions
1177template <typename Opnd> inline auto m_SExtOrSelf(const Opnd &Op) {
1178 return m_AnyOf(m_SExt(Op), Op);
1179}
1180
1181template <typename Opnd> inline auto m_SExtLike(const Opnd &Op) {
1182 return m_AnyOf(m_SExt(Op), m_NNegZExt(Op));
1183}
1184
1185/// Match a zext or sext
1186template <typename Opnd> inline auto m_ZExtOrSExt(const Opnd &Op) {
1187 return m_AnyOf(m_ZExt(Op), m_SExt(Op));
1188}
1189
1190/// Match a aext or identity
1191/// Allows to peek through optional extensions
1192template <typename Opnd>
1193inline Or<UnaryOpc_match<Opnd>, Opnd> m_AExtOrSelf(const Opnd &Op) {
1194 return Or<UnaryOpc_match<Opnd>, Opnd>(m_AnyExt(Op), Op);
1195}
1196
1197/// Match a trunc or identity
1198/// Allows to peek through optional truncations
1199template <typename Opnd>
1200inline Or<UnaryOpc_match<Opnd>, Opnd> m_TruncOrSelf(const Opnd &Op) {
1201 return Or<UnaryOpc_match<Opnd>, Opnd>(m_Trunc(Op), Op);
1202}
1203
1204template <typename Opnd> inline UnaryOpc_match<Opnd> m_VScale(const Opnd &Op) {
1205 return UnaryOpc_match<Opnd>(ISD::VSCALE, Op);
1206}
1207
1208template <typename Opnd> inline UnaryOpc_match<Opnd> m_FPToUI(const Opnd &Op) {
1209 return UnaryOpc_match<Opnd>(ISD::FP_TO_UINT, Op);
1210}
1211
1212template <typename Opnd> inline UnaryOpc_match<Opnd> m_FPToSI(const Opnd &Op) {
1213 return UnaryOpc_match<Opnd>(ISD::FP_TO_SINT, Op);
1214}
1215
1216template <typename Opnd> inline UnaryOpc_match<Opnd> m_Ctpop(const Opnd &Op) {
1217 return UnaryOpc_match<Opnd>(ISD::CTPOP, Op);
1218}
1219
1220template <typename Opnd> inline UnaryOpc_match<Opnd> m_Ctlz(const Opnd &Op) {
1221 return UnaryOpc_match<Opnd>(ISD::CTLZ, Op);
1222}
1223
1224template <typename Opnd> inline UnaryOpc_match<Opnd> m_Cttz(const Opnd &Op) {
1225 return UnaryOpc_match<Opnd>(ISD::CTTZ, Op);
1226}
1227
1228template <typename Opnd> inline UnaryOpc_match<Opnd> m_FNeg(const Opnd &Op) {
1229 return UnaryOpc_match<Opnd>(ISD::FNEG, Op);
1230}
1231
1232template <typename Opnd>
1233inline UnaryOpc_match<Opnd> m_VectorReverse(const Opnd &Op) {
1234 return UnaryOpc_match<Opnd>(ISD::VECTOR_REVERSE, Op);
1235}
1236
1237// === Constants ===
1238struct ConstantInt_match {
1239 APInt *BindVal;
1240
1241 explicit ConstantInt_match(APInt *V) : BindVal(V) {}
1242
1243 bool match(SDValue N) {
1244 // The logics here are similar to that in
1245 // SelectionDAG::isConstantIntBuildVectorOrConstantInt, but the latter also
1246 // treats GlobalAddressSDNode as a constant, which is difficult to turn into
1247 // APInt.
1248 if (auto *C = dyn_cast_or_null<ConstantSDNode>(Val: N.getNode())) {
1249 if (BindVal)
1250 *BindVal = C->getAPIntValue();
1251 return true;
1252 }
1253
1254 APInt Discard;
1255 return ISD::isConstantSplatVector(N: N.getNode(),
1256 SplatValue&: BindVal ? *BindVal : Discard);
1257 }
1258};
1259
1260template <typename T> struct Constant64_match {
1261 static_assert(sizeof(T) == 8, "T must be 64 bits wide");
1262
1263 T &BindVal;
1264
1265 explicit Constant64_match(T &V) : BindVal(V) {}
1266
1267 bool match(SDValue N) {
1268 APInt V;
1269 if (!ConstantInt_match(&V).match(N))
1270 return false;
1271
1272 if constexpr (std::is_signed_v<T>) {
1273 if (std::optional<int64_t> TrySExt = V.trySExtValue()) {
1274 BindVal = *TrySExt;
1275 return true;
1276 }
1277 }
1278
1279 if constexpr (std::is_unsigned_v<T>) {
1280 if (std::optional<uint64_t> TryZExt = V.tryZExtValue()) {
1281 BindVal = *TryZExt;
1282 return true;
1283 }
1284 }
1285
1286 return false;
1287 }
1288};
1289
1290/// Match any integer constants or splat of an integer constant.
1291inline ConstantInt_match m_ConstInt() { return ConstantInt_match(nullptr); }
1292/// Match any integer constants or splat of an integer constant; return the
1293/// specific constant or constant splat value.
1294inline ConstantInt_match m_ConstInt(APInt &V) { return ConstantInt_match(&V); }
1295/// Match any integer constants or splat of an integer constant that can fit in
1296/// 64 bits; return the specific constant or constant splat value, zero-extended
1297/// to 64 bits.
1298inline Constant64_match<uint64_t> m_ConstInt(uint64_t &V) {
1299 return Constant64_match<uint64_t>(V);
1300}
1301/// Match any integer constants or splat of an integer constant that can fit in
1302/// 64 bits; return the specific constant or constant splat value, sign-extended
1303/// to 64 bits.
1304inline Constant64_match<int64_t> m_ConstInt(int64_t &V) {
1305 return Constant64_match<int64_t>(V);
1306}
1307
1308template <typename T0_P, typename T1_P, typename T2_P, bool Left>
1309bool FunnelShiftLike_match<T0_P, T1_P, T2_P, Left>::matchShiftOr(
1310 SDValue N, unsigned BitWidth) {
1311 SDValue X, Y, ShlAmt, SrlAmt;
1312 APInt ShlConst, SrlConst;
1313 if (!sd_match(
1314 N, P: m_Or(L: m_Shl(L: m_Value(N&: X), R: m_Value(N&: ShlAmt, P: m_ConstInt(V&: ShlConst))),
1315 R: m_Srl(L: m_Value(N&: Y), R: m_Value(N&: SrlAmt, P: m_ConstInt(V&: SrlConst))))) ||
1316 !hasComplementaryConstantShifts(ShlV: ShlConst, SrlV: SrlConst, BitWidth))
1317 return false;
1318
1319 return matchOperands(X, Y, Z: Left ? ShlAmt : SrlAmt);
1320}
1321
1322struct SpecificInt_match {
1323 APInt IntVal;
1324
1325 explicit SpecificInt_match(APInt APV) : IntVal(std::move(APV)) {}
1326
1327 bool match(SDValue N) {
1328 APInt ConstInt;
1329 if (sd_match(N, P: m_ConstInt(V&: ConstInt)))
1330 return APInt::isSameValue(I1: IntVal, I2: ConstInt);
1331 return false;
1332 }
1333};
1334
1335/// Match a specific integer constant or constant splat value.
1336inline SpecificInt_match m_SpecificInt(APInt V) {
1337 return SpecificInt_match(std::move(V));
1338}
1339inline SpecificInt_match m_SpecificInt(uint64_t V) {
1340 return SpecificInt_match(APInt(64, V));
1341}
1342
1343struct SpecificFP_match {
1344 APFloat Val;
1345
1346 explicit SpecificFP_match(APFloat V) : Val(V) {}
1347
1348 bool match(SDValue V) {
1349 if (const auto *CFP = dyn_cast<ConstantFPSDNode>(Val: V.getNode()))
1350 return CFP->isExactlyValue(V: Val);
1351 if (ConstantFPSDNode *C = isConstOrConstSplatFP(N: V, /*AllowUndefs=*/AllowUndefs: true))
1352 return C->getValueAPF().compare(RHS: Val) == APFloat::cmpEqual;
1353 return false;
1354 }
1355};
1356
1357/// Match a specific float constant.
1358inline SpecificFP_match m_SpecificFP(APFloat V) { return SpecificFP_match(V); }
1359
1360inline SpecificFP_match m_SpecificFP(double V) {
1361 return SpecificFP_match(APFloat(V));
1362}
1363
1364struct AnyZeroFP_match {
1365 bool match(SDValue N) {
1366 if (ConstantFPSDNode *C = isConstOrConstSplatFP(N))
1367 return C->isZero();
1368 return false;
1369 }
1370};
1371
1372/// Match a floating-point +0.0 or -0.0 constant or splat.
1373inline AnyZeroFP_match m_AnyZeroFP() { return AnyZeroFP_match(); }
1374
1375struct Zero_match {
1376 bool AllowUndefs;
1377
1378 explicit Zero_match(bool AllowUndefs) : AllowUndefs(AllowUndefs) {}
1379
1380 bool match(SDValue N) const { return isZeroOrZeroSplat(N, AllowUndefs); }
1381};
1382
1383struct Ones_match {
1384 bool AllowUndefs;
1385
1386 Ones_match(bool AllowUndefs) : AllowUndefs(AllowUndefs) {}
1387
1388 bool match(SDValue N) { return isOnesOrOnesSplat(N, AllowUndefs); }
1389};
1390
1391struct AllOnes_match {
1392 bool AllowUndefs;
1393
1394 AllOnes_match(bool AllowUndefs) : AllowUndefs(AllowUndefs) {}
1395
1396 bool match(SDValue N) { return isAllOnesOrAllOnesSplat(V: N, AllowUndefs); }
1397};
1398
1399inline Ones_match m_One(bool AllowUndefs = false) {
1400 return Ones_match(AllowUndefs);
1401}
1402inline Zero_match m_Zero(bool AllowUndefs = false) {
1403 return Zero_match(AllowUndefs);
1404}
1405inline AllOnes_match m_AllOnes(bool AllowUndefs = false) {
1406 return AllOnes_match(AllowUndefs);
1407}
1408
1409template <bool Expected> struct Bool_match {
1410 const SelectionDAG &DAG;
1411
1412 Bool_match(const SelectionDAG &DAG) : DAG(DAG) {}
1413
1414 bool match(SDValue N) {
1415 auto Res = DAG.isBoolConstant(N);
1416 return Res && *Res == Expected;
1417 }
1418};
1419
1420/// Match true boolean value based on the information provided by
1421/// TargetLowering.
1422inline auto m_True(const SelectionDAG &DAG) { return Bool_match<true>(DAG); }
1423
1424/// Match false boolean value based on the information provided by
1425/// TargetLowering.
1426inline auto m_False(const SelectionDAG &DAG) { return Bool_match<false>(DAG); }
1427
1428/// Match a negate as a sub(0, v)
1429template <typename ValTy>
1430inline BinaryOpc_match<Zero_match, ValTy, false> m_Neg(const ValTy &V) {
1431 return m_Sub(m_Zero(), V);
1432}
1433
1434/// Match a Not as a xor(v, -1) or xor(-1, v)
1435template <typename ValTy>
1436inline BinaryOpc_match<ValTy, AllOnes_match, true> m_Not(const ValTy &V) {
1437 return m_Xor(V, m_AllOnes());
1438}
1439
1440template <unsigned IntrinsicId, typename... OpndPreds>
1441inline auto m_IntrinsicWOChain(const OpndPreds &...Opnds) {
1442 return m_Node(ISD::INTRINSIC_WO_CHAIN, m_SpecificInt(V: IntrinsicId), Opnds...);
1443}
1444
1445struct SpecificNeg_match {
1446 SDValue V;
1447
1448 explicit SpecificNeg_match(SDValue V) : V(V) {}
1449
1450 bool match(SDValue N) {
1451 if (sd_match(N, P: m_Neg(V: m_Specific(N: V))))
1452 return true;
1453
1454 return ISD::matchBinaryPredicate(
1455 LHS: V, RHS: N, Match: [](ConstantSDNode *LHS, ConstantSDNode *RHS) {
1456 return LHS->getAPIntValue() == -RHS->getAPIntValue();
1457 });
1458 }
1459};
1460
1461/// Match a negation of a specific value V, either as sub(0, V) or as
1462/// constant(s) that are the negation of V's constant(s).
1463inline SpecificNeg_match m_SpecificNeg(SDValue V) {
1464 return SpecificNeg_match(V);
1465}
1466
1467template <typename... PatternTs> struct ReassociatableOpc_match {
1468 unsigned Opcode;
1469 std::tuple<PatternTs...> Patterns;
1470 constexpr static size_t NumPatterns =
1471 std::tuple_size_v<std::tuple<PatternTs...>>;
1472
1473 SDNodeFlags Flags;
1474
1475 ReassociatableOpc_match(unsigned Opcode, const PatternTs &...Patterns)
1476 : Opcode(Opcode), Patterns(Patterns...) {}
1477
1478 ReassociatableOpc_match(unsigned Opcode, SDNodeFlags Flags,
1479 const PatternTs &...Patterns)
1480 : Opcode(Opcode), Patterns(Patterns...), Flags(Flags) {}
1481
1482 bool match(SDValue N) {
1483 std::array<SDValue, NumPatterns> Leaves;
1484 size_t LeavesIdx = 0;
1485 if (!(collectLeaves(V: N, Leaves, LeafIdx&: LeavesIdx) && (LeavesIdx == NumPatterns)))
1486 return false;
1487
1488 Bitset<NumPatterns> Used;
1489 return std::apply(
1490 [&](auto &...P) -> bool {
1491 return reassociatableMatchHelper(Leaves, Used, P...);
1492 },
1493 Patterns);
1494 }
1495
1496 bool collectLeaves(SDValue V, std::array<SDValue, NumPatterns> &Leaves,
1497 std::size_t &LeafIdx) {
1498 if (V->getOpcode() == Opcode && (Flags & V->getFlags()) == Flags) {
1499 for (size_t I = 0, N = V->getNumOperands(); I < N; I++)
1500 if ((LeafIdx == NumPatterns) ||
1501 !collectLeaves(V: V->getOperand(Num: I), Leaves, LeafIdx))
1502 return false;
1503 } else {
1504 Leaves[LeafIdx] = V;
1505 LeafIdx++;
1506 }
1507 return true;
1508 }
1509
1510 // Searchs for a matching leaf for every sub-pattern.
1511 template <typename PatternHd, typename... PatternTl>
1512 [[nodiscard]] inline bool
1513 reassociatableMatchHelper(ArrayRef<SDValue> Leaves, Bitset<NumPatterns> &Used,
1514 PatternHd &HeadPattern,
1515 PatternTl &...TailPatterns) {
1516 for (size_t Match = 0, N = Used.size(); Match < N; Match++) {
1517 if (Used[Match] || !(sd_match(Leaves[Match], HeadPattern)))
1518 continue;
1519 Used.set(Match);
1520 if (reassociatableMatchHelper(Leaves, Used, TailPatterns...))
1521 return true;
1522 Used.reset(Match);
1523 }
1524 return false;
1525 }
1526
1527 [[nodiscard]] inline bool
1528 reassociatableMatchHelper(ArrayRef<SDValue> Leaves,
1529 Bitset<NumPatterns> &Used) {
1530 return true;
1531 }
1532};
1533
1534template <typename... PatternTs>
1535inline ReassociatableOpc_match<PatternTs...>
1536m_ReassociatableAdd(const PatternTs &...Patterns) {
1537 return ReassociatableOpc_match<PatternTs...>(ISD::ADD, Patterns...);
1538}
1539
1540template <typename... PatternTs>
1541inline ReassociatableOpc_match<PatternTs...>
1542m_ReassociatableOr(const PatternTs &...Patterns) {
1543 return ReassociatableOpc_match<PatternTs...>(ISD::OR, Patterns...);
1544}
1545
1546template <typename... PatternTs>
1547inline ReassociatableOpc_match<PatternTs...>
1548m_ReassociatableAnd(const PatternTs &...Patterns) {
1549 return ReassociatableOpc_match<PatternTs...>(ISD::AND, Patterns...);
1550}
1551
1552template <typename... PatternTs>
1553inline ReassociatableOpc_match<PatternTs...>
1554m_ReassociatableMul(const PatternTs &...Patterns) {
1555 return ReassociatableOpc_match<PatternTs...>(ISD::MUL, Patterns...);
1556}
1557
1558template <typename... PatternTs>
1559inline ReassociatableOpc_match<PatternTs...>
1560m_ReassociatableNSWAdd(const PatternTs &...Patterns) {
1561 return ReassociatableOpc_match<PatternTs...>(
1562 ISD::ADD, SDNodeFlags::NoSignedWrap, Patterns...);
1563}
1564
1565template <typename... PatternTs>
1566inline ReassociatableOpc_match<PatternTs...>
1567m_ReassociatableNUWAdd(const PatternTs &...Patterns) {
1568 return ReassociatableOpc_match<PatternTs...>(
1569 ISD::ADD, SDNodeFlags::NoUnsignedWrap, Patterns...);
1570}
1571
1572} // namespace SDPatternMatch
1573} // namespace llvm
1574#endif
1575