1//===-- lib/CodeGen/GlobalISel/GICombinerHelper.cpp -----------------------===//
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#include "llvm/CodeGen/GlobalISel/CombinerHelper.h"
9#include "llvm/ADT/APFloat.h"
10#include "llvm/ADT/STLExtras.h"
11#include "llvm/ADT/SetVector.h"
12#include "llvm/ADT/SmallBitVector.h"
13#include "llvm/Analysis/CmpInstAnalysis.h"
14#include "llvm/CodeGen/GlobalISel/GISelChangeObserver.h"
15#include "llvm/CodeGen/GlobalISel/GISelValueTracking.h"
16#include "llvm/CodeGen/GlobalISel/GenericMachineInstrs.h"
17#include "llvm/CodeGen/GlobalISel/LegalizerHelper.h"
18#include "llvm/CodeGen/GlobalISel/LegalizerInfo.h"
19#include "llvm/CodeGen/GlobalISel/MIPatternMatch.h"
20#include "llvm/CodeGen/GlobalISel/MachineIRBuilder.h"
21#include "llvm/CodeGen/GlobalISel/Utils.h"
22#include "llvm/CodeGen/LowLevelTypeUtils.h"
23#include "llvm/CodeGen/MachineBasicBlock.h"
24#include "llvm/CodeGen/MachineDominators.h"
25#include "llvm/CodeGen/MachineInstr.h"
26#include "llvm/CodeGen/MachineMemOperand.h"
27#include "llvm/CodeGen/MachineRegisterInfo.h"
28#include "llvm/CodeGen/Register.h"
29#include "llvm/CodeGen/RegisterBankInfo.h"
30#include "llvm/CodeGen/TargetInstrInfo.h"
31#include "llvm/CodeGen/TargetLowering.h"
32#include "llvm/CodeGen/TargetOpcodes.h"
33#include "llvm/IR/ConstantRange.h"
34#include "llvm/IR/DataLayout.h"
35#include "llvm/IR/InstrTypes.h"
36#include "llvm/Support/Casting.h"
37#include "llvm/Support/DivisionByConstantInfo.h"
38#include "llvm/Support/ErrorHandling.h"
39#include "llvm/Support/MathExtras.h"
40#include "llvm/Target/TargetMachine.h"
41#include <cmath>
42#include <optional>
43#include <tuple>
44
45#define DEBUG_TYPE "gi-combiner"
46
47using namespace llvm;
48using namespace MIPatternMatch;
49
50// Option to allow testing of the combiner while no targets know about indexed
51// addressing.
52static cl::opt<bool>
53 ForceLegalIndexing("force-legal-indexing", cl::Hidden, cl::init(Val: false),
54 cl::desc("Force all indexed operations to be "
55 "legal for the GlobalISel combiner"));
56
57CombinerHelper::CombinerHelper(GISelChangeObserver &Observer,
58 MachineIRBuilder &B, bool IsPreLegalize,
59 GISelValueTracking *VT,
60 MachineDominatorTree *MDT,
61 const LegalizerInfo *LI)
62 : Builder(B), MRI(Builder.getMF().getRegInfo()), Observer(Observer), VT(VT),
63 MDT(MDT), IsPreLegalize(IsPreLegalize), LI(LI),
64 TII(Builder.getMF().getSubtarget().getInstrInfo()),
65 RBI(Builder.getMF().getSubtarget().getRegBankInfo()),
66 TRI(Builder.getMF().getSubtarget().getRegisterInfo()) {
67 (void)this->VT;
68}
69
70const TargetLowering &CombinerHelper::getTargetLowering() const {
71 return *Builder.getMF().getSubtarget().getTargetLowering();
72}
73
74const MachineFunction &CombinerHelper::getMachineFunction() const {
75 return Builder.getMF();
76}
77
78LLVMContext &CombinerHelper::getContext() const { return Builder.getContext(); }
79
80/// \returns The little endian in-memory byte position of byte \p I in a
81/// \p ByteWidth bytes wide type.
82///
83/// E.g. Given a 4-byte type x, x[0] -> byte 0
84static unsigned littleEndianByteAt(const unsigned ByteWidth, const unsigned I) {
85 assert(I < ByteWidth && "I must be in [0, ByteWidth)");
86 return I;
87}
88
89/// Determines the LogBase2 value for a non-null input value using the
90/// transform: LogBase2(V) = (EltBits - 1) - ctlz(V).
91static Register buildLogBase2(Register V, MachineIRBuilder &MIB) {
92 auto &MRI = *MIB.getMRI();
93 LLT Ty = MRI.getType(Reg: V);
94 auto Ctlz = MIB.buildCTLZ(Dst: Ty, Src0: V);
95 auto Base = MIB.buildConstant(Res: Ty, Val: Ty.getScalarSizeInBits() - 1);
96 return MIB.buildSub(Dst: Ty, Src0: Base, Src1: Ctlz).getReg(Idx: 0);
97}
98
99/// \returns The big endian in-memory byte position of byte \p I in a
100/// \p ByteWidth bytes wide type.
101///
102/// E.g. Given a 4-byte type x, x[0] -> byte 3
103static unsigned bigEndianByteAt(const unsigned ByteWidth, const unsigned I) {
104 assert(I < ByteWidth && "I must be in [0, ByteWidth)");
105 return ByteWidth - I - 1;
106}
107
108/// Given a map from byte offsets in memory to indices in a load/store,
109/// determine if that map corresponds to a little or big endian byte pattern.
110///
111/// \param MemOffset2Idx maps memory offsets to address offsets.
112/// \param LowestIdx is the lowest index in \p MemOffset2Idx.
113///
114/// \returns true if the map corresponds to a big endian byte pattern, false if
115/// it corresponds to a little endian byte pattern, and std::nullopt otherwise.
116///
117/// E.g. given a 32-bit type x, and x[AddrOffset], the in-memory byte patterns
118/// are as follows:
119///
120/// AddrOffset Little endian Big endian
121/// 0 0 3
122/// 1 1 2
123/// 2 2 1
124/// 3 3 0
125static std::optional<bool>
126isBigEndian(const SmallDenseMap<int64_t, int64_t, 8> &MemOffset2Idx,
127 int64_t LowestIdx) {
128 // Need at least two byte positions to decide on endianness.
129 unsigned Width = MemOffset2Idx.size();
130 if (Width < 2)
131 return std::nullopt;
132 bool BigEndian = true, LittleEndian = true;
133 for (unsigned MemOffset = 0; MemOffset < Width; ++ MemOffset) {
134 auto MemOffsetAndIdx = MemOffset2Idx.find(Val: MemOffset);
135 if (MemOffsetAndIdx == MemOffset2Idx.end())
136 return std::nullopt;
137 const int64_t Idx = MemOffsetAndIdx->second - LowestIdx;
138 assert(Idx >= 0 && "Expected non-negative byte offset?");
139 LittleEndian &= Idx == littleEndianByteAt(ByteWidth: Width, I: MemOffset);
140 BigEndian &= Idx == bigEndianByteAt(ByteWidth: Width, I: MemOffset);
141 if (!BigEndian && !LittleEndian)
142 return std::nullopt;
143 }
144
145 assert((BigEndian != LittleEndian) &&
146 "Pattern cannot be both big and little endian!");
147 return BigEndian;
148}
149
150bool CombinerHelper::isPreLegalize() const { return IsPreLegalize; }
151
152bool CombinerHelper::isLegal(const LegalityQuery &Query) const {
153 assert(LI && "Must have LegalizerInfo to query isLegal!");
154 return LI->getAction(Query).Action == LegalizeActions::Legal;
155}
156
157bool CombinerHelper::isLegalOrBeforeLegalizer(
158 const LegalityQuery &Query) const {
159 return isPreLegalize() || isLegal(Query);
160}
161
162bool CombinerHelper::isLegalOrHasWidenScalar(const LegalityQuery &Query) const {
163 return isLegal(Query) ||
164 LI->getAction(Query).Action == LegalizeActions::WidenScalar;
165}
166
167bool CombinerHelper::isLegalOrHasFewerElements(
168 const LegalityQuery &Query) const {
169 LegalizeAction Action = LI->getAction(Query).Action;
170 return Action == LegalizeActions::Legal ||
171 Action == LegalizeActions::FewerElements;
172}
173
174bool CombinerHelper::isConstantLegalOrBeforeLegalizer(const LLT Ty) const {
175 if (!Ty.isVector())
176 return isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_CONSTANT, {Ty}});
177 // Vector constants are represented as a G_BUILD_VECTOR of scalar G_CONSTANTs.
178 if (isPreLegalize())
179 return true;
180 LLT EltTy = Ty.getElementType();
181 return isLegal(Query: {TargetOpcode::G_BUILD_VECTOR, {Ty, EltTy}}) &&
182 isLegal(Query: {TargetOpcode::G_CONSTANT, {EltTy}});
183}
184
185void CombinerHelper::replaceRegWith(MachineRegisterInfo &MRI, Register FromReg,
186 Register ToReg) const {
187 Observer.changingAllUsesOfReg(MRI, Reg: FromReg);
188
189 if (MRI.constrainRegAttrs(Reg: ToReg, ConstrainingReg: FromReg))
190 MRI.replaceRegWith(FromReg, ToReg);
191 else
192 Builder.buildCopy(Res: FromReg, Op: ToReg);
193
194 Observer.finishedChangingAllUsesOfReg();
195}
196
197void CombinerHelper::replaceRegOpWith(MachineRegisterInfo &MRI,
198 MachineOperand &FromRegOp,
199 Register ToReg) const {
200 assert(FromRegOp.getParent() && "Expected an operand in an MI");
201 Observer.changingInstr(MI&: *FromRegOp.getParent());
202
203 FromRegOp.setReg(ToReg);
204
205 Observer.changedInstr(MI&: *FromRegOp.getParent());
206}
207
208void CombinerHelper::replaceOpcodeWith(MachineInstr &FromMI,
209 unsigned ToOpcode) const {
210 Observer.changingInstr(MI&: FromMI);
211
212 FromMI.setDesc(Builder.getTII().get(Opcode: ToOpcode));
213
214 Observer.changedInstr(MI&: FromMI);
215}
216
217const RegisterBank *CombinerHelper::getRegBank(Register Reg) const {
218 return RBI->getRegBank(Reg, MRI, TRI: *TRI);
219}
220
221void CombinerHelper::setRegBank(Register Reg,
222 const RegisterBank *RegBank) const {
223 if (RegBank)
224 MRI.setRegBank(Reg, RegBank: *RegBank);
225}
226
227bool CombinerHelper::matchCombineCopy(MachineInstr &MI) const {
228 if (MI.getOpcode() != TargetOpcode::COPY)
229 return false;
230 Register DstReg = MI.getOperand(i: 0).getReg();
231 Register SrcReg = MI.getOperand(i: 1).getReg();
232 return canReplaceReg(DstReg, SrcReg, MRI);
233}
234void CombinerHelper::applyCombineCopy(MachineInstr &MI) const {
235 Register DstReg = MI.getOperand(i: 0).getReg();
236 Register SrcReg = MI.getOperand(i: 1).getReg();
237 replaceRegWith(MRI, FromReg: DstReg, ToReg: SrcReg);
238 MI.eraseFromParent();
239}
240
241bool CombinerHelper::matchFreezeOfSingleMaybePoisonOperand(
242 MachineInstr &MI, BuildFnTy &MatchInfo) const {
243 assert(MI.getOpcode() == TargetOpcode::G_FREEZE && "Invalid instruction");
244
245 // Ported from InstCombinerImpl::pushFreezeToPreventPoisonFromPropagating.
246 Register DstOp = MI.getOperand(i: 0).getReg();
247 Register OrigOp = MI.getOperand(i: 1).getReg();
248
249 if (!MRI.hasOneNonDBGUse(RegNo: OrigOp))
250 return false;
251
252 MachineInstr *OrigDef;
253 if (!mi_match(R: OrigOp, MRI, P: m_MInstr(MI&: OrigDef)))
254 return false;
255 // Even if only a single operand of the PHI is not guaranteed non-poison,
256 // moving freeze() backwards across a PHI can cause optimization issues for
257 // other users of that operand.
258 //
259 // Moving freeze() from one of the output registers of a G_UNMERGE_VALUES to
260 // the source register is unprofitable because it makes the freeze() more
261 // strict than is necessary (it would affect the whole register instead of
262 // just the subreg being frozen).
263 if (OrigDef->isPHI() || isa<GUnmerge>(Val: OrigDef))
264 return false;
265
266 if (canCreateUndefOrPoison(Reg: OrigOp, MRI,
267 /*ConsiderFlagsAndMetadata=*/false))
268 return false;
269
270 std::optional<MachineOperand> MaybePoisonOperand;
271 for (MachineOperand &Operand : OrigDef->uses()) {
272 if (!Operand.isReg())
273 return false;
274
275 if (isGuaranteedNotToBeUndefOrPoison(Reg: Operand.getReg(), MRI))
276 continue;
277
278 if (!MaybePoisonOperand)
279 MaybePoisonOperand = Operand;
280 else {
281 // We have more than one maybe-poison operand. Moving the freeze is
282 // unsafe.
283 return false;
284 }
285 }
286
287 // Eliminate freeze if all operands are guaranteed non-poison.
288 if (!MaybePoisonOperand) {
289 MatchInfo = [=](MachineIRBuilder &B) {
290 Observer.changingInstr(MI&: *OrigDef);
291 cast<GenericMachineInstr>(Val: OrigDef)->dropPoisonGeneratingFlags();
292 Observer.changedInstr(MI&: *OrigDef);
293 B.buildCopy(Res: DstOp, Op: OrigOp);
294 };
295 return true;
296 }
297
298 Register MaybePoisonOperandReg = MaybePoisonOperand->getReg();
299 LLT MaybePoisonOperandRegTy = MRI.getType(Reg: MaybePoisonOperandReg);
300
301 if (!isLegalOrBeforeLegalizer(
302 Query: {TargetOpcode::G_FREEZE, {MaybePoisonOperandRegTy}}))
303 return false;
304
305 MatchInfo = [=](MachineIRBuilder &B) mutable {
306 Observer.changingInstr(MI&: *OrigDef);
307 cast<GenericMachineInstr>(Val: OrigDef)->dropPoisonGeneratingFlags();
308 Observer.changedInstr(MI&: *OrigDef);
309 B.setInsertPt(MBB&: *OrigDef->getParent(), II: OrigDef->getIterator());
310 auto Freeze = B.buildFreeze(Dst: MaybePoisonOperandRegTy, Src: MaybePoisonOperandReg);
311 replaceRegOpWith(
312 MRI, FromRegOp&: *OrigDef->findRegisterUseOperand(Reg: MaybePoisonOperandReg, TRI),
313 ToReg: Freeze.getReg(Idx: 0));
314 replaceRegWith(MRI, FromReg: DstOp, ToReg: OrigOp);
315 };
316 return true;
317}
318
319bool CombinerHelper::matchCombineConcatVectors(
320 MachineInstr &MI, SmallVector<Register> &Ops) const {
321 assert(MI.getOpcode() == TargetOpcode::G_CONCAT_VECTORS &&
322 "Invalid instruction");
323 bool IsUndef = true;
324 MachineInstr *Undef = nullptr;
325
326 // Walk over all the operands of concat vectors and check if they are
327 // build_vector themselves or undef.
328 // Then collect their operands in Ops.
329 for (const MachineOperand &MO : MI.uses()) {
330 Register Reg = MO.getReg();
331 MachineInstr *Def;
332 if (!mi_match(R: Reg, MRI, P: m_MInstr(MI&: Def)))
333 return false;
334 if (!MRI.hasOneNonDBGUse(RegNo: Reg))
335 return false;
336 switch (Def->getOpcode()) {
337 case TargetOpcode::G_BUILD_VECTOR:
338 IsUndef = false;
339 // Remember the operands of the build_vector to fold
340 // them into the yet-to-build flattened concat vectors.
341 for (const MachineOperand &BuildVecMO : Def->uses())
342 Ops.push_back(Elt: BuildVecMO.getReg());
343 break;
344 case TargetOpcode::G_IMPLICIT_DEF: {
345 LLT OpType = MRI.getType(Reg);
346 // Keep one undef value for all the undef operands.
347 if (!Undef) {
348 Builder.setInsertPt(MBB&: *MI.getParent(), II: MI);
349 Undef = Builder.buildUndef(Res: OpType.getScalarType());
350 }
351 assert(MRI.getType(Undef->getOperand(0).getReg()) ==
352 OpType.getScalarType() &&
353 "All undefs should have the same type");
354 // Break the undef vector in as many scalar elements as needed
355 // for the flattening.
356 for (unsigned EltIdx = 0, EltEnd = OpType.getNumElements();
357 EltIdx != EltEnd; ++EltIdx)
358 Ops.push_back(Elt: Undef->getOperand(i: 0).getReg());
359 break;
360 }
361 default:
362 return false;
363 }
364 }
365
366 // Check if the combine is illegal
367 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
368 if (!isLegalOrBeforeLegalizer(
369 Query: {TargetOpcode::G_BUILD_VECTOR, {DstTy, MRI.getType(Reg: Ops[0])}})) {
370 return false;
371 }
372
373 if (IsUndef)
374 Ops.clear();
375
376 return true;
377}
378void CombinerHelper::applyCombineConcatVectors(
379 MachineInstr &MI, SmallVector<Register> &Ops) const {
380 // We determined that the concat_vectors can be flatten.
381 // Generate the flattened build_vector.
382 Register DstReg = MI.getOperand(i: 0).getReg();
383 Builder.setInsertPt(MBB&: *MI.getParent(), II: MI);
384 Register NewDstReg = MRI.cloneVirtualRegister(VReg: DstReg);
385
386 // Note: IsUndef is sort of redundant. We could have determine it by
387 // checking that at all Ops are undef. Alternatively, we could have
388 // generate a build_vector of undefs and rely on another combine to
389 // clean that up. For now, given we already gather this information
390 // in matchCombineConcatVectors, just save compile time and issue the
391 // right thing.
392 if (Ops.empty())
393 Builder.buildUndef(Res: NewDstReg);
394 else
395 Builder.buildBuildVector(Res: NewDstReg, Ops);
396 replaceRegWith(MRI, FromReg: DstReg, ToReg: NewDstReg);
397 MI.eraseFromParent();
398}
399
400bool CombinerHelper::matchCombineBuildVectorOfBitcast(
401 MachineInstr &MI, SmallVector<Register> &Ops) const {
402 auto &BV = cast<GBuildVector>(Val&: MI);
403
404 // Look at the first operand for a unmerge(bitcast) from a scalar type.
405 GUnmerge *Unmerge = getOpcodeDef<GUnmerge>(Reg: BV.getSourceReg(I: 0), MRI);
406 if (!Unmerge || Unmerge->getReg(Idx: 0) != BV.getSourceReg(I: 0))
407 return false;
408 Register BCSrc;
409 if (!mi_match(R: Unmerge->getSourceReg(), MRI, P: m_GBitcast(Src: m_Reg(R&: BCSrc))))
410 return false;
411 LLT InputTy = MRI.getType(Reg: BCSrc);
412 unsigned Factor = Unmerge->getNumDefs();
413 if (!InputTy.isScalar() || BV.getNumSources() % Factor != 0)
414 return false;
415
416 // Check if the build_vector is legal
417 LLT BVDstTy = LLT::fixed_vector(NumElements: BV.getNumSources() / Factor, ScalarTy: InputTy);
418 if (!isLegal(Query: {TargetOpcode::G_BUILD_VECTOR, {BVDstTy, InputTy}}))
419 return false;
420
421 // Check all other operands are bitcasts or undef.
422 for (unsigned Idx = 0; Idx < BV.getNumSources(); Idx += Factor) {
423 GUnmerge *Unmerge = getOpcodeDef<GUnmerge>(Reg: BV.getSourceReg(I: Idx), MRI);
424 if (!all_of(Range: iota_range<unsigned>(0, Factor, false), P: [&](unsigned J) {
425 if (mi_match(R: BV.getSourceReg(I: Idx + J), MRI, P: m_GImplicitDef()))
426 return true;
427 return Unmerge && BV.getSourceReg(I: Idx + J) == Unmerge->getReg(Idx: J);
428 }))
429 return false;
430 if (!Unmerge)
431 Ops.push_back(Elt: 0);
432 else {
433 Register BCSrc;
434 if (!mi_match(
435 R: Unmerge->getSourceReg(), MRI,
436 P: m_GBitcast(Src: m_all_of(preds: m_Reg(R&: BCSrc), preds: m_SpecificType(Ty: InputTy)))))
437 return false;
438 Ops.push_back(Elt: BCSrc);
439 }
440 }
441
442 return true;
443}
444
445void CombinerHelper::applyCombineBuildVectorOfBitcast(
446 MachineInstr &MI, SmallVector<Register> &Ops) const {
447 LLT SrcTy = MRI.getType(Reg: Ops[0]);
448 // Build undef if any operations require it.
449 Register Undef = 0;
450 for (Register &Op : Ops) {
451 if (!Op) {
452 if (!Undef)
453 Undef = Builder.buildUndef(Res: SrcTy).getReg(Idx: 0);
454 Op = Undef;
455 }
456 }
457
458 LLT BVDstTy = LLT::fixed_vector(NumElements: Ops.size(), ScalarTy: SrcTy);
459 auto BV = Builder.buildBuildVector(Res: BVDstTy, Ops);
460 Builder.buildBitcast(Dst: MI.getOperand(i: 0).getReg(), Src: BV);
461 MI.eraseFromParent();
462}
463
464void CombinerHelper::applyCombineShuffleToBuildVector(MachineInstr &MI) const {
465 auto &Shuffle = cast<GShuffleVector>(Val&: MI);
466
467 Register SrcVec1 = Shuffle.getSrc1Reg();
468 Register SrcVec2 = Shuffle.getSrc2Reg();
469 LLT EltTy = MRI.getType(Reg: SrcVec1).getElementType();
470 int Width = MRI.getType(Reg: SrcVec1).getNumElements();
471
472 auto Unmerge1 = Builder.buildUnmerge(Res: EltTy, Op: SrcVec1);
473 auto Unmerge2 = Builder.buildUnmerge(Res: EltTy, Op: SrcVec2);
474
475 SmallVector<Register> Extracts;
476 // Select only applicable elements from unmerged values.
477 for (int Val : Shuffle.getMask()) {
478 if (Val == -1)
479 Extracts.push_back(Elt: Builder.buildUndef(Res: EltTy).getReg(Idx: 0));
480 else if (Val < Width)
481 Extracts.push_back(Elt: Unmerge1.getReg(Idx: Val));
482 else
483 Extracts.push_back(Elt: Unmerge2.getReg(Idx: Val - Width));
484 }
485 assert(Extracts.size() > 0 && "Expected at least one element in the shuffle");
486 if (Extracts.size() == 1)
487 Builder.buildCopy(Res: MI.getOperand(i: 0).getReg(), Op: Extracts[0]);
488 else
489 Builder.buildBuildVector(Res: MI.getOperand(i: 0).getReg(), Ops: Extracts);
490 MI.eraseFromParent();
491}
492
493bool CombinerHelper::matchCombineShuffleConcat(
494 MachineInstr &MI, SmallVector<Register> &Ops) const {
495 ArrayRef<int> Mask = MI.getOperand(i: 3).getShuffleMask();
496 GConcatVectors *ConcatMI1, *ConcatMI2;
497 if (!mi_match(R: MI.getOperand(i: 1).getReg(), MRI, P: m_GConcatVectors(Inst&: ConcatMI1)) ||
498 !mi_match(R: MI.getOperand(i: 2).getReg(), MRI, P: m_GConcatVectors(Inst&: ConcatMI2)))
499 return false;
500
501 // Check that the sources of the Concat instructions have the same type
502 if (MRI.getType(Reg: ConcatMI1->getSourceReg(I: 0)) !=
503 MRI.getType(Reg: ConcatMI2->getSourceReg(I: 0)))
504 return false;
505
506 LLT ConcatSrcTy = MRI.getType(Reg: ConcatMI1->getReg(Idx: 1));
507 LLT ShuffleSrcTy1 = MRI.getType(Reg: MI.getOperand(i: 1).getReg());
508 unsigned ConcatSrcNumElt = ConcatSrcTy.getNumElements();
509 for (unsigned i = 0; i < Mask.size(); i += ConcatSrcNumElt) {
510 // Check if the index takes a whole source register from G_CONCAT_VECTORS
511 // Assumes that all Sources of G_CONCAT_VECTORS are the same type
512 if (Mask[i] == -1) {
513 for (unsigned j = 1; j < ConcatSrcNumElt; j++) {
514 if (i + j >= Mask.size())
515 return false;
516 if (Mask[i + j] != -1)
517 return false;
518 }
519 if (!isLegalOrBeforeLegalizer(
520 Query: {TargetOpcode::G_IMPLICIT_DEF, {ConcatSrcTy}}))
521 return false;
522 Ops.push_back(Elt: 0);
523 } else if (Mask[i] % ConcatSrcNumElt == 0) {
524 for (unsigned j = 1; j < ConcatSrcNumElt; j++) {
525 if (i + j >= Mask.size())
526 return false;
527 if (Mask[i + j] != Mask[i] + static_cast<int>(j))
528 return false;
529 }
530 // Retrieve the source register from its respective G_CONCAT_VECTORS
531 // instruction
532 if (Mask[i] < ShuffleSrcTy1.getNumElements()) {
533 Ops.push_back(Elt: ConcatMI1->getSourceReg(I: Mask[i] / ConcatSrcNumElt));
534 } else {
535 Ops.push_back(Elt: ConcatMI2->getSourceReg(I: Mask[i] / ConcatSrcNumElt -
536 ConcatMI1->getNumSources()));
537 }
538 } else {
539 return false;
540 }
541 }
542
543 if (!isLegalOrBeforeLegalizer(
544 Query: {TargetOpcode::G_CONCAT_VECTORS,
545 {MRI.getType(Reg: MI.getOperand(i: 0).getReg()), ConcatSrcTy}}))
546 return false;
547
548 return !Ops.empty();
549}
550
551void CombinerHelper::applyCombineShuffleConcat(
552 MachineInstr &MI, SmallVector<Register> &Ops) const {
553 LLT SrcTy;
554 for (Register &Reg : Ops) {
555 if (Reg != 0)
556 SrcTy = MRI.getType(Reg);
557 }
558 assert(SrcTy.isValid() && "Unexpected full undef vector in concat combine");
559
560 Register UndefReg = 0;
561
562 for (Register &Reg : Ops) {
563 if (Reg == 0) {
564 if (UndefReg == 0)
565 UndefReg = Builder.buildUndef(Res: SrcTy).getReg(Idx: 0);
566 Reg = UndefReg;
567 }
568 }
569
570 if (Ops.size() > 1)
571 Builder.buildConcatVectors(Res: MI.getOperand(i: 0).getReg(), Ops);
572 else
573 Builder.buildCopy(Res: MI.getOperand(i: 0).getReg(), Op: Ops[0]);
574 MI.eraseFromParent();
575}
576
577bool CombinerHelper::matchCombineShuffleVector(
578 MachineInstr &MI, SmallVectorImpl<Register> &Ops) const {
579 assert(MI.getOpcode() == TargetOpcode::G_SHUFFLE_VECTOR &&
580 "Invalid instruction kind");
581 LLT DstType = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
582 Register Src1 = MI.getOperand(i: 1).getReg();
583 LLT SrcType = MRI.getType(Reg: Src1);
584
585 unsigned DstNumElts = DstType.getNumElements();
586 unsigned SrcNumElts = SrcType.getNumElements();
587
588 // If the resulting vector is smaller than the size of the source
589 // vectors being concatenated, we won't be able to replace the
590 // shuffle vector into a concat_vectors.
591 //
592 // Note: We may still be able to produce a concat_vectors fed by
593 // extract_vector_elt and so on. It is less clear that would
594 // be better though, so don't bother for now.
595 //
596 // If the destination is a scalar, the size of the sources doesn't
597 // matter. we will lower the shuffle to a plain copy. This will
598 // work only if the source and destination have the same size. But
599 // that's covered by the next condition.
600 //
601 // TODO: If the size between the source and destination don't match
602 // we could still emit an extract vector element in that case.
603 if (DstNumElts < 2 * SrcNumElts)
604 return false;
605
606 // Check that the shuffle mask can be broken evenly between the
607 // different sources.
608 if (DstNumElts % SrcNumElts != 0)
609 return false;
610
611 // Mask length is a multiple of the source vector length.
612 // Check if the shuffle is some kind of concatenation of the input
613 // vectors.
614 unsigned NumConcat = DstNumElts / SrcNumElts;
615 SmallVector<int, 8> ConcatSrcs(NumConcat, -1);
616 ArrayRef<int> Mask = MI.getOperand(i: 3).getShuffleMask();
617 for (unsigned i = 0; i != DstNumElts; ++i) {
618 int Idx = Mask[i];
619 // Undef value.
620 if (Idx < 0)
621 continue;
622 // Ensure the indices in each SrcType sized piece are sequential and that
623 // the same source is used for the whole piece.
624 if ((Idx % SrcNumElts != (i % SrcNumElts)) ||
625 (ConcatSrcs[i / SrcNumElts] >= 0 &&
626 ConcatSrcs[i / SrcNumElts] != (int)(Idx / SrcNumElts)))
627 return false;
628 // Remember which source this index came from.
629 ConcatSrcs[i / SrcNumElts] = Idx / SrcNumElts;
630 }
631
632 // The shuffle is concatenating multiple vectors together.
633 // Collect the different operands for that.
634 Register UndefReg;
635 Register Src2 = MI.getOperand(i: 2).getReg();
636 for (auto Src : ConcatSrcs) {
637 if (Src < 0) {
638 if (!UndefReg) {
639 Builder.setInsertPt(MBB&: *MI.getParent(), II: MI);
640 UndefReg = Builder.buildUndef(Res: SrcType).getReg(Idx: 0);
641 }
642 Ops.push_back(Elt: UndefReg);
643 } else if (Src == 0)
644 Ops.push_back(Elt: Src1);
645 else
646 Ops.push_back(Elt: Src2);
647 }
648 return true;
649}
650
651void CombinerHelper::applyCombineShuffleVector(MachineInstr &MI,
652 ArrayRef<Register> Ops) const {
653 Register DstReg = MI.getOperand(i: 0).getReg();
654 Builder.setInsertPt(MBB&: *MI.getParent(), II: MI);
655 Register NewDstReg = MRI.cloneVirtualRegister(VReg: DstReg);
656
657 if (Ops.size() == 1)
658 Builder.buildCopy(Res: NewDstReg, Op: Ops[0]);
659 else
660 Builder.buildMergeLikeInstr(Res: NewDstReg, Ops);
661
662 replaceRegWith(MRI, FromReg: DstReg, ToReg: NewDstReg);
663 MI.eraseFromParent();
664}
665
666namespace {
667
668/// Select a preference between two uses. CurrentUse is the current preference
669/// while *ForCandidate is attributes of the candidate under consideration.
670PreferredTuple ChoosePreferredUse(MachineInstr &LoadMI,
671 PreferredTuple &CurrentUse,
672 const LLT TyForCandidate,
673 unsigned OpcodeForCandidate,
674 MachineInstr *MIForCandidate) {
675 if (!CurrentUse.Ty.isValid()) {
676 if (CurrentUse.ExtendOpcode == OpcodeForCandidate ||
677 CurrentUse.ExtendOpcode == TargetOpcode::G_ANYEXT)
678 return {.Ty: TyForCandidate, .ExtendOpcode: OpcodeForCandidate, .MI: MIForCandidate};
679 return CurrentUse;
680 }
681
682 // We permit the extend to hoist through basic blocks but this is only
683 // sensible if the target has extending loads. If you end up lowering back
684 // into a load and extend during the legalizer then the end result is
685 // hoisting the extend up to the load.
686
687 // Prefer defined extensions to undefined extensions as these are more
688 // likely to reduce the number of instructions.
689 if (OpcodeForCandidate == TargetOpcode::G_ANYEXT &&
690 CurrentUse.ExtendOpcode != TargetOpcode::G_ANYEXT)
691 return CurrentUse;
692 else if (CurrentUse.ExtendOpcode == TargetOpcode::G_ANYEXT &&
693 OpcodeForCandidate != TargetOpcode::G_ANYEXT)
694 return {.Ty: TyForCandidate, .ExtendOpcode: OpcodeForCandidate, .MI: MIForCandidate};
695
696 // Prefer sign extensions to zero extensions as sign-extensions tend to be
697 // more expensive. Don't do this if the load is already a zero-extend load
698 // though, otherwise we'll rewrite a zero-extend load into a sign-extend
699 // later.
700 if (!isa<GZExtLoad>(Val: LoadMI) && CurrentUse.Ty == TyForCandidate) {
701 if (CurrentUse.ExtendOpcode == TargetOpcode::G_SEXT &&
702 OpcodeForCandidate == TargetOpcode::G_ZEXT)
703 return CurrentUse;
704 else if (CurrentUse.ExtendOpcode == TargetOpcode::G_ZEXT &&
705 OpcodeForCandidate == TargetOpcode::G_SEXT)
706 return {.Ty: TyForCandidate, .ExtendOpcode: OpcodeForCandidate, .MI: MIForCandidate};
707 }
708
709 // This is potentially target specific. We've chosen the largest type
710 // because G_TRUNC is usually free. One potential catch with this is that
711 // some targets have a reduced number of larger registers than smaller
712 // registers and this choice potentially increases the live-range for the
713 // larger value.
714 if (TyForCandidate.getSizeInBits() > CurrentUse.Ty.getSizeInBits()) {
715 return {.Ty: TyForCandidate, .ExtendOpcode: OpcodeForCandidate, .MI: MIForCandidate};
716 }
717 return CurrentUse;
718}
719
720/// Find a suitable place to insert some instructions and insert them. This
721/// function accounts for special cases like inserting before a PHI node.
722/// The current strategy for inserting before PHI's is to duplicate the
723/// instructions for each predecessor. However, while that's ok for G_TRUNC
724/// on most targets since it generally requires no code, other targets/cases may
725/// want to try harder to find a dominating block.
726static void InsertInsnsWithoutSideEffectsBeforeUse(
727 MachineIRBuilder &Builder, MachineInstr &DefMI, MachineOperand &UseMO,
728 std::function<void(MachineBasicBlock *, MachineBasicBlock::iterator,
729 MachineOperand &UseMO)>
730 Inserter) {
731 MachineInstr &UseMI = *UseMO.getParent();
732
733 MachineBasicBlock *InsertBB = UseMI.getParent();
734
735 // If the use is a PHI then we want the predecessor block instead.
736 if (UseMI.isPHI()) {
737 MachineOperand *PredBB = std::next(x: &UseMO);
738 InsertBB = PredBB->getMBB();
739 }
740
741 // If the block is the same block as the def then we want to insert just after
742 // the def instead of at the start of the block.
743 if (InsertBB == DefMI.getParent()) {
744 MachineBasicBlock::iterator InsertPt = &DefMI;
745 Inserter(InsertBB, std::next(x: InsertPt), UseMO);
746 return;
747 }
748
749 // Otherwise we want the start of the BB
750 Inserter(InsertBB, InsertBB->getFirstNonPHI(), UseMO);
751}
752} // end anonymous namespace
753
754bool CombinerHelper::tryCombineExtendingLoads(MachineInstr &MI) const {
755 PreferredTuple Preferred;
756 if (matchCombineExtendingLoads(MI, MatchInfo&: Preferred)) {
757 applyCombineExtendingLoads(MI, MatchInfo&: Preferred);
758 return true;
759 }
760 return false;
761}
762
763static unsigned getExtLoadOpcForExtend(unsigned ExtOpc) {
764 unsigned CandidateLoadOpc;
765 switch (ExtOpc) {
766 case TargetOpcode::G_ANYEXT:
767 CandidateLoadOpc = TargetOpcode::G_LOAD;
768 break;
769 case TargetOpcode::G_SEXT:
770 CandidateLoadOpc = TargetOpcode::G_SEXTLOAD;
771 break;
772 case TargetOpcode::G_ZEXT:
773 CandidateLoadOpc = TargetOpcode::G_ZEXTLOAD;
774 break;
775 default:
776 llvm_unreachable("Unexpected extend opc");
777 }
778 return CandidateLoadOpc;
779}
780
781bool CombinerHelper::matchCombineExtendingLoads(
782 MachineInstr &MI, PreferredTuple &Preferred) const {
783 // We match the loads and follow the uses to the extend instead of matching
784 // the extends and following the def to the load. This is because the load
785 // must remain in the same position for correctness (unless we also add code
786 // to find a safe place to sink it) whereas the extend is freely movable.
787 // It also prevents us from duplicating the load for the volatile case or just
788 // for performance.
789 GAnyLoad *LoadMI = dyn_cast<GAnyLoad>(Val: &MI);
790 if (!LoadMI)
791 return false;
792
793 Register LoadReg = LoadMI->getDstReg();
794
795 LLT LoadValueTy = MRI.getType(Reg: LoadReg);
796 if (!LoadValueTy.isScalar())
797 return false;
798
799 // Most architectures are going to legalize <s8 loads into at least a 1 byte
800 // load, and the MMOs can only describe memory accesses in multiples of bytes.
801 // If we try to perform extload combining on those, we can end up with
802 // %a(s8) = extload %ptr (load 1 byte from %ptr)
803 // ... which is an illegal extload instruction.
804 if (LoadValueTy.getSizeInBits() < 8)
805 return false;
806
807 // For non power-of-2 types, they will very likely be legalized into multiple
808 // loads. Don't bother trying to match them into extending loads.
809 if (!llvm::has_single_bit<uint32_t>(Value: LoadValueTy.getSizeInBits()))
810 return false;
811
812 // Find the preferred type aside from the any-extends (unless it's the only
813 // one) and non-extending ops. We'll emit an extending load to that type and
814 // and emit a variant of (extend (trunc X)) for the others according to the
815 // relative type sizes. At the same time, pick an extend to use based on the
816 // extend involved in the chosen type.
817 unsigned PreferredOpcode =
818 isa<GLoad>(Val: &MI)
819 ? TargetOpcode::G_ANYEXT
820 : isa<GSExtLoad>(Val: &MI) ? TargetOpcode::G_SEXT : TargetOpcode::G_ZEXT;
821 Preferred = {.Ty: LLT(), .ExtendOpcode: PreferredOpcode, .MI: nullptr};
822 for (auto &UseMI : MRI.use_nodbg_instructions(Reg: LoadReg)) {
823 if (UseMI.getOpcode() == TargetOpcode::G_SEXT ||
824 UseMI.getOpcode() == TargetOpcode::G_ZEXT ||
825 (UseMI.getOpcode() == TargetOpcode::G_ANYEXT)) {
826 const auto &MMO = LoadMI->getMMO();
827 // Don't do anything for atomics.
828 if (MMO.isAtomic())
829 continue;
830 // Check for legality.
831 if (!isPreLegalize()) {
832 LegalityQuery::MemDesc MMDesc(MMO);
833 unsigned CandidateLoadOpc = getExtLoadOpcForExtend(ExtOpc: UseMI.getOpcode());
834 LLT UseTy = MRI.getType(Reg: UseMI.getOperand(i: 0).getReg());
835 LLT SrcTy = MRI.getType(Reg: LoadMI->getPointerReg());
836 if (LI->getAction(Query: {CandidateLoadOpc, {UseTy, SrcTy}, {MMDesc}})
837 .Action != LegalizeActions::Legal)
838 continue;
839 }
840 Preferred = ChoosePreferredUse(LoadMI&: MI, CurrentUse&: Preferred,
841 TyForCandidate: MRI.getType(Reg: UseMI.getOperand(i: 0).getReg()),
842 OpcodeForCandidate: UseMI.getOpcode(), MIForCandidate: &UseMI);
843 }
844 }
845
846 // There were no extends
847 if (!Preferred.MI)
848 return false;
849 // It should be impossible to chose an extend without selecting a different
850 // type since by definition the result of an extend is larger.
851 assert(Preferred.Ty != LoadValueTy && "Extending to same type?");
852
853 LLVM_DEBUG(dbgs() << "Preferred use is: " << *Preferred.MI);
854 return true;
855}
856
857void CombinerHelper::applyCombineExtendingLoads(
858 MachineInstr &MI, PreferredTuple &Preferred) const {
859 // Rewrite the load to the chosen extending load.
860 Register ChosenDstReg = Preferred.MI->getOperand(i: 0).getReg();
861
862 // Inserter to insert a truncate back to the original type at a given point
863 // with some basic CSE to limit truncate duplication to one per BB.
864 DenseMap<MachineBasicBlock *, MachineInstr *> EmittedInsns;
865 auto InsertTruncAt = [&](MachineBasicBlock *InsertIntoBB,
866 MachineBasicBlock::iterator InsertBefore,
867 MachineOperand &UseMO) {
868 MachineInstr *PreviouslyEmitted = EmittedInsns.lookup(Val: InsertIntoBB);
869 if (PreviouslyEmitted) {
870 Observer.changingInstr(MI&: *UseMO.getParent());
871 UseMO.setReg(PreviouslyEmitted->getOperand(i: 0).getReg());
872 Observer.changedInstr(MI&: *UseMO.getParent());
873 return;
874 }
875
876 Builder.setInsertPt(MBB&: *InsertIntoBB, II: InsertBefore);
877 Register NewDstReg = MRI.cloneVirtualRegister(VReg: MI.getOperand(i: 0).getReg());
878 MachineInstr *NewMI = Builder.buildTrunc(Res: NewDstReg, Op: ChosenDstReg);
879 EmittedInsns[InsertIntoBB] = NewMI;
880 replaceRegOpWith(MRI, FromRegOp&: UseMO, ToReg: NewDstReg);
881 };
882
883 Observer.changingInstr(MI);
884 unsigned LoadOpc = getExtLoadOpcForExtend(ExtOpc: Preferred.ExtendOpcode);
885 MI.setDesc(Builder.getTII().get(Opcode: LoadOpc));
886
887 // Rewrite all the uses to fix up the types.
888 auto &LoadValue = MI.getOperand(i: 0);
889 SmallVector<MachineOperand *, 4> Uses(
890 llvm::make_pointer_range(Range: MRI.use_operands(Reg: LoadValue.getReg())));
891
892 for (auto *UseMO : Uses) {
893 MachineInstr *UseMI = UseMO->getParent();
894
895 // If the extend is compatible with the preferred extend then we should fix
896 // up the type and extend so that it uses the preferred use.
897 if (UseMI->getOpcode() == Preferred.ExtendOpcode ||
898 UseMI->getOpcode() == TargetOpcode::G_ANYEXT) {
899 Register UseDstReg = UseMI->getOperand(i: 0).getReg();
900 MachineOperand &UseSrcMO = UseMI->getOperand(i: 1);
901 const LLT UseDstTy = MRI.getType(Reg: UseDstReg);
902 if (UseDstReg != ChosenDstReg) {
903 if (Preferred.Ty == UseDstTy) {
904 // If the use has the same type as the preferred use, then merge
905 // the vregs and erase the extend. For example:
906 // %1:_(s8) = G_LOAD ...
907 // %2:_(s32) = G_SEXT %1(s8)
908 // %3:_(s32) = G_ANYEXT %1(s8)
909 // ... = ... %3(s32)
910 // rewrites to:
911 // %2:_(s32) = G_SEXTLOAD ...
912 // ... = ... %2(s32)
913 replaceRegWith(MRI, FromReg: UseDstReg, ToReg: ChosenDstReg);
914 Observer.erasingInstr(MI&: *UseMO->getParent());
915 UseMO->getParent()->eraseFromParent();
916 } else if (Preferred.Ty.getSizeInBits() < UseDstTy.getSizeInBits()) {
917 // If the preferred size is smaller, then keep the extend but extend
918 // from the result of the extending load. For example:
919 // %1:_(s8) = G_LOAD ...
920 // %2:_(s32) = G_SEXT %1(s8)
921 // %3:_(s64) = G_ANYEXT %1(s8)
922 // ... = ... %3(s64)
923 /// rewrites to:
924 // %2:_(s32) = G_SEXTLOAD ...
925 // %3:_(s64) = G_ANYEXT %2:_(s32)
926 // ... = ... %3(s64)
927 replaceRegOpWith(MRI, FromRegOp&: UseSrcMO, ToReg: ChosenDstReg);
928 } else {
929 // If the preferred size is large, then insert a truncate. For
930 // example:
931 // %1:_(s8) = G_LOAD ...
932 // %2:_(s64) = G_SEXT %1(s8)
933 // %3:_(s32) = G_ZEXT %1(s8)
934 // ... = ... %3(s32)
935 /// rewrites to:
936 // %2:_(s64) = G_SEXTLOAD ...
937 // %4:_(s8) = G_TRUNC %2:_(s32)
938 // %3:_(s64) = G_ZEXT %2:_(s8)
939 // ... = ... %3(s64)
940 InsertInsnsWithoutSideEffectsBeforeUse(Builder, DefMI&: MI, UseMO&: *UseMO,
941 Inserter: InsertTruncAt);
942 }
943 continue;
944 }
945 // The use is (one of) the uses of the preferred use we chose earlier.
946 // We're going to update the load to def this value later so just erase
947 // the old extend.
948 Observer.erasingInstr(MI&: *UseMO->getParent());
949 UseMO->getParent()->eraseFromParent();
950 continue;
951 }
952
953 // The use isn't an extend. Truncate back to the type we originally loaded.
954 // This is free on many targets.
955 InsertInsnsWithoutSideEffectsBeforeUse(Builder, DefMI&: MI, UseMO&: *UseMO, Inserter: InsertTruncAt);
956 }
957
958 MI.getOperand(i: 0).setReg(ChosenDstReg);
959 Observer.changedInstr(MI);
960}
961
962bool CombinerHelper::matchCombineLoadWithAndMask(MachineInstr &MI,
963 BuildFnTy &MatchInfo) const {
964 assert(MI.getOpcode() == TargetOpcode::G_AND);
965
966 // If we have the following code:
967 // %mask = G_CONSTANT 255
968 // %ld = G_LOAD %ptr, (load s16)
969 // %and = G_AND %ld, %mask
970 //
971 // Try to fold it into
972 // %ld = G_ZEXTLOAD %ptr, (load s8)
973
974 Register Dst = MI.getOperand(i: 0).getReg();
975 if (MRI.getType(Reg: Dst).isVector())
976 return false;
977
978 auto MaybeMask =
979 getIConstantVRegValWithLookThrough(VReg: MI.getOperand(i: 2).getReg(), MRI);
980 if (!MaybeMask)
981 return false;
982
983 APInt MaskVal = MaybeMask->Value;
984
985 if (!MaskVal.isMask())
986 return false;
987
988 Register SrcReg = MI.getOperand(i: 1).getReg();
989 // Don't use getOpcodeDef() here since intermediate instructions may have
990 // multiple users.
991 GAnyLoad *LoadMI;
992 Register PtrReg;
993 const MachineMemOperand *MMO;
994 if (!mi_match(R: SrcReg, MRI, P: m_GAnyLoad(Inst&: LoadMI, Ptr: m_Reg(R&: PtrReg), MMO: m_MMO(MMO))))
995 return false;
996
997 Register LoadReg = LoadMI->getDstReg();
998 LLT RegTy = MRI.getType(Reg: LoadReg);
999 unsigned RegSize = RegTy.getSizeInBits();
1000 unsigned LoadSizeBits = MMO->getSizeInBits().getValue();
1001 unsigned MaskSizeBits = MaskVal.countr_one();
1002
1003 if ((isa<GSExtLoad>(Val: LoadMI) || MaskSizeBits < LoadSizeBits) &&
1004 !MRI.hasOneNonDBGUse(RegNo: LoadReg))
1005 return false;
1006
1007 // The mask may not be larger than the in-memory type, as it might cover sign
1008 // extended bits
1009 if (MaskSizeBits > LoadSizeBits)
1010 return false;
1011
1012 // If the mask covers the whole destination register, there's nothing to
1013 // extend
1014 if (MaskSizeBits >= RegSize)
1015 return false;
1016
1017 // Most targets cannot deal with loads of size < 8 and need to re-legalize to
1018 // at least byte loads. Avoid creating such loads here
1019 if (MaskSizeBits < 8 || !isPowerOf2_32(Value: MaskSizeBits))
1020 return false;
1021
1022 LegalityQuery::MemDesc MemDesc(*MMO);
1023
1024 // Don't modify the memory access size if this is atomic/volatile, but we can
1025 // still adjust the opcode to indicate the high bit behavior.
1026 if (!MMO->isAtomic() && !MMO->isVolatile())
1027 MemDesc.MemoryTy = LLT::scalar(SizeInBits: MaskSizeBits);
1028 else if (LoadSizeBits > MaskSizeBits || LoadSizeBits == RegSize)
1029 return false;
1030
1031 // TODO: Could check if it's legal with the reduced or original memory size.
1032 if (!isLegalOrBeforeLegalizer(
1033 Query: {TargetOpcode::G_ZEXTLOAD, {RegTy, MRI.getType(Reg: PtrReg)}, {MemDesc}}))
1034 return false;
1035
1036 MatchInfo = [=](MachineIRBuilder &B) {
1037 B.setInstrAndDebugLoc(*LoadMI);
1038 auto &MF = B.getMF();
1039 auto PtrInfo = MMO->getPointerInfo();
1040 auto *NewMMO = MF.getMachineMemOperand(MMO, PtrInfo, Ty: MemDesc.MemoryTy);
1041 B.buildLoadInstr(Opcode: TargetOpcode::G_ZEXTLOAD, Res: Dst, Addr: PtrReg, MMO&: *NewMMO);
1042 replaceRegWith(MRI, FromReg: LoadReg, ToReg: Dst);
1043 LoadMI->eraseFromParent();
1044 };
1045 return true;
1046}
1047
1048bool CombinerHelper::isPredecessor(const MachineInstr &DefMI,
1049 const MachineInstr &UseMI) const {
1050 assert(!DefMI.isDebugInstr() && !UseMI.isDebugInstr() &&
1051 "shouldn't consider debug uses");
1052 assert(DefMI.getParent() == UseMI.getParent());
1053 if (&DefMI == &UseMI)
1054 return true;
1055 const MachineBasicBlock &MBB = *DefMI.getParent();
1056 auto DefOrUse = find_if(Range: MBB, P: [&DefMI, &UseMI](const MachineInstr &MI) {
1057 return &MI == &DefMI || &MI == &UseMI;
1058 });
1059 if (DefOrUse == MBB.end())
1060 llvm_unreachable("Block must contain both DefMI and UseMI!");
1061 return &*DefOrUse == &DefMI;
1062}
1063
1064bool CombinerHelper::dominates(const MachineInstr &DefMI,
1065 const MachineInstr &UseMI) const {
1066 assert(!DefMI.isDebugInstr() && !UseMI.isDebugInstr() &&
1067 "shouldn't consider debug uses");
1068 if (MDT)
1069 return MDT->dominates(A: &DefMI, B: &UseMI);
1070 else if (DefMI.getParent() != UseMI.getParent())
1071 return false;
1072
1073 return isPredecessor(DefMI, UseMI);
1074}
1075
1076bool CombinerHelper::matchSextTruncSextLoad(MachineInstr &MI) const {
1077 assert(MI.getOpcode() == TargetOpcode::G_SEXT_INREG);
1078 Register SrcReg = MI.getOperand(i: 1).getReg();
1079 Register LoadUser = SrcReg;
1080
1081 if (MRI.getType(Reg: SrcReg).isVector())
1082 return false;
1083
1084 Register TruncSrc;
1085 if (mi_match(R: SrcReg, MRI, P: m_GTrunc(Src: m_Reg(R&: TruncSrc))))
1086 LoadUser = TruncSrc;
1087
1088 uint64_t SizeInBits = MI.getOperand(i: 2).getImm();
1089 // If the source is a G_SEXTLOAD from the same bit width, then we don't
1090 // need any extend at all, just a truncate.
1091 if (auto *LoadMI = getOpcodeDef<GSExtLoad>(Reg: LoadUser, MRI)) {
1092 // If truncating more than the original extended value, abort.
1093 auto LoadSizeBits = LoadMI->getMemSizeInBits();
1094 if (TruncSrc &&
1095 MRI.getType(Reg: TruncSrc).getSizeInBits() < LoadSizeBits.getValue())
1096 return false;
1097 if (LoadSizeBits == SizeInBits)
1098 return true;
1099 }
1100 return false;
1101}
1102
1103bool CombinerHelper::matchSextInRegOfLoad(
1104 MachineInstr &MI, std::tuple<Register, unsigned> &MatchInfo) const {
1105 assert(MI.getOpcode() == TargetOpcode::G_SEXT_INREG);
1106
1107 Register DstReg = MI.getOperand(i: 0).getReg();
1108 LLT RegTy = MRI.getType(Reg: DstReg);
1109
1110 // Only supports scalars for now.
1111 if (RegTy.isVector())
1112 return false;
1113
1114 Register SrcReg = MI.getOperand(i: 1).getReg();
1115 Register PtrReg;
1116 const MachineMemOperand *MMO;
1117 if (!mi_match(R: SrcReg, MRI, P: m_GLoad(Ptr: m_Reg(R&: PtrReg), MMO: m_MMO(MMO))))
1118 return false;
1119
1120 uint64_t MemBits = MMO->getSizeInBits().getValue();
1121 uint64_t ExtFrom = MI.getOperand(i: 2).getImm();
1122
1123 if (MemBits > ExtFrom && !MRI.hasOneNonDBGUse(RegNo: SrcReg))
1124 return false;
1125
1126 // If the sign extend extends from a narrower width than the load's width,
1127 // then we can narrow the load width when we combine to a G_SEXTLOAD.
1128 // Avoid widening the load at all.
1129 unsigned NewSizeBits = std::min(a: ExtFrom, b: MemBits);
1130
1131 // Don't generate G_SEXTLOADs with a < 1 byte width.
1132 if (NewSizeBits < 8)
1133 return false;
1134 // Don't bother creating a non-power-2 sextload, it will likely be broken up
1135 // anyway for most targets.
1136 if (!isPowerOf2_32(Value: NewSizeBits))
1137 return false;
1138
1139 LegalityQuery::MemDesc MMDesc(*MMO);
1140
1141 // Don't modify the memory access size if this is atomic/volatile, but we can
1142 // still adjust the opcode to indicate the high bit behavior.
1143 if (!MMO->isAtomic() && !MMO->isVolatile())
1144 MMDesc.MemoryTy = LLT::scalar(SizeInBits: NewSizeBits);
1145 else if (MemBits > NewSizeBits || MemBits == RegTy.getSizeInBits())
1146 return false;
1147
1148 // TODO: Could check if it's legal with the reduced or original memory size.
1149 if (!isLegalOrBeforeLegalizer(
1150 Query: {TargetOpcode::G_SEXTLOAD, {RegTy, MRI.getType(Reg: PtrReg)}, {MMDesc}}))
1151 return false;
1152
1153 MatchInfo = std::make_tuple(args&: SrcReg, args&: NewSizeBits);
1154 return true;
1155}
1156
1157void CombinerHelper::applySextInRegOfLoad(
1158 MachineInstr &MI, std::tuple<Register, unsigned> &MatchInfo) const {
1159 assert(MI.getOpcode() == TargetOpcode::G_SEXT_INREG);
1160 Register LoadReg;
1161 unsigned ScalarSizeBits;
1162 std::tie(args&: LoadReg, args&: ScalarSizeBits) = MatchInfo;
1163 GLoad *LoadDef = cast<GLoad>(Val: MRI.getVRegDef(Reg: LoadReg));
1164
1165 // If we have the following:
1166 // %ld = G_LOAD %ptr, (load 2)
1167 // %ext = G_SEXT_INREG %ld, 8
1168 // ==>
1169 // %ld = G_SEXTLOAD %ptr (load 1)
1170
1171 auto &MMO = LoadDef->getMMO();
1172 Builder.setInstrAndDebugLoc(*LoadDef);
1173 auto &MF = Builder.getMF();
1174 auto PtrInfo = MMO.getPointerInfo();
1175 auto *NewMMO = MF.getMachineMemOperand(MMO: &MMO, PtrInfo, Size: ScalarSizeBits / 8);
1176 Builder.buildLoadInstr(Opcode: TargetOpcode::G_SEXTLOAD, Res: MI.getOperand(i: 0).getReg(),
1177 Addr: LoadDef->getPointerReg(), MMO&: *NewMMO);
1178 replaceRegWith(MRI, FromReg: LoadReg, ToReg: MI.getOperand(i: 0).getReg());
1179 MI.eraseFromParent();
1180
1181 // Not all loads can be deleted, so make sure the old one is removed.
1182 LoadDef->eraseFromParent();
1183}
1184
1185/// Return true if 'MI' is a load or a store that may be fold it's address
1186/// operand into the load / store addressing mode.
1187static bool canFoldInAddressingMode(GLoadStore *MI, const TargetLowering &TLI,
1188 MachineRegisterInfo &MRI) {
1189 TargetLowering::AddrMode AM;
1190 auto *MF = MI->getMF();
1191 auto *Addr = getOpcodeDef<GPtrAdd>(Reg: MI->getPointerReg(), MRI);
1192 if (!Addr)
1193 return false;
1194
1195 AM.HasBaseReg = true;
1196 if (auto CstOff = getIConstantVRegVal(VReg: Addr->getOffsetReg(), MRI))
1197 AM.BaseOffs = CstOff->getSExtValue(); // [reg +/- imm]
1198 else
1199 AM.Scale = 1; // [reg +/- reg]
1200
1201 return TLI.isLegalAddressingMode(
1202 DL: MF->getDataLayout(), AM,
1203 Ty: getTypeForLLT(Ty: MI->getMMO().getMemoryType(),
1204 C&: MF->getFunction().getContext()),
1205 AddrSpace: MI->getMMO().getAddrSpace());
1206}
1207
1208static unsigned getIndexedOpc(unsigned LdStOpc) {
1209 switch (LdStOpc) {
1210 case TargetOpcode::G_LOAD:
1211 return TargetOpcode::G_INDEXED_LOAD;
1212 case TargetOpcode::G_STORE:
1213 return TargetOpcode::G_INDEXED_STORE;
1214 case TargetOpcode::G_ZEXTLOAD:
1215 return TargetOpcode::G_INDEXED_ZEXTLOAD;
1216 case TargetOpcode::G_SEXTLOAD:
1217 return TargetOpcode::G_INDEXED_SEXTLOAD;
1218 default:
1219 llvm_unreachable("Unexpected opcode");
1220 }
1221}
1222
1223bool CombinerHelper::isIndexedLoadStoreLegal(GLoadStore &LdSt) const {
1224 // Check for legality.
1225 LLT PtrTy = MRI.getType(Reg: LdSt.getPointerReg());
1226 LLT Ty = MRI.getType(Reg: LdSt.getReg(Idx: 0));
1227 LLT MemTy = LdSt.getMMO().getMemoryType();
1228 SmallVector<LegalityQuery::MemDesc, 2> MemDescrs(
1229 {{MemTy, MemTy.getSizeInBits().getKnownMinValue(),
1230 AtomicOrdering::NotAtomic, AtomicOrdering::NotAtomic}});
1231 unsigned IndexedOpc = getIndexedOpc(LdStOpc: LdSt.getOpcode());
1232 SmallVector<LLT> OpTys;
1233 if (IndexedOpc == TargetOpcode::G_INDEXED_STORE)
1234 OpTys = {PtrTy, Ty, Ty};
1235 else
1236 OpTys = {Ty, PtrTy}; // For G_INDEXED_LOAD, G_INDEXED_[SZ]EXTLOAD
1237
1238 LegalityQuery Q(IndexedOpc, OpTys, MemDescrs);
1239 return isLegal(Query: Q);
1240}
1241
1242static cl::opt<unsigned> PostIndexUseThreshold(
1243 "post-index-use-threshold", cl::Hidden, cl::init(Val: 32),
1244 cl::desc("Number of uses of a base pointer to check before it is no longer "
1245 "considered for post-indexing."));
1246
1247bool CombinerHelper::findPostIndexCandidate(GLoadStore &LdSt, Register &Addr,
1248 Register &Base, Register &Offset,
1249 bool &RematOffset) const {
1250 // We're looking for the following pattern, for either load or store:
1251 // %baseptr:_(p0) = ...
1252 // G_STORE %val(s64), %baseptr(p0)
1253 // %offset:_(s64) = G_CONSTANT i64 -256
1254 // %new_addr:_(p0) = G_PTR_ADD %baseptr, %offset(s64)
1255 const auto &TLI = getTargetLowering();
1256
1257 Register Ptr = LdSt.getPointerReg();
1258 // If the store is the only use, don't bother.
1259 if (MRI.hasOneNonDBGUse(RegNo: Ptr))
1260 return false;
1261
1262 if (!isIndexedLoadStoreLegal(LdSt))
1263 return false;
1264
1265 if (getOpcodeDef(Opcode: TargetOpcode::G_FRAME_INDEX, Reg: Ptr, MRI))
1266 return false;
1267
1268 MachineInstr *StoredValDef = getDefIgnoringCopies(Reg: LdSt.getReg(Idx: 0), MRI);
1269 MachineInstr *PtrDef;
1270 if (!mi_match(R: Ptr, MRI, P: m_MInstr(MI&: PtrDef)))
1271 return false;
1272
1273 unsigned NumUsesChecked = 0;
1274 for (auto &Use : MRI.use_nodbg_instructions(Reg: Ptr)) {
1275 if (++NumUsesChecked > PostIndexUseThreshold)
1276 return false; // Try to avoid exploding compile time.
1277
1278 auto *PtrAdd = dyn_cast<GPtrAdd>(Val: &Use);
1279 // The use itself might be dead. This can happen during combines if DCE
1280 // hasn't had a chance to run yet. Don't allow it to form an indexed op.
1281 if (!PtrAdd || MRI.use_nodbg_empty(RegNo: PtrAdd->getReg(Idx: 0)))
1282 continue;
1283
1284 // Check the user of this isn't the store, otherwise we'd be generate a
1285 // indexed store defining its own use.
1286 if (StoredValDef == &Use)
1287 continue;
1288
1289 Offset = PtrAdd->getOffsetReg();
1290 if (!ForceLegalIndexing &&
1291 !TLI.isIndexingLegal(MI&: LdSt, Base: PtrAdd->getBaseReg(), Offset,
1292 /*IsPre*/ false, MRI))
1293 continue;
1294
1295 // Make sure the offset calculation is before the potentially indexed op.
1296 MachineInstr *OffsetDef;
1297 if (!mi_match(R: Offset, MRI, P: m_MInstr(MI&: OffsetDef)))
1298 continue;
1299 RematOffset = false;
1300 if (!dominates(DefMI: *OffsetDef, UseMI: LdSt)) {
1301 // If the offset however is just a G_CONSTANT, we can always just
1302 // rematerialize it where we need it.
1303 if (OffsetDef->getOpcode() != TargetOpcode::G_CONSTANT)
1304 continue;
1305 RematOffset = true;
1306 }
1307
1308 for (auto &BasePtrUse : MRI.use_nodbg_instructions(Reg: PtrAdd->getBaseReg())) {
1309 if (&BasePtrUse == PtrDef)
1310 continue;
1311
1312 // If the user is a later load/store that can be post-indexed, then don't
1313 // combine this one.
1314 auto *BasePtrLdSt = dyn_cast<GLoadStore>(Val: &BasePtrUse);
1315 if (BasePtrLdSt && BasePtrLdSt != &LdSt &&
1316 dominates(DefMI: LdSt, UseMI: *BasePtrLdSt) &&
1317 isIndexedLoadStoreLegal(LdSt&: *BasePtrLdSt))
1318 return false;
1319
1320 // Now we're looking for the key G_PTR_ADD instruction, which contains
1321 // the offset add that we want to fold.
1322 if (auto *BasePtrUseDef = dyn_cast<GPtrAdd>(Val: &BasePtrUse)) {
1323 Register PtrAddDefReg = BasePtrUseDef->getReg(Idx: 0);
1324 for (auto &BaseUseUse : MRI.use_nodbg_instructions(Reg: PtrAddDefReg)) {
1325 // If the use is in a different block, then we may produce worse code
1326 // due to the extra register pressure.
1327 if (BaseUseUse.getParent() != LdSt.getParent())
1328 return false;
1329
1330 if (auto *UseUseLdSt = dyn_cast<GLoadStore>(Val: &BaseUseUse))
1331 if (canFoldInAddressingMode(MI: UseUseLdSt, TLI, MRI))
1332 return false;
1333 }
1334 if (!dominates(DefMI: LdSt, UseMI: BasePtrUse))
1335 return false; // All use must be dominated by the load/store.
1336 }
1337 }
1338
1339 Addr = PtrAdd->getReg(Idx: 0);
1340 Base = PtrAdd->getBaseReg();
1341 return true;
1342 }
1343
1344 return false;
1345}
1346
1347bool CombinerHelper::findPreIndexCandidate(GLoadStore &LdSt, Register &Addr,
1348 Register &Base,
1349 Register &Offset) const {
1350 auto &MF = *LdSt.getParent()->getParent();
1351 const auto &TLI = *MF.getSubtarget().getTargetLowering();
1352
1353 Addr = LdSt.getPointerReg();
1354 if (!mi_match(R: Addr, MRI, P: m_GPtrAdd(L: m_Reg(R&: Base), R: m_Reg(R&: Offset))) ||
1355 MRI.hasOneNonDBGUse(RegNo: Addr))
1356 return false;
1357
1358 if (!ForceLegalIndexing &&
1359 !TLI.isIndexingLegal(MI&: LdSt, Base, Offset, /*IsPre*/ true, MRI))
1360 return false;
1361
1362 if (!isIndexedLoadStoreLegal(LdSt))
1363 return false;
1364
1365 MachineInstr *BaseDef = getDefIgnoringCopies(Reg: Base, MRI);
1366 if (BaseDef->getOpcode() == TargetOpcode::G_FRAME_INDEX)
1367 return false;
1368
1369 if (auto *St = dyn_cast<GStore>(Val: &LdSt)) {
1370 // Would require a copy.
1371 if (Base == St->getValueReg())
1372 return false;
1373
1374 // We're expecting one use of Addr in MI, but it could also be the
1375 // value stored, which isn't actually dominated by the instruction.
1376 if (St->getValueReg() == Addr)
1377 return false;
1378 }
1379
1380 // Avoid increasing cross-block register pressure.
1381 for (auto &AddrUse : MRI.use_nodbg_instructions(Reg: Addr))
1382 if (AddrUse.getParent() != LdSt.getParent())
1383 return false;
1384
1385 // FIXME: check whether all uses of the base pointer are constant PtrAdds.
1386 // That might allow us to end base's liveness here by adjusting the constant.
1387 bool RealUse = false;
1388 for (auto &AddrUse : MRI.use_nodbg_instructions(Reg: Addr)) {
1389 if (!dominates(DefMI: LdSt, UseMI: AddrUse))
1390 return false; // All use must be dominated by the load/store.
1391
1392 // If Ptr may be folded in addressing mode of other use, then it's
1393 // not profitable to do this transformation.
1394 if (auto *UseLdSt = dyn_cast<GLoadStore>(Val: &AddrUse)) {
1395 if (!canFoldInAddressingMode(MI: UseLdSt, TLI, MRI))
1396 RealUse = true;
1397 } else {
1398 RealUse = true;
1399 }
1400 }
1401 return RealUse;
1402}
1403
1404bool CombinerHelper::matchCombineExtractedVectorLoad(
1405 MachineInstr &MI, BuildFnTy &MatchInfo) const {
1406 assert(MI.getOpcode() == TargetOpcode::G_EXTRACT_VECTOR_ELT);
1407
1408 // Check if there is a load that defines the vector being extracted from.
1409 auto *LoadMI = getOpcodeDef<GLoad>(Reg: MI.getOperand(i: 1).getReg(), MRI);
1410 if (!LoadMI)
1411 return false;
1412
1413 Register Vector = MI.getOperand(i: 1).getReg();
1414 LLT VecEltTy = MRI.getType(Reg: Vector).getElementType();
1415
1416 assert(MRI.getType(MI.getOperand(0).getReg()) == VecEltTy);
1417
1418 // Checking whether we should reduce the load width.
1419 if (!MRI.hasOneNonDBGUse(RegNo: Vector))
1420 return false;
1421
1422 // Check if the defining load is simple.
1423 if (!LoadMI->isSimple())
1424 return false;
1425
1426 // If the vector element type is not a multiple of a byte then we are unable
1427 // to correctly compute an address to load only the extracted element as a
1428 // scalar.
1429 if (!VecEltTy.isByteSized())
1430 return false;
1431
1432 // Check for load fold barriers between the extraction and the load.
1433 if (MI.getParent() != LoadMI->getParent())
1434 return false;
1435 const unsigned MaxIter = 20;
1436 unsigned Iter = 0;
1437 for (auto II = LoadMI->getIterator(), IE = MI.getIterator(); II != IE; ++II) {
1438 if (II->isLoadFoldBarrier())
1439 return false;
1440 if (Iter++ == MaxIter)
1441 return false;
1442 }
1443
1444 // Check if the new load that we are going to create is legal
1445 // if we are in the post-legalization phase.
1446 MachineMemOperand MMO = LoadMI->getMMO();
1447 Align Alignment = MMO.getAlign();
1448 MachinePointerInfo PtrInfo;
1449 uint64_t Offset;
1450
1451 // Finding the appropriate PtrInfo if offset is a known constant.
1452 // This is required to create the memory operand for the narrowed load.
1453 // This machine memory operand object helps us infer about legality
1454 // before we proceed to combine the instruction.
1455 if (auto CVal = getIConstantVRegVal(VReg: Vector, MRI)) {
1456 int Elt = CVal->getZExtValue();
1457 // FIXME: should be (ABI size)*Elt.
1458 Offset = VecEltTy.getSizeInBits() * Elt / 8;
1459 PtrInfo = MMO.getPointerInfo().getWithOffset(O: Offset);
1460 } else {
1461 // Discard the pointer info except the address space because the memory
1462 // operand can't represent this new access since the offset is variable.
1463 Offset = VecEltTy.getSizeInBits() / 8;
1464 PtrInfo = MachinePointerInfo(MMO.getPointerInfo().getAddrSpace());
1465 }
1466
1467 Alignment = commonAlignment(A: Alignment, Offset);
1468
1469 Register VecPtr = LoadMI->getPointerReg();
1470 LLT PtrTy = MRI.getType(Reg: VecPtr);
1471
1472 MachineFunction &MF = *MI.getMF();
1473 auto *NewMMO = MF.getMachineMemOperand(MMO: &MMO, PtrInfo, Ty: VecEltTy);
1474
1475 LegalityQuery::MemDesc MMDesc(*NewMMO);
1476
1477 if (!isLegalOrBeforeLegalizer(
1478 Query: {TargetOpcode::G_LOAD, {VecEltTy, PtrTy}, {MMDesc}}))
1479 return false;
1480
1481 // Load must be allowed and fast on the target.
1482 LLVMContext &C = MF.getFunction().getContext();
1483 auto &DL = MF.getDataLayout();
1484 unsigned Fast = 0;
1485 if (!getTargetLowering().allowsMemoryAccess(Context&: C, DL, Ty: VecEltTy, MMO: *NewMMO,
1486 Fast: &Fast) ||
1487 !Fast)
1488 return false;
1489
1490 Register Result = MI.getOperand(i: 0).getReg();
1491 Register Index = MI.getOperand(i: 2).getReg();
1492
1493 MatchInfo = [=](MachineIRBuilder &B) {
1494 GISelObserverWrapper DummyObserver;
1495 LegalizerHelper Helper(B.getMF(), DummyObserver, B);
1496 //// Get pointer to the vector element.
1497 Register finalPtr = Helper.getVectorElementPointer(
1498 VecPtr: LoadMI->getPointerReg(), VecTy: MRI.getType(Reg: LoadMI->getOperand(i: 0).getReg()),
1499 Index);
1500 // New G_LOAD instruction.
1501 B.buildLoad(Res: Result, Addr: finalPtr, PtrInfo, Alignment);
1502 // Remove original GLOAD instruction.
1503 LoadMI->eraseFromParent();
1504 };
1505
1506 return true;
1507}
1508
1509bool CombinerHelper::matchCombineIndexedLoadStore(
1510 MachineInstr &MI, IndexedLoadStoreMatchInfo &MatchInfo) const {
1511 auto &LdSt = cast<GLoadStore>(Val&: MI);
1512
1513 if (LdSt.isAtomic())
1514 return false;
1515
1516 MatchInfo.IsPre = findPreIndexCandidate(LdSt, Addr&: MatchInfo.Addr, Base&: MatchInfo.Base,
1517 Offset&: MatchInfo.Offset);
1518 if (!MatchInfo.IsPre &&
1519 !findPostIndexCandidate(LdSt, Addr&: MatchInfo.Addr, Base&: MatchInfo.Base,
1520 Offset&: MatchInfo.Offset, RematOffset&: MatchInfo.RematOffset))
1521 return false;
1522
1523 return true;
1524}
1525
1526void CombinerHelper::applyCombineIndexedLoadStore(
1527 MachineInstr &MI, IndexedLoadStoreMatchInfo &MatchInfo) const {
1528 MachineInstr &AddrDef = *MRI.getVRegDef(Reg: MatchInfo.Addr);
1529 unsigned Opcode = MI.getOpcode();
1530 bool IsStore = Opcode == TargetOpcode::G_STORE;
1531 unsigned NewOpcode = getIndexedOpc(LdStOpc: Opcode);
1532
1533 // If the offset constant didn't happen to dominate the load/store, we can
1534 // just clone it as needed.
1535 if (MatchInfo.RematOffset) {
1536 auto *OldCst = MRI.getVRegDef(Reg: MatchInfo.Offset);
1537 auto NewCst = Builder.buildConstant(Res: MRI.getType(Reg: MatchInfo.Offset),
1538 Val: *OldCst->getOperand(i: 1).getCImm());
1539 MatchInfo.Offset = NewCst.getReg(Idx: 0);
1540 }
1541
1542 auto MIB = Builder.buildInstr(Opcode: NewOpcode);
1543 if (IsStore) {
1544 MIB.addDef(RegNo: MatchInfo.Addr);
1545 MIB.addUse(RegNo: MI.getOperand(i: 0).getReg());
1546 } else {
1547 MIB.addDef(RegNo: MI.getOperand(i: 0).getReg());
1548 MIB.addDef(RegNo: MatchInfo.Addr);
1549 }
1550
1551 MIB.addUse(RegNo: MatchInfo.Base);
1552 MIB.addUse(RegNo: MatchInfo.Offset);
1553 MIB.addImm(Val: MatchInfo.IsPre);
1554 MIB->cloneMemRefs(MF&: *MI.getMF(), MI);
1555 MI.eraseFromParent();
1556 AddrDef.eraseFromParent();
1557
1558 LLVM_DEBUG(dbgs() << " Combinined to indexed operation");
1559}
1560
1561bool CombinerHelper::matchCombineDivRem(MachineInstr &MI,
1562 MachineInstr *&OtherMI) const {
1563 unsigned Opcode = MI.getOpcode();
1564 bool IsDiv, IsSigned;
1565
1566 switch (Opcode) {
1567 default:
1568 llvm_unreachable("Unexpected opcode!");
1569 case TargetOpcode::G_SDIV:
1570 case TargetOpcode::G_UDIV: {
1571 IsDiv = true;
1572 IsSigned = Opcode == TargetOpcode::G_SDIV;
1573 break;
1574 }
1575 case TargetOpcode::G_SREM:
1576 case TargetOpcode::G_UREM: {
1577 IsDiv = false;
1578 IsSigned = Opcode == TargetOpcode::G_SREM;
1579 break;
1580 }
1581 }
1582
1583 Register Src1 = MI.getOperand(i: 1).getReg();
1584 unsigned DivOpcode, RemOpcode, DivremOpcode;
1585 if (IsSigned) {
1586 DivOpcode = TargetOpcode::G_SDIV;
1587 RemOpcode = TargetOpcode::G_SREM;
1588 DivremOpcode = TargetOpcode::G_SDIVREM;
1589 } else {
1590 DivOpcode = TargetOpcode::G_UDIV;
1591 RemOpcode = TargetOpcode::G_UREM;
1592 DivremOpcode = TargetOpcode::G_UDIVREM;
1593 }
1594
1595 if (!isLegalOrBeforeLegalizer(Query: {DivremOpcode, {MRI.getType(Reg: Src1)}}))
1596 return false;
1597
1598 // Combine:
1599 // %div:_ = G_[SU]DIV %src1:_, %src2:_
1600 // %rem:_ = G_[SU]REM %src1:_, %src2:_
1601 // into:
1602 // %div:_, %rem:_ = G_[SU]DIVREM %src1:_, %src2:_
1603
1604 // Combine:
1605 // %rem:_ = G_[SU]REM %src1:_, %src2:_
1606 // %div:_ = G_[SU]DIV %src1:_, %src2:_
1607 // into:
1608 // %div:_, %rem:_ = G_[SU]DIVREM %src1:_, %src2:_
1609
1610 for (auto &UseMI : MRI.use_nodbg_instructions(Reg: Src1)) {
1611 if (MI.getParent() == UseMI.getParent() &&
1612 ((IsDiv && UseMI.getOpcode() == RemOpcode) ||
1613 (!IsDiv && UseMI.getOpcode() == DivOpcode)) &&
1614 matchEqualDefs(MOP1: MI.getOperand(i: 2), MOP2: UseMI.getOperand(i: 2)) &&
1615 matchEqualDefs(MOP1: MI.getOperand(i: 1), MOP2: UseMI.getOperand(i: 1))) {
1616 OtherMI = &UseMI;
1617 return true;
1618 }
1619 }
1620
1621 return false;
1622}
1623
1624void CombinerHelper::applyCombineDivRem(MachineInstr &MI,
1625 MachineInstr *&OtherMI) const {
1626 unsigned Opcode = MI.getOpcode();
1627 assert(OtherMI && "OtherMI shouldn't be empty.");
1628
1629 Register DestDivReg, DestRemReg;
1630 if (Opcode == TargetOpcode::G_SDIV || Opcode == TargetOpcode::G_UDIV) {
1631 DestDivReg = MI.getOperand(i: 0).getReg();
1632 DestRemReg = OtherMI->getOperand(i: 0).getReg();
1633 } else {
1634 DestDivReg = OtherMI->getOperand(i: 0).getReg();
1635 DestRemReg = MI.getOperand(i: 0).getReg();
1636 }
1637
1638 bool IsSigned =
1639 Opcode == TargetOpcode::G_SDIV || Opcode == TargetOpcode::G_SREM;
1640
1641 // Check which instruction is first in the block so we don't break def-use
1642 // deps by "moving" the instruction incorrectly. Also keep track of which
1643 // instruction is first so we pick it's operands, avoiding use-before-def
1644 // bugs.
1645 MachineInstr *FirstInst = dominates(DefMI: MI, UseMI: *OtherMI) ? &MI : OtherMI;
1646 Builder.setInstrAndDebugLoc(*FirstInst);
1647
1648 Builder.buildInstr(Opc: IsSigned ? TargetOpcode::G_SDIVREM
1649 : TargetOpcode::G_UDIVREM,
1650 DstOps: {DestDivReg, DestRemReg},
1651 SrcOps: { FirstInst->getOperand(i: 1), FirstInst->getOperand(i: 2) });
1652 MI.eraseFromParent();
1653 OtherMI->eraseFromParent();
1654}
1655
1656bool CombinerHelper::matchOptBrCondByInvertingCond(
1657 MachineInstr &MI, MachineInstr *&BrCond) const {
1658 assert(MI.getOpcode() == TargetOpcode::G_BR);
1659
1660 // Try to match the following:
1661 // bb1:
1662 // G_BRCOND %c1, %bb2
1663 // G_BR %bb3
1664 // bb2:
1665 // ...
1666 // bb3:
1667
1668 // The above pattern does not have a fall through to the successor bb2, always
1669 // resulting in a branch no matter which path is taken. Here we try to find
1670 // and replace that pattern with conditional branch to bb3 and otherwise
1671 // fallthrough to bb2. This is generally better for branch predictors.
1672
1673 MachineBasicBlock *MBB = MI.getParent();
1674 MachineBasicBlock::iterator BrIt(MI);
1675 if (BrIt == MBB->begin())
1676 return false;
1677 assert(std::next(BrIt) == MBB->end() && "expected G_BR to be a terminator");
1678
1679 BrCond = &*std::prev(x: BrIt);
1680 if (BrCond->getOpcode() != TargetOpcode::G_BRCOND)
1681 return false;
1682
1683 // Check that the next block is the conditional branch target. Also make sure
1684 // that it isn't the same as the G_BR's target (otherwise, this will loop.)
1685 MachineBasicBlock *BrCondTarget = BrCond->getOperand(i: 1).getMBB();
1686 return BrCondTarget != MI.getOperand(i: 0).getMBB() &&
1687 MBB->isLayoutSuccessor(MBB: BrCondTarget);
1688}
1689
1690void CombinerHelper::applyOptBrCondByInvertingCond(
1691 MachineInstr &MI, MachineInstr *&BrCond) const {
1692 MachineBasicBlock *BrTarget = MI.getOperand(i: 0).getMBB();
1693 Builder.setInstrAndDebugLoc(*BrCond);
1694 LLT Ty = MRI.getType(Reg: BrCond->getOperand(i: 0).getReg());
1695 // FIXME: Does int/fp matter for this? If so, we might need to restrict
1696 // this to i1 only since we might not know for sure what kind of
1697 // compare generated the condition value.
1698 auto True = Builder.buildConstant(
1699 Res: Ty, Val: getICmpTrueVal(TLI: getTargetLowering(), IsVector: false, IsFP: false));
1700 auto Xor = Builder.buildXor(Dst: Ty, Src0: BrCond->getOperand(i: 0), Src1: True);
1701
1702 auto *FallthroughBB = BrCond->getOperand(i: 1).getMBB();
1703 Observer.changingInstr(MI);
1704 MI.getOperand(i: 0).setMBB(FallthroughBB);
1705 Observer.changedInstr(MI);
1706
1707 // Change the conditional branch to use the inverted condition and
1708 // new target block.
1709 Observer.changingInstr(MI&: *BrCond);
1710 BrCond->getOperand(i: 0).setReg(Xor.getReg(Idx: 0));
1711 BrCond->getOperand(i: 1).setMBB(BrTarget);
1712 Observer.changedInstr(MI&: *BrCond);
1713}
1714
1715bool CombinerHelper::matchCombineMemCpyFamily(
1716 MachineInstr &MI, MemCpyFamilyLoweringInfo &MatchInfo,
1717 unsigned MaxLen) const {
1718 auto &[Dst, Src, KnownLen, Alignment, DstAlignCanChange, MemOps] = MatchInfo;
1719 return canLowerMemCpyFamily(MI, MRI, MaxLen, Dst, Src, KnownLen, Alignment,
1720 DstAlignCanChange, MemOps);
1721}
1722
1723void CombinerHelper::applyCombineMemCpyFamily(
1724 MachineInstr &MI, MemCpyFamilyLoweringInfo &MatchInfo) const {
1725 auto &[Dst, Src, KnownLen, Alignment, DstAlignCanChange, MemOps] = MatchInfo;
1726 MachineIRBuilder HelperBuilder(MI);
1727 GISelObserverWrapper DummyObserver;
1728 LegalizerHelper Helper(HelperBuilder.getMF(), DummyObserver, HelperBuilder);
1729 bool Changed = Helper.lowerMemCpyFamily(MI, Dst, Src, KnownLen, Alignment,
1730 DstAlignCanChange, MemOps) ==
1731 LegalizerHelper::LegalizeResult::Legalized;
1732 assert(Changed && "expected memcpy-family instruction to lower");
1733 (void)Changed;
1734}
1735
1736bool CombinerHelper::tryCombineMemCpyFamily(MachineInstr &MI,
1737 unsigned MaxLen) const {
1738 MachineIRBuilder HelperBuilder(MI);
1739 GISelObserverWrapper DummyObserver;
1740 LegalizerHelper Helper(HelperBuilder.getMF(), DummyObserver, HelperBuilder);
1741 return Helper.lowerMemCpyFamily(MI, MaxLen) ==
1742 LegalizerHelper::LegalizeResult::Legalized;
1743}
1744
1745static APFloat constantFoldFpUnary(const MachineInstr &MI,
1746 const MachineRegisterInfo &MRI,
1747 const APFloat &Val) {
1748 APFloat Result(Val);
1749 switch (MI.getOpcode()) {
1750 default:
1751 llvm_unreachable("Unexpected opcode!");
1752 case TargetOpcode::G_FNEG: {
1753 Result.changeSign();
1754 return Result;
1755 }
1756 case TargetOpcode::G_FABS: {
1757 Result.clearSign();
1758 return Result;
1759 }
1760 case TargetOpcode::G_FCEIL:
1761 Result.roundToIntegral(RM: APFloat::rmTowardPositive);
1762 return Result;
1763 case TargetOpcode::G_FFLOOR:
1764 Result.roundToIntegral(RM: APFloat::rmTowardNegative);
1765 return Result;
1766 case TargetOpcode::G_INTRINSIC_TRUNC:
1767 Result.roundToIntegral(RM: APFloat::rmTowardZero);
1768 return Result;
1769 case TargetOpcode::G_INTRINSIC_ROUND:
1770 Result.roundToIntegral(RM: APFloat::rmNearestTiesToAway);
1771 return Result;
1772 case TargetOpcode::G_INTRINSIC_ROUNDEVEN:
1773 Result.roundToIntegral(RM: APFloat::rmNearestTiesToEven);
1774 return Result;
1775 case TargetOpcode::G_FRINT:
1776 case TargetOpcode::G_FNEARBYINT:
1777 // Use default rounding mode (round to nearest, ties to even)
1778 Result.roundToIntegral(RM: APFloat::rmNearestTiesToEven);
1779 return Result;
1780 case TargetOpcode::G_FPEXT:
1781 case TargetOpcode::G_FPTRUNC: {
1782 bool Unused;
1783 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
1784 Result.convert(ToSemantics: getFltSemanticForLLT(Ty: DstTy), RM: APFloat::rmNearestTiesToEven,
1785 losesInfo: &Unused);
1786 return Result;
1787 }
1788 case TargetOpcode::G_FSQRT: {
1789 bool Unused;
1790 Result.convert(ToSemantics: APFloat::IEEEdouble(), RM: APFloat::rmNearestTiesToEven,
1791 losesInfo: &Unused);
1792 Result = APFloat(sqrt(x: Result.convertToDouble()));
1793 break;
1794 }
1795 case TargetOpcode::G_FLOG2: {
1796 bool Unused;
1797 Result.convert(ToSemantics: APFloat::IEEEdouble(), RM: APFloat::rmNearestTiesToEven,
1798 losesInfo: &Unused);
1799 Result = APFloat(log2(x: Result.convertToDouble()));
1800 break;
1801 }
1802 }
1803 // Convert `APFloat` to appropriate IEEE type depending on `DstTy`. Otherwise,
1804 // `buildFConstant` will assert on size mismatch. Only `G_FSQRT`, and
1805 // `G_FLOG2` reach here.
1806 bool Unused;
1807 Result.convert(ToSemantics: Val.getSemantics(), RM: APFloat::rmNearestTiesToEven, losesInfo: &Unused);
1808 return Result;
1809}
1810
1811void CombinerHelper::applyCombineConstantFoldFpUnary(
1812 MachineInstr &MI, const ConstantFP *Cst) const {
1813 APFloat Folded = constantFoldFpUnary(MI, MRI, Val: Cst->getValue());
1814 const ConstantFP *NewCst = ConstantFP::get(Context&: Builder.getContext(), V: Folded);
1815 Builder.buildFConstant(Res: MI.getOperand(i: 0), Val: *NewCst);
1816 MI.eraseFromParent();
1817}
1818
1819bool CombinerHelper::matchPtrAddImmedChain(MachineInstr &MI,
1820 PtrAddChain &MatchInfo) const {
1821 // We're trying to match the following pattern:
1822 // %t1 = G_PTR_ADD %base, G_CONSTANT imm1
1823 // %root = G_PTR_ADD %t1, G_CONSTANT imm2
1824 // -->
1825 // %root = G_PTR_ADD %base, G_CONSTANT (imm1 + imm2)
1826
1827 if (MI.getOpcode() != TargetOpcode::G_PTR_ADD)
1828 return false;
1829
1830 Register Add2 = MI.getOperand(i: 1).getReg();
1831 Register Imm1 = MI.getOperand(i: 2).getReg();
1832 auto MaybeImmVal = getIConstantVRegValWithLookThrough(VReg: Imm1, MRI);
1833 if (!MaybeImmVal)
1834 return false;
1835
1836 Register Base, Imm2;
1837 uint32_t LHSPtrAddFlags;
1838 if (!mi_match(R: Add2, MRI,
1839 P: m_GPtrAdd(L: m_Reg(R&: Base), R: m_Reg(R&: Imm2), Flags: m_MIFlags(Flags&: LHSPtrAddFlags))))
1840 return false;
1841
1842 auto MaybeImm2Val = getIConstantVRegValWithLookThrough(VReg: Imm2, MRI);
1843 if (!MaybeImm2Val)
1844 return false;
1845
1846 // Check if the new combined immediate forms an illegal addressing mode.
1847 // Do not combine if it was legal before but would get illegal.
1848 // To do so, we need to find a load/store user of the pointer to get
1849 // the access type.
1850 Type *AccessTy = nullptr;
1851 auto &MF = *MI.getMF();
1852 for (auto &UseMI : MRI.use_nodbg_instructions(Reg: MI.getOperand(i: 0).getReg())) {
1853 if (auto *LdSt = dyn_cast<GLoadStore>(Val: &UseMI)) {
1854 AccessTy = getTypeForLLT(Ty: MRI.getType(Reg: LdSt->getReg(Idx: 0)),
1855 C&: MF.getFunction().getContext());
1856 break;
1857 }
1858 }
1859 TargetLoweringBase::AddrMode AMNew;
1860 APInt CombinedImm = MaybeImmVal->Value + MaybeImm2Val->Value;
1861 AMNew.BaseOffs = CombinedImm.getSExtValue();
1862 if (AccessTy) {
1863 AMNew.HasBaseReg = true;
1864 TargetLoweringBase::AddrMode AMOld;
1865 AMOld.BaseOffs = MaybeImmVal->Value.getSExtValue();
1866 AMOld.HasBaseReg = true;
1867 unsigned AS = MRI.getType(Reg: Add2).getAddressSpace();
1868 const auto &TLI = *MF.getSubtarget().getTargetLowering();
1869 if (TLI.isLegalAddressingMode(DL: MF.getDataLayout(), AM: AMOld, Ty: AccessTy, AddrSpace: AS) &&
1870 !TLI.isLegalAddressingMode(DL: MF.getDataLayout(), AM: AMNew, Ty: AccessTy, AddrSpace: AS))
1871 return false;
1872 }
1873
1874 // Reassociating nuw additions preserves nuw. If both original G_PTR_ADDs are
1875 // inbounds, reaching the same result in one G_PTR_ADD is also inbounds.
1876 // The nusw constraints are satisfied because imm1+imm2 cannot exceed the
1877 // largest signed integer that fits into the index type, which is the maximum
1878 // size of allocated objects according to the IR Language Reference.
1879 unsigned PtrAddFlags = MI.getFlags();
1880 bool IsNoUWrap = PtrAddFlags & LHSPtrAddFlags & MachineInstr::MIFlag::NoUWrap;
1881 bool IsInBounds =
1882 PtrAddFlags & LHSPtrAddFlags & MachineInstr::MIFlag::InBounds;
1883 unsigned Flags = 0;
1884 if (IsNoUWrap)
1885 Flags |= MachineInstr::MIFlag::NoUWrap;
1886 if (IsInBounds) {
1887 Flags |= MachineInstr::MIFlag::InBounds;
1888 Flags |= MachineInstr::MIFlag::NoUSWrap;
1889 }
1890
1891 // Pass the combined immediate to the apply function.
1892 MatchInfo.Imm = AMNew.BaseOffs;
1893 MatchInfo.Base = Base;
1894 MatchInfo.Bank = getRegBank(Reg: Imm2);
1895 MatchInfo.Flags = Flags;
1896 return true;
1897}
1898
1899void CombinerHelper::applyPtrAddImmedChain(MachineInstr &MI,
1900 PtrAddChain &MatchInfo) const {
1901 assert(MI.getOpcode() == TargetOpcode::G_PTR_ADD && "Expected G_PTR_ADD");
1902 MachineIRBuilder MIB(MI);
1903 LLT OffsetTy = MRI.getType(Reg: MI.getOperand(i: 2).getReg());
1904 auto NewOffset = MIB.buildConstant(Res: OffsetTy, Val: MatchInfo.Imm);
1905 setRegBank(Reg: NewOffset.getReg(Idx: 0), RegBank: MatchInfo.Bank);
1906 Observer.changingInstr(MI);
1907 MI.getOperand(i: 1).setReg(MatchInfo.Base);
1908 MI.getOperand(i: 2).setReg(NewOffset.getReg(Idx: 0));
1909 MI.setFlags(MatchInfo.Flags);
1910 Observer.changedInstr(MI);
1911}
1912
1913bool CombinerHelper::matchShiftImmedChain(MachineInstr &MI,
1914 RegisterImmPair &MatchInfo) const {
1915 // We're trying to match the following pattern with any of
1916 // G_SHL/G_ASHR/G_LSHR/G_SSHLSAT/G_USHLSAT shift instructions:
1917 // %t1 = SHIFT %base, G_CONSTANT imm1
1918 // %root = SHIFT %t1, G_CONSTANT imm2
1919 // -->
1920 // %root = SHIFT %base, G_CONSTANT (imm1 + imm2)
1921
1922 unsigned Opcode = MI.getOpcode();
1923 assert((Opcode == TargetOpcode::G_SHL || Opcode == TargetOpcode::G_ASHR ||
1924 Opcode == TargetOpcode::G_LSHR || Opcode == TargetOpcode::G_SSHLSAT ||
1925 Opcode == TargetOpcode::G_USHLSAT) &&
1926 "Expected G_SHL, G_ASHR, G_LSHR, G_SSHLSAT or G_USHLSAT");
1927
1928 Register Shl2 = MI.getOperand(i: 1).getReg();
1929 Register Imm1 = MI.getOperand(i: 2).getReg();
1930 auto MaybeImmVal = getIConstantVRegValWithLookThrough(VReg: Imm1, MRI);
1931 if (!MaybeImmVal)
1932 return false;
1933
1934 MachineInstr *Shl2Def;
1935 if (!mi_match(R: Shl2, MRI, P: m_MInstr(MI&: Shl2Def)) || Shl2Def->getOpcode() != Opcode)
1936 return false;
1937
1938 Register Base = Shl2Def->getOperand(i: 1).getReg();
1939 Register Imm2 = Shl2Def->getOperand(i: 2).getReg();
1940 auto MaybeImm2Val = getIConstantVRegValWithLookThrough(VReg: Imm2, MRI);
1941 if (!MaybeImm2Val)
1942 return false;
1943
1944 // Pass the combined immediate to the apply function.
1945 MatchInfo.Imm =
1946 (MaybeImmVal->Value.getZExtValue() + MaybeImm2Val->Value).getZExtValue();
1947 MatchInfo.Reg = Base;
1948
1949 // There is no simple replacement for a saturating unsigned left shift that
1950 // exceeds the scalar size.
1951 if (Opcode == TargetOpcode::G_USHLSAT &&
1952 MatchInfo.Imm >= MRI.getType(Reg: Shl2).getScalarSizeInBits())
1953 return false;
1954
1955 return true;
1956}
1957
1958void CombinerHelper::applyShiftImmedChain(MachineInstr &MI,
1959 RegisterImmPair &MatchInfo) const {
1960 unsigned Opcode = MI.getOpcode();
1961 assert((Opcode == TargetOpcode::G_SHL || Opcode == TargetOpcode::G_ASHR ||
1962 Opcode == TargetOpcode::G_LSHR || Opcode == TargetOpcode::G_SSHLSAT ||
1963 Opcode == TargetOpcode::G_USHLSAT) &&
1964 "Expected G_SHL, G_ASHR, G_LSHR, G_SSHLSAT or G_USHLSAT");
1965
1966 LLT Ty = MRI.getType(Reg: MI.getOperand(i: 1).getReg());
1967 unsigned const ScalarSizeInBits = Ty.getScalarSizeInBits();
1968 auto Imm = MatchInfo.Imm;
1969
1970 if (Imm >= ScalarSizeInBits) {
1971 // Any logical shift that exceeds scalar size will produce zero.
1972 if (Opcode == TargetOpcode::G_SHL || Opcode == TargetOpcode::G_LSHR) {
1973 Builder.buildConstant(Res: MI.getOperand(i: 0), Val: 0);
1974 MI.eraseFromParent();
1975 return;
1976 }
1977 // Arithmetic shift and saturating signed left shift have no effect beyond
1978 // scalar size.
1979 Imm = ScalarSizeInBits - 1;
1980 }
1981
1982 LLT ImmTy = MRI.getType(Reg: MI.getOperand(i: 2).getReg());
1983 Register NewImm = Builder.buildConstant(Res: ImmTy, Val: Imm).getReg(Idx: 0);
1984 Observer.changingInstr(MI);
1985 MI.getOperand(i: 1).setReg(MatchInfo.Reg);
1986 MI.getOperand(i: 2).setReg(NewImm);
1987 Observer.changedInstr(MI);
1988}
1989
1990bool CombinerHelper::matchShiftOfShiftedLogic(
1991 MachineInstr &MI, ShiftOfShiftedLogic &MatchInfo) const {
1992 // We're trying to match the following pattern with any of
1993 // G_SHL/G_ASHR/G_LSHR/G_USHLSAT/G_SSHLSAT shift instructions in combination
1994 // with any of G_AND/G_OR/G_XOR logic instructions.
1995 // %t1 = SHIFT %X, G_CONSTANT C0
1996 // %t2 = LOGIC %t1, %Y
1997 // %root = SHIFT %t2, G_CONSTANT C1
1998 // -->
1999 // %t3 = SHIFT %X, G_CONSTANT (C0+C1)
2000 // %t4 = SHIFT %Y, G_CONSTANT C1
2001 // %root = LOGIC %t3, %t4
2002 unsigned ShiftOpcode = MI.getOpcode();
2003 assert((ShiftOpcode == TargetOpcode::G_SHL ||
2004 ShiftOpcode == TargetOpcode::G_ASHR ||
2005 ShiftOpcode == TargetOpcode::G_LSHR ||
2006 ShiftOpcode == TargetOpcode::G_USHLSAT ||
2007 ShiftOpcode == TargetOpcode::G_SSHLSAT) &&
2008 "Expected G_SHL, G_ASHR, G_LSHR, G_USHLSAT and G_SSHLSAT");
2009
2010 // Match a one-use bitwise logic op.
2011 Register LogicDest = MI.getOperand(i: 1).getReg();
2012 if (!MRI.hasOneNonDBGUse(RegNo: LogicDest))
2013 return false;
2014
2015 MachineInstr *LogicMI;
2016 if (!mi_match(R: LogicDest, MRI, P: m_MInstr(MI&: LogicMI)))
2017 return false;
2018 unsigned LogicOpcode = LogicMI->getOpcode();
2019 if (LogicOpcode != TargetOpcode::G_AND && LogicOpcode != TargetOpcode::G_OR &&
2020 LogicOpcode != TargetOpcode::G_XOR)
2021 return false;
2022
2023 // Find a matching one-use shift by constant.
2024 const Register C1 = MI.getOperand(i: 2).getReg();
2025 auto MaybeImmVal = getIConstantVRegValWithLookThrough(VReg: C1, MRI);
2026 if (!MaybeImmVal || MaybeImmVal->Value == 0)
2027 return false;
2028
2029 const uint64_t C1Val = MaybeImmVal->Value.getZExtValue();
2030
2031 auto matchFirstShift = [&](const MachineInstr *MI, uint64_t &ShiftVal) {
2032 // Shift should match previous one and should be a one-use.
2033 if (MI->getOpcode() != ShiftOpcode ||
2034 !MRI.hasOneNonDBGUse(RegNo: MI->getOperand(i: 0).getReg()))
2035 return false;
2036
2037 // Must be a constant.
2038 auto MaybeImmVal =
2039 getIConstantVRegValWithLookThrough(VReg: MI->getOperand(i: 2).getReg(), MRI);
2040 if (!MaybeImmVal)
2041 return false;
2042
2043 ShiftVal = MaybeImmVal->Value.getSExtValue();
2044 return true;
2045 };
2046
2047 // Logic ops are commutative, so check each operand for a match.
2048 Register LogicMIReg1 = LogicMI->getOperand(i: 1).getReg();
2049 MachineInstr *LogicMIOp1;
2050 Register LogicMIReg2 = LogicMI->getOperand(i: 2).getReg();
2051 MachineInstr *LogicMIOp2;
2052 if (!mi_match(R: LogicMIReg1, MRI, P: m_MInstr(MI&: LogicMIOp1)) ||
2053 !mi_match(R: LogicMIReg2, MRI, P: m_MInstr(MI&: LogicMIOp2)))
2054 return false;
2055 uint64_t C0Val;
2056
2057 if (matchFirstShift(LogicMIOp1, C0Val)) {
2058 MatchInfo.LogicNonShiftReg = LogicMIReg2;
2059 MatchInfo.Shift2 = LogicMIOp1;
2060 } else if (matchFirstShift(LogicMIOp2, C0Val)) {
2061 MatchInfo.LogicNonShiftReg = LogicMIReg1;
2062 MatchInfo.Shift2 = LogicMIOp2;
2063 } else
2064 return false;
2065
2066 MatchInfo.ValSum = C0Val + C1Val;
2067
2068 // The fold is not valid if the sum of the shift values exceeds bitwidth.
2069 if (MatchInfo.ValSum >= MRI.getType(Reg: LogicDest).getScalarSizeInBits())
2070 return false;
2071
2072 MatchInfo.Logic = LogicMI;
2073 return true;
2074}
2075
2076void CombinerHelper::applyShiftOfShiftedLogic(
2077 MachineInstr &MI, ShiftOfShiftedLogic &MatchInfo) const {
2078 unsigned Opcode = MI.getOpcode();
2079 assert((Opcode == TargetOpcode::G_SHL || Opcode == TargetOpcode::G_ASHR ||
2080 Opcode == TargetOpcode::G_LSHR || Opcode == TargetOpcode::G_USHLSAT ||
2081 Opcode == TargetOpcode::G_SSHLSAT) &&
2082 "Expected G_SHL, G_ASHR, G_LSHR, G_USHLSAT and G_SSHLSAT");
2083
2084 LLT ShlType = MRI.getType(Reg: MI.getOperand(i: 2).getReg());
2085 LLT DestType = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
2086
2087 Register Const = Builder.buildConstant(Res: ShlType, Val: MatchInfo.ValSum).getReg(Idx: 0);
2088
2089 Register Shift1Base = MatchInfo.Shift2->getOperand(i: 1).getReg();
2090 Register Shift1 =
2091 Builder.buildInstr(Opc: Opcode, DstOps: {DestType}, SrcOps: {Shift1Base, Const}).getReg(Idx: 0);
2092
2093 // If LogicNonShiftReg is the same to Shift1Base, and shift1 const is the same
2094 // to MatchInfo.Shift2 const, CSEMIRBuilder will reuse the old shift1 when
2095 // build shift2. So, if we erase MatchInfo.Shift2 at the end, actually we
2096 // remove old shift1. And it will cause crash later. So erase it earlier to
2097 // avoid the crash.
2098 MatchInfo.Shift2->eraseFromParent();
2099
2100 Register Shift2Const = MI.getOperand(i: 2).getReg();
2101 Register Shift2 = Builder
2102 .buildInstr(Opc: Opcode, DstOps: {DestType},
2103 SrcOps: {MatchInfo.LogicNonShiftReg, Shift2Const})
2104 .getReg(Idx: 0);
2105
2106 Register Dest = MI.getOperand(i: 0).getReg();
2107 Builder.buildInstr(Opc: MatchInfo.Logic->getOpcode(), DstOps: {Dest}, SrcOps: {Shift1, Shift2});
2108
2109 // This was one use so it's safe to remove it.
2110 MatchInfo.Logic->eraseFromParent();
2111
2112 MI.eraseFromParent();
2113}
2114
2115bool CombinerHelper::isDesirableToCommuteWithShift(
2116 const MachineInstr &MI) const {
2117 return getTargetLowering().isDesirableToCommuteWithShift(MI,
2118 IsAfterLegal: !isPreLegalize());
2119}
2120
2121bool CombinerHelper::matchLshrOfTruncOfLshr(MachineInstr &MI,
2122 LshrOfTruncOfLshr &MatchInfo,
2123 MachineInstr &ShiftMI) const {
2124 assert(MI.getOpcode() == TargetOpcode::G_LSHR && "Expected a G_LSHR");
2125
2126 Register N0 = MI.getOperand(i: 1).getReg();
2127 Register N1 = MI.getOperand(i: 2).getReg();
2128 unsigned OpSizeInBits = MRI.getType(Reg: N0).getScalarSizeInBits();
2129
2130 APInt N1C, N001C;
2131 if (!mi_match(R: N1, MRI, P: m_ICstOrSplat(Cst&: N1C)))
2132 return false;
2133 auto N001 = ShiftMI.getOperand(i: 2).getReg();
2134 if (!mi_match(R: N001, MRI, P: m_ICstOrSplat(Cst&: N001C)))
2135 return false;
2136
2137 if (N001C.getBitWidth() > N1C.getBitWidth())
2138 N1C = N1C.zext(width: N001C.getBitWidth());
2139 else
2140 N001C = N001C.zext(width: N1C.getBitWidth());
2141
2142 Register InnerShift = ShiftMI.getOperand(i: 0).getReg();
2143 LLT InnerShiftTy = MRI.getType(Reg: InnerShift);
2144 uint64_t InnerShiftSize = InnerShiftTy.getScalarSizeInBits();
2145 if ((N1C + N001C).ult(RHS: InnerShiftSize)) {
2146 MatchInfo.Src = ShiftMI.getOperand(i: 1).getReg();
2147 MatchInfo.ShiftAmt = N1C + N001C;
2148 MatchInfo.ShiftAmtTy = MRI.getType(Reg: N001);
2149 MatchInfo.InnerShiftTy = InnerShiftTy;
2150
2151 if ((N001C + OpSizeInBits) == InnerShiftSize)
2152 return true;
2153 if (MRI.hasOneUse(RegNo: N0) && MRI.hasOneUse(RegNo: InnerShift)) {
2154 MatchInfo.Mask = true;
2155 MatchInfo.MaskVal = APInt(N1C.getBitWidth(), OpSizeInBits) - N1C;
2156 return true;
2157 }
2158 }
2159 return false;
2160}
2161
2162void CombinerHelper::applyLshrOfTruncOfLshr(
2163 MachineInstr &MI, LshrOfTruncOfLshr &MatchInfo) const {
2164 assert(MI.getOpcode() == TargetOpcode::G_LSHR && "Expected a G_LSHR");
2165
2166 Register Dst = MI.getOperand(i: 0).getReg();
2167 auto ShiftAmt =
2168 Builder.buildConstant(Res: MatchInfo.ShiftAmtTy, Val: MatchInfo.ShiftAmt);
2169 auto Shift =
2170 Builder.buildLShr(Dst: MatchInfo.InnerShiftTy, Src0: MatchInfo.Src, Src1: ShiftAmt);
2171 if (MatchInfo.Mask == true) {
2172 APInt MaskVal =
2173 APInt::getLowBitsSet(numBits: MatchInfo.InnerShiftTy.getScalarSizeInBits(),
2174 loBitsSet: MatchInfo.MaskVal.getZExtValue());
2175 auto Mask = Builder.buildConstant(Res: MatchInfo.InnerShiftTy, Val: MaskVal);
2176 auto And = Builder.buildAnd(Dst: MatchInfo.InnerShiftTy, Src0: Shift, Src1: Mask);
2177 Builder.buildTrunc(Res: Dst, Op: And);
2178 } else
2179 Builder.buildTrunc(Res: Dst, Op: Shift);
2180 MI.eraseFromParent();
2181}
2182
2183bool CombinerHelper::matchCombineMulToShl(MachineInstr &MI,
2184 unsigned &ShiftVal) const {
2185 assert(MI.getOpcode() == TargetOpcode::G_MUL && "Expected a G_MUL");
2186 auto MaybeImmVal =
2187 getIConstantVRegValWithLookThrough(VReg: MI.getOperand(i: 2).getReg(), MRI);
2188 if (!MaybeImmVal)
2189 return false;
2190
2191 ShiftVal = MaybeImmVal->Value.exactLogBase2();
2192 return (static_cast<int32_t>(ShiftVal) != -1);
2193}
2194
2195void CombinerHelper::applyCombineMulToShl(MachineInstr &MI,
2196 unsigned &ShiftVal) const {
2197 assert(MI.getOpcode() == TargetOpcode::G_MUL && "Expected a G_MUL");
2198 MachineIRBuilder MIB(MI);
2199 LLT ShiftTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
2200 auto ShiftCst = MIB.buildConstant(Res: ShiftTy, Val: ShiftVal);
2201 Observer.changingInstr(MI);
2202 MI.setDesc(MIB.getTII().get(Opcode: TargetOpcode::G_SHL));
2203 MI.getOperand(i: 2).setReg(ShiftCst.getReg(Idx: 0));
2204 if (ShiftVal == ShiftTy.getScalarSizeInBits() - 1)
2205 MI.clearFlag(Flag: MachineInstr::MIFlag::NoSWrap);
2206 Observer.changedInstr(MI);
2207}
2208
2209bool CombinerHelper::matchCombineSubToAdd(MachineInstr &MI,
2210 BuildFnTy &MatchInfo) const {
2211 GSub &Sub = cast<GSub>(Val&: MI);
2212
2213 LLT Ty = MRI.getType(Reg: Sub.getReg(Idx: 0));
2214
2215 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_ADD, {Ty}}))
2216 return false;
2217
2218 if (!isConstantLegalOrBeforeLegalizer(Ty))
2219 return false;
2220
2221 APInt Imm = getIConstantFromReg(VReg: Sub.getRHSReg(), MRI);
2222
2223 MatchInfo = [=, &MI](MachineIRBuilder &B) {
2224 auto NegCst = B.buildConstant(Res: Ty, Val: -Imm);
2225 Observer.changingInstr(MI);
2226 MI.setDesc(B.getTII().get(Opcode: TargetOpcode::G_ADD));
2227 MI.getOperand(i: 2).setReg(NegCst.getReg(Idx: 0));
2228 MI.clearFlag(Flag: MachineInstr::MIFlag::NoUWrap);
2229 if (Imm.isMinSignedValue())
2230 MI.clearFlags(flags: MachineInstr::MIFlag::NoSWrap);
2231 Observer.changedInstr(MI);
2232 };
2233 return true;
2234}
2235
2236// shl ([sza]ext x), y => zext (shl x, y), if shift does not overflow source
2237bool CombinerHelper::matchCombineShlOfExtend(MachineInstr &MI,
2238 RegisterImmPair &MatchData) const {
2239 assert(MI.getOpcode() == TargetOpcode::G_SHL && VT);
2240 if (!getTargetLowering().isDesirableToPullExtFromShl(MI))
2241 return false;
2242
2243 Register LHS = MI.getOperand(i: 1).getReg();
2244
2245 Register ExtSrc;
2246 if (!mi_match(R: LHS, MRI, P: m_GAnyExt(Src: m_Reg(R&: ExtSrc))) &&
2247 !mi_match(R: LHS, MRI, P: m_GZExt(Src: m_Reg(R&: ExtSrc))) &&
2248 !mi_match(R: LHS, MRI, P: m_GSExt(Src: m_Reg(R&: ExtSrc))))
2249 return false;
2250
2251 Register RHS = MI.getOperand(i: 2).getReg();
2252 auto MaybeShiftAmtVal = isConstantOrConstantSplatVector(Def: RHS, MRI);
2253 if (!MaybeShiftAmtVal)
2254 return false;
2255
2256 if (LI) {
2257 LLT SrcTy = MRI.getType(Reg: ExtSrc);
2258
2259 // We only really care about the legality with the shifted value. We can
2260 // pick any type the constant shift amount, so ask the target what to
2261 // use. Otherwise we would have to guess and hope it is reported as legal.
2262 LLT ShiftAmtTy = getTargetLowering().getPreferredShiftAmountTy(ShiftValueTy: SrcTy);
2263 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_SHL, {SrcTy, ShiftAmtTy}}))
2264 return false;
2265 }
2266
2267 int64_t ShiftAmt = MaybeShiftAmtVal->getSExtValue();
2268 MatchData.Reg = ExtSrc;
2269 MatchData.Imm = ShiftAmt;
2270
2271 unsigned MinLeadingZeros = VT->getKnownZeroes(R: ExtSrc).countl_one();
2272 unsigned SrcTySize = MRI.getType(Reg: ExtSrc).getScalarSizeInBits();
2273 return MinLeadingZeros >= ShiftAmt && ShiftAmt < SrcTySize;
2274}
2275
2276void CombinerHelper::applyCombineShlOfExtend(
2277 MachineInstr &MI, const RegisterImmPair &MatchData) const {
2278 Register ExtSrcReg = MatchData.Reg;
2279 int64_t ShiftAmtVal = MatchData.Imm;
2280
2281 LLT ExtSrcTy = MRI.getType(Reg: ExtSrcReg);
2282 auto ShiftAmt = Builder.buildConstant(Res: ExtSrcTy, Val: ShiftAmtVal);
2283 auto NarrowShift =
2284 Builder.buildShl(Dst: ExtSrcTy, Src0: ExtSrcReg, Src1: ShiftAmt, Flags: MI.getFlags());
2285 Builder.buildZExt(Res: MI.getOperand(i: 0), Op: NarrowShift);
2286 MI.eraseFromParent();
2287}
2288
2289static Register peekThroughBitcast(Register Reg,
2290 const MachineRegisterInfo &MRI) {
2291 while (mi_match(R: Reg, MRI, P: m_GBitcast(Src: m_Reg(R&: Reg))))
2292 ;
2293
2294 return Reg;
2295}
2296
2297bool CombinerHelper::matchCombineUnmergeMergeToPlainValues(
2298 MachineInstr &MI, SmallVectorImpl<Register> &Operands) const {
2299 assert(MI.getOpcode() == TargetOpcode::G_UNMERGE_VALUES &&
2300 "Expected an unmerge");
2301 auto &Unmerge = cast<GUnmerge>(Val&: MI);
2302 Register SrcReg = peekThroughBitcast(Reg: Unmerge.getSourceReg(), MRI);
2303
2304 auto *SrcInstr = getOpcodeDef<GMergeLikeInstr>(Reg: SrcReg, MRI);
2305 if (!SrcInstr)
2306 return false;
2307
2308 // Check the source type of the merge.
2309 LLT SrcMergeTy = MRI.getType(Reg: SrcInstr->getSourceReg(I: 0));
2310 LLT Dst0Ty = MRI.getType(Reg: Unmerge.getReg(Idx: 0));
2311 bool SameSize = Dst0Ty.getSizeInBits() == SrcMergeTy.getSizeInBits();
2312 if (SrcMergeTy != Dst0Ty && !SameSize)
2313 return false;
2314 // They are the same now (modulo a bitcast).
2315 // We can collect all the src registers.
2316 for (unsigned Idx = 0; Idx < SrcInstr->getNumSources(); ++Idx)
2317 Operands.push_back(Elt: SrcInstr->getSourceReg(I: Idx));
2318 return true;
2319}
2320
2321void CombinerHelper::applyCombineUnmergeMergeToPlainValues(
2322 MachineInstr &MI, SmallVectorImpl<Register> &Operands) const {
2323 assert(MI.getOpcode() == TargetOpcode::G_UNMERGE_VALUES &&
2324 "Expected an unmerge");
2325 assert((MI.getNumOperands() - 1 == Operands.size()) &&
2326 "Not enough operands to replace all defs");
2327 unsigned NumElems = MI.getNumOperands() - 1;
2328
2329 LLT SrcTy = MRI.getType(Reg: Operands[0]);
2330 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
2331 bool CanReuseInputDirectly = DstTy == SrcTy;
2332 for (unsigned Idx = 0; Idx < NumElems; ++Idx) {
2333 Register DstReg = MI.getOperand(i: Idx).getReg();
2334 Register SrcReg = Operands[Idx];
2335
2336 // This combine may run after RegBankSelect, so we need to be aware of
2337 // register banks.
2338 const auto &DstCB = MRI.getRegClassOrRegBank(Reg: DstReg);
2339 if (!DstCB.isNull() && DstCB != MRI.getRegClassOrRegBank(Reg: SrcReg)) {
2340 SrcReg = Builder.buildCopy(Res: MRI.getType(Reg: SrcReg), Op: SrcReg).getReg(Idx: 0);
2341 MRI.setRegClassOrRegBank(Reg: SrcReg, RCOrRB: DstCB);
2342 }
2343
2344 if (CanReuseInputDirectly)
2345 replaceRegWith(MRI, FromReg: DstReg, ToReg: SrcReg);
2346 else
2347 Builder.buildCast(Dst: DstReg, Src: SrcReg);
2348 }
2349 MI.eraseFromParent();
2350}
2351
2352bool CombinerHelper::matchCombineUnmergeConstant(
2353 MachineInstr &MI, SmallVectorImpl<APInt> &Csts) const {
2354 unsigned SrcIdx = MI.getNumOperands() - 1;
2355 Register SrcReg = MI.getOperand(i: SrcIdx).getReg();
2356 // Break down the big constant in smaller ones.
2357 APInt Val;
2358 if (!mi_match(R: SrcReg, MRI, P: m_GConstantOrFConstantBits(Bits&: Val)))
2359 return false;
2360
2361 LLT Dst0Ty = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
2362 unsigned ShiftAmt = Dst0Ty.getSizeInBits();
2363 // Unmerge a constant.
2364 for (unsigned Idx = 0; Idx != SrcIdx; ++Idx) {
2365 Csts.emplace_back(Args: Val.trunc(width: ShiftAmt));
2366 Val = Val.lshr(shiftAmt: ShiftAmt);
2367 }
2368
2369 return true;
2370}
2371
2372void CombinerHelper::applyCombineUnmergeConstant(
2373 MachineInstr &MI, SmallVectorImpl<APInt> &Csts) const {
2374 assert(MI.getOpcode() == TargetOpcode::G_UNMERGE_VALUES &&
2375 "Expected an unmerge");
2376 assert((MI.getNumOperands() - 1 == Csts.size()) &&
2377 "Not enough operands to replace all defs");
2378 unsigned NumElems = MI.getNumOperands() - 1;
2379 for (unsigned Idx = 0; Idx < NumElems; ++Idx) {
2380 Register DstReg = MI.getOperand(i: Idx).getReg();
2381 Builder.buildConstant(Res: DstReg, Val: Csts[Idx]);
2382 }
2383
2384 MI.eraseFromParent();
2385}
2386
2387bool CombinerHelper::matchCombineUnmergeUndef(
2388 MachineInstr &MI,
2389 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
2390 unsigned SrcIdx = MI.getNumOperands() - 1;
2391 Register SrcReg = MI.getOperand(i: SrcIdx).getReg();
2392 MatchInfo = [&MI](MachineIRBuilder &B) {
2393 unsigned NumElems = MI.getNumOperands() - 1;
2394 for (unsigned Idx = 0; Idx < NumElems; ++Idx) {
2395 Register DstReg = MI.getOperand(i: Idx).getReg();
2396 B.buildUndef(Res: DstReg);
2397 }
2398 };
2399 return mi_match(R: SrcReg, MRI, P: m_GImplicitDef());
2400}
2401
2402bool CombinerHelper::matchCombineUnmergeWithDeadLanesToTrunc(
2403 MachineInstr &MI) const {
2404 assert(MI.getOpcode() == TargetOpcode::G_UNMERGE_VALUES &&
2405 "Expected an unmerge");
2406 if (!MRI.getType(Reg: MI.getOperand(i: 0).getReg()).isScalar() ||
2407 !MRI.getType(Reg: MI.getOperand(i: MI.getNumDefs()).getReg()).isScalar())
2408 return false;
2409 // Check that all the lanes are dead except the first one.
2410 for (unsigned Idx = 1, EndIdx = MI.getNumDefs(); Idx != EndIdx; ++Idx) {
2411 if (!MRI.use_nodbg_empty(RegNo: MI.getOperand(i: Idx).getReg()))
2412 return false;
2413 }
2414 return true;
2415}
2416
2417void CombinerHelper::applyCombineUnmergeWithDeadLanesToTrunc(
2418 MachineInstr &MI) const {
2419 Register SrcReg = MI.getOperand(i: MI.getNumDefs()).getReg();
2420 Register Dst0Reg = MI.getOperand(i: 0).getReg();
2421 Builder.buildTrunc(Res: Dst0Reg, Op: SrcReg);
2422 MI.eraseFromParent();
2423}
2424
2425bool CombinerHelper::matchCombineUnmergeZExtToZExt(MachineInstr &MI) const {
2426 assert(MI.getOpcode() == TargetOpcode::G_UNMERGE_VALUES &&
2427 "Expected an unmerge");
2428 Register Dst0Reg = MI.getOperand(i: 0).getReg();
2429 LLT Dst0Ty = MRI.getType(Reg: Dst0Reg);
2430 // G_ZEXT on vector applies to each lane, so it will
2431 // affect all destinations. Therefore we won't be able
2432 // to simplify the unmerge to just the first definition.
2433 if (Dst0Ty.isVector())
2434 return false;
2435 Register SrcReg = MI.getOperand(i: MI.getNumDefs()).getReg();
2436 LLT SrcTy = MRI.getType(Reg: SrcReg);
2437 if (SrcTy.isVector())
2438 return false;
2439
2440 Register ZExtSrcReg;
2441 if (!mi_match(R: SrcReg, MRI, P: m_GZExt(Src: m_Reg(R&: ZExtSrcReg))))
2442 return false;
2443
2444 // Finally we can replace the first definition with
2445 // a zext of the source if the definition is big enough to hold
2446 // all of ZExtSrc bits.
2447 LLT ZExtSrcTy = MRI.getType(Reg: ZExtSrcReg);
2448 return ZExtSrcTy.getSizeInBits() <= Dst0Ty.getSizeInBits();
2449}
2450
2451void CombinerHelper::applyCombineUnmergeZExtToZExt(MachineInstr &MI) const {
2452 assert(MI.getOpcode() == TargetOpcode::G_UNMERGE_VALUES &&
2453 "Expected an unmerge");
2454
2455 Register Dst0Reg = MI.getOperand(i: 0).getReg();
2456
2457 GZext *ZExtInstr =
2458 cast<GZext>(Val: MRI.getVRegDef(Reg: MI.getOperand(i: MI.getNumDefs()).getReg()));
2459 Register ZExtSrcReg = ZExtInstr->getSrcReg();
2460 LLT Dst0Ty = MRI.getType(Reg: Dst0Reg);
2461 LLT ZExtSrcTy = MRI.getType(Reg: ZExtSrcReg);
2462
2463 if (Dst0Ty.getSizeInBits() > ZExtSrcTy.getSizeInBits()) {
2464 Builder.buildZExt(Res: Dst0Reg, Op: ZExtSrcReg);
2465 } else {
2466 assert(Dst0Ty.getSizeInBits() == ZExtSrcTy.getSizeInBits() &&
2467 "ZExt src doesn't fit in destination");
2468 replaceRegWith(MRI, FromReg: Dst0Reg, ToReg: ZExtSrcReg);
2469 }
2470
2471 Register ZeroReg;
2472 for (unsigned Idx = 1, EndIdx = MI.getNumDefs(); Idx != EndIdx; ++Idx) {
2473 if (!ZeroReg)
2474 ZeroReg = Builder.buildConstant(Res: Dst0Ty, Val: 0).getReg(Idx: 0);
2475 replaceRegWith(MRI, FromReg: MI.getOperand(i: Idx).getReg(), ToReg: ZeroReg);
2476 }
2477 MI.eraseFromParent();
2478}
2479
2480bool CombinerHelper::matchCombineShiftToUnmerge(MachineInstr &MI,
2481 unsigned TargetShiftSize,
2482 unsigned &ShiftVal) const {
2483 assert((MI.getOpcode() == TargetOpcode::G_SHL ||
2484 MI.getOpcode() == TargetOpcode::G_LSHR ||
2485 MI.getOpcode() == TargetOpcode::G_ASHR) && "Expected a shift");
2486
2487 LLT Ty = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
2488 if (Ty.isVector()) // TODO:
2489 return false;
2490
2491 // Don't narrow further than the requested size.
2492 unsigned Size = Ty.getSizeInBits();
2493 if (Size <= TargetShiftSize)
2494 return false;
2495
2496 auto MaybeImmVal =
2497 getIConstantVRegValWithLookThrough(VReg: MI.getOperand(i: 2).getReg(), MRI);
2498 if (!MaybeImmVal)
2499 return false;
2500
2501 ShiftVal = MaybeImmVal->Value.getSExtValue();
2502 return ShiftVal >= Size / 2 && ShiftVal < Size;
2503}
2504
2505void CombinerHelper::applyCombineShiftToUnmerge(
2506 MachineInstr &MI, const unsigned &ShiftVal) const {
2507 Register DstReg = MI.getOperand(i: 0).getReg();
2508 Register SrcReg = MI.getOperand(i: 1).getReg();
2509 LLT Ty = MRI.getType(Reg: SrcReg);
2510 unsigned Size = Ty.getSizeInBits();
2511 unsigned HalfSize = Size / 2;
2512 assert(ShiftVal >= HalfSize);
2513
2514 LLT HalfTy = Ty.changeElementSize(NewEltSize: HalfSize);
2515
2516 auto Unmerge = Builder.buildUnmerge(Res: HalfTy, Op: SrcReg);
2517 unsigned NarrowShiftAmt = ShiftVal - HalfSize;
2518
2519 if (MI.getOpcode() == TargetOpcode::G_LSHR) {
2520 Register Narrowed = Unmerge.getReg(Idx: 1);
2521
2522 // dst = G_LSHR s64:x, C for C >= 32
2523 // =>
2524 // lo, hi = G_UNMERGE_VALUES x
2525 // dst = G_MERGE_VALUES (G_LSHR hi, C - 32), 0
2526
2527 if (NarrowShiftAmt != 0) {
2528 Narrowed = Builder.buildLShr(Dst: HalfTy, Src0: Narrowed,
2529 Src1: Builder.buildConstant(Res: HalfTy, Val: NarrowShiftAmt)).getReg(Idx: 0);
2530 }
2531
2532 auto Zero = Builder.buildConstant(Res: HalfTy, Val: 0);
2533 Builder.buildMergeLikeInstr(Res: DstReg, Ops: {Narrowed, Zero});
2534 } else if (MI.getOpcode() == TargetOpcode::G_SHL) {
2535 Register Narrowed = Unmerge.getReg(Idx: 0);
2536 // dst = G_SHL s64:x, C for C >= 32
2537 // =>
2538 // lo, hi = G_UNMERGE_VALUES x
2539 // dst = G_MERGE_VALUES 0, (G_SHL hi, C - 32)
2540 if (NarrowShiftAmt != 0) {
2541 Narrowed = Builder.buildShl(Dst: HalfTy, Src0: Narrowed,
2542 Src1: Builder.buildConstant(Res: HalfTy, Val: NarrowShiftAmt)).getReg(Idx: 0);
2543 }
2544
2545 auto Zero = Builder.buildConstant(Res: HalfTy, Val: 0);
2546 Builder.buildMergeLikeInstr(Res: DstReg, Ops: {Zero, Narrowed});
2547 } else {
2548 assert(MI.getOpcode() == TargetOpcode::G_ASHR);
2549 auto Hi = Builder.buildAShr(
2550 Dst: HalfTy, Src0: Unmerge.getReg(Idx: 1),
2551 Src1: Builder.buildConstant(Res: HalfTy, Val: HalfSize - 1));
2552
2553 if (ShiftVal == HalfSize) {
2554 // (G_ASHR i64:x, 32) ->
2555 // G_MERGE_VALUES hi_32(x), (G_ASHR hi_32(x), 31)
2556 Builder.buildMergeLikeInstr(Res: DstReg, Ops: {Unmerge.getReg(Idx: 1), Hi});
2557 } else if (ShiftVal == Size - 1) {
2558 // Don't need a second shift.
2559 // (G_ASHR i64:x, 63) ->
2560 // %narrowed = (G_ASHR hi_32(x), 31)
2561 // G_MERGE_VALUES %narrowed, %narrowed
2562 Builder.buildMergeLikeInstr(Res: DstReg, Ops: {Hi, Hi});
2563 } else {
2564 auto Lo = Builder.buildAShr(
2565 Dst: HalfTy, Src0: Unmerge.getReg(Idx: 1),
2566 Src1: Builder.buildConstant(Res: HalfTy, Val: ShiftVal - HalfSize));
2567
2568 // (G_ASHR i64:x, C) ->, for C >= 32
2569 // G_MERGE_VALUES (G_ASHR hi_32(x), C - 32), (G_ASHR hi_32(x), 31)
2570 Builder.buildMergeLikeInstr(Res: DstReg, Ops: {Lo, Hi});
2571 }
2572 }
2573
2574 MI.eraseFromParent();
2575}
2576
2577bool CombinerHelper::tryCombineShiftToUnmerge(
2578 MachineInstr &MI, unsigned TargetShiftAmount) const {
2579 unsigned ShiftAmt;
2580 if (matchCombineShiftToUnmerge(MI, TargetShiftSize: TargetShiftAmount, ShiftVal&: ShiftAmt)) {
2581 applyCombineShiftToUnmerge(MI, ShiftVal: ShiftAmt);
2582 return true;
2583 }
2584
2585 return false;
2586}
2587
2588void CombinerHelper::applyCombineP2IToI2P(MachineInstr &MI,
2589 Register &Reg) const {
2590 assert(MI.getOpcode() == TargetOpcode::G_PTRTOINT && "Expected a G_PTRTOINT");
2591 Register DstReg = MI.getOperand(i: 0).getReg();
2592 Builder.buildZExtOrTrunc(Res: DstReg, Op: Reg);
2593 MI.eraseFromParent();
2594}
2595
2596bool CombinerHelper::matchCombineAnyExtTrunc(MachineInstr &MI,
2597 Register &Reg) const {
2598 assert(MI.getOpcode() == TargetOpcode::G_ANYEXT && "Expected a G_ANYEXT");
2599 Register DstReg = MI.getOperand(i: 0).getReg();
2600 Register SrcReg = MI.getOperand(i: 1).getReg();
2601 Register OriginalSrcReg = getSrcRegIgnoringCopies(Reg: SrcReg, MRI);
2602 if (OriginalSrcReg.isValid())
2603 SrcReg = OriginalSrcReg;
2604 LLT DstTy = MRI.getType(Reg: DstReg);
2605 return mi_match(R: SrcReg, MRI,
2606 P: m_GTrunc(Src: m_all_of(preds: m_Reg(R&: Reg), preds: m_SpecificType(Ty: DstTy)))) &&
2607 canReplaceReg(DstReg, SrcReg: Reg, MRI);
2608}
2609
2610bool CombinerHelper::matchCombineZextTrunc(MachineInstr &MI,
2611 Register &Reg) const {
2612 assert(MI.getOpcode() == TargetOpcode::G_ZEXT && "Expected a G_ZEXT");
2613 Register DstReg = MI.getOperand(i: 0).getReg();
2614 Register SrcReg = MI.getOperand(i: 1).getReg();
2615 LLT DstTy = MRI.getType(Reg: DstReg);
2616 if (mi_match(R: SrcReg, MRI,
2617 P: m_GTrunc(Src: m_all_of(preds: m_Reg(R&: Reg), preds: m_SpecificType(Ty: DstTy)))) &&
2618 canReplaceReg(DstReg, SrcReg: Reg, MRI)) {
2619 unsigned DstSize = DstTy.getScalarSizeInBits();
2620 unsigned SrcSize = MRI.getType(Reg: SrcReg).getScalarSizeInBits();
2621 return VT->getKnownBits(R: Reg).countMinLeadingZeros() >= DstSize - SrcSize;
2622 }
2623 return false;
2624}
2625
2626static LLT getMidVTForTruncRightShiftCombine(LLT ShiftTy, LLT TruncTy) {
2627 const unsigned ShiftSize = ShiftTy.getScalarSizeInBits();
2628 const unsigned TruncSize = TruncTy.getScalarSizeInBits();
2629
2630 // ShiftTy > 32 > TruncTy -> 32
2631 if (ShiftSize > 32 && TruncSize < 32)
2632 return ShiftTy.changeElementSize(NewEltSize: 32);
2633
2634 // TODO: We could also reduce to 16 bits, but that's more target-dependent.
2635 // Some targets like it, some don't, some only like it under certain
2636 // conditions/processor versions, etc.
2637 // A TL hook might be needed for this.
2638
2639 // Don't combine
2640 return ShiftTy;
2641}
2642
2643bool CombinerHelper::matchCombineTruncOfShift(
2644 MachineInstr &MI, std::pair<MachineInstr *, LLT> &MatchInfo) const {
2645 assert(MI.getOpcode() == TargetOpcode::G_TRUNC && "Expected a G_TRUNC");
2646 Register DstReg = MI.getOperand(i: 0).getReg();
2647 Register SrcReg = MI.getOperand(i: 1).getReg();
2648
2649 if (!MRI.hasOneNonDBGUse(RegNo: SrcReg))
2650 return false;
2651
2652 LLT SrcTy = MRI.getType(Reg: SrcReg);
2653 LLT DstTy = MRI.getType(Reg: DstReg);
2654
2655 MachineInstr *SrcMI = getDefIgnoringCopies(Reg: SrcReg, MRI);
2656 const auto &TL = getTargetLowering();
2657
2658 LLT NewShiftTy;
2659 switch (SrcMI->getOpcode()) {
2660 default:
2661 return false;
2662 case TargetOpcode::G_SHL: {
2663 NewShiftTy = DstTy;
2664
2665 // Make sure new shift amount is legal.
2666 KnownBits Known = VT->getKnownBits(R: SrcMI->getOperand(i: 2).getReg());
2667 if (Known.getMaxValue().uge(RHS: NewShiftTy.getScalarSizeInBits()))
2668 return false;
2669 break;
2670 }
2671 case TargetOpcode::G_LSHR:
2672 case TargetOpcode::G_ASHR: {
2673 // For right shifts, we conservatively do not do the transform if the TRUNC
2674 // has any STORE users. The reason is that if we change the type of the
2675 // shift, we may break the truncstore combine.
2676 //
2677 // TODO: Fix truncstore combine to handle (trunc(lshr (trunc x), k)).
2678 for (auto &User : MRI.use_instructions(Reg: DstReg))
2679 if (User.getOpcode() == TargetOpcode::G_STORE)
2680 return false;
2681
2682 NewShiftTy = getMidVTForTruncRightShiftCombine(ShiftTy: SrcTy, TruncTy: DstTy);
2683 if (NewShiftTy == SrcTy)
2684 return false;
2685
2686 // Make sure we won't lose information by truncating the high bits.
2687 KnownBits Known = VT->getKnownBits(R: SrcMI->getOperand(i: 2).getReg());
2688 if (Known.getMaxValue().ugt(RHS: NewShiftTy.getScalarSizeInBits() -
2689 DstTy.getScalarSizeInBits()))
2690 return false;
2691 break;
2692 }
2693 }
2694
2695 if (!isLegalOrBeforeLegalizer(
2696 Query: {SrcMI->getOpcode(),
2697 {NewShiftTy, TL.getPreferredShiftAmountTy(ShiftValueTy: NewShiftTy)}}))
2698 return false;
2699
2700 MatchInfo = std::make_pair(x&: SrcMI, y&: NewShiftTy);
2701 return true;
2702}
2703
2704void CombinerHelper::applyCombineTruncOfShift(
2705 MachineInstr &MI, std::pair<MachineInstr *, LLT> &MatchInfo) const {
2706 MachineInstr *ShiftMI = MatchInfo.first;
2707 LLT NewShiftTy = MatchInfo.second;
2708
2709 Register Dst = MI.getOperand(i: 0).getReg();
2710 LLT DstTy = MRI.getType(Reg: Dst);
2711
2712 Register ShiftAmt = ShiftMI->getOperand(i: 2).getReg();
2713 Register ShiftSrc = ShiftMI->getOperand(i: 1).getReg();
2714 ShiftSrc = Builder.buildTrunc(Res: NewShiftTy, Op: ShiftSrc).getReg(Idx: 0);
2715
2716 const auto &TL = getTargetLowering();
2717 LLT PrefShiftTy = TL.getPreferredShiftAmountTy(ShiftValueTy: NewShiftTy);
2718 if (MRI.getType(Reg: ShiftAmt) != PrefShiftTy)
2719 ShiftAmt = Builder.buildZExtOrTrunc(Res: PrefShiftTy, Op: ShiftAmt).getReg(Idx: 0);
2720
2721 Register NewShift =
2722 Builder
2723 .buildInstr(Opc: ShiftMI->getOpcode(), DstOps: {NewShiftTy}, SrcOps: {ShiftSrc, ShiftAmt})
2724 .getReg(Idx: 0);
2725
2726 if (NewShiftTy == DstTy)
2727 replaceRegWith(MRI, FromReg: Dst, ToReg: NewShift);
2728 else
2729 Builder.buildTrunc(Res: Dst, Op: NewShift);
2730
2731 eraseInst(MI);
2732}
2733
2734bool CombinerHelper::matchAllExplicitUsesAreUndef(MachineInstr &MI) const {
2735 return all_of(Range: MI.explicit_uses(), P: [this](const MachineOperand &MO) {
2736 return !MO.isReg() ||
2737 getOpcodeDef(Opcode: TargetOpcode::G_IMPLICIT_DEF, Reg: MO.getReg(), MRI);
2738 });
2739}
2740
2741bool CombinerHelper::matchUndefShuffleVectorMask(MachineInstr &MI) const {
2742 assert(MI.getOpcode() == TargetOpcode::G_SHUFFLE_VECTOR);
2743 ArrayRef<int> Mask = MI.getOperand(i: 3).getShuffleMask();
2744 return all_of(Range&: Mask, P: [](int Elt) { return Elt < 0; });
2745}
2746
2747bool CombinerHelper::matchUndefStore(MachineInstr &MI) const {
2748 assert(MI.getOpcode() == TargetOpcode::G_STORE);
2749 return getOpcodeDef(Opcode: TargetOpcode::G_IMPLICIT_DEF, Reg: MI.getOperand(i: 0).getReg(),
2750 MRI);
2751}
2752
2753bool CombinerHelper::matchInsertExtractVecEltOutOfBounds(
2754 MachineInstr &MI) const {
2755 assert((MI.getOpcode() == TargetOpcode::G_INSERT_VECTOR_ELT ||
2756 MI.getOpcode() == TargetOpcode::G_EXTRACT_VECTOR_ELT) &&
2757 "Expected an insert/extract element op");
2758 LLT VecTy = MRI.getType(Reg: MI.getOperand(i: 1).getReg());
2759 if (VecTy.isScalableVector())
2760 return false;
2761
2762 unsigned IdxIdx =
2763 MI.getOpcode() == TargetOpcode::G_EXTRACT_VECTOR_ELT ? 2 : 3;
2764 auto Idx = getIConstantVRegVal(VReg: MI.getOperand(i: IdxIdx).getReg(), MRI);
2765 if (!Idx)
2766 return false;
2767 return Idx->getZExtValue() >= VecTy.getNumElements();
2768}
2769
2770bool CombinerHelper::matchConstantSelectCmp(MachineInstr &MI,
2771 unsigned &OpIdx) const {
2772 GSelect &SelMI = cast<GSelect>(Val&: MI);
2773 auto Cst = isConstantOrConstantSplatVector(Def: SelMI.getCondReg(), MRI);
2774 if (!Cst)
2775 return false;
2776 OpIdx = Cst->isZero() ? 3 : 2;
2777 return true;
2778}
2779
2780void CombinerHelper::eraseInst(MachineInstr &MI) const { MI.eraseFromParent(); }
2781
2782bool CombinerHelper::matchEqualDefs(const MachineOperand &MOP1,
2783 const MachineOperand &MOP2) const {
2784 if (!MOP1.isReg() || !MOP2.isReg())
2785 return false;
2786 auto InstAndDef1 = getDefSrcRegIgnoringCopies(Reg: MOP1.getReg(), MRI);
2787 if (!InstAndDef1)
2788 return false;
2789 auto InstAndDef2 = getDefSrcRegIgnoringCopies(Reg: MOP2.getReg(), MRI);
2790 if (!InstAndDef2)
2791 return false;
2792 MachineInstr *I1 = InstAndDef1->MI;
2793 MachineInstr *I2 = InstAndDef2->MI;
2794
2795 // Handle a case like this:
2796 //
2797 // %0:_(s64), %1:_(s64) = G_UNMERGE_VALUES %2:_(<2 x s64>)
2798 //
2799 // Even though %0 and %1 are produced by the same instruction they are not
2800 // the same values.
2801 if (I1 == I2)
2802 return MOP1.getReg() == MOP2.getReg();
2803
2804 // If we have an instruction which loads or stores, we can't guarantee that
2805 // it is identical.
2806 //
2807 // For example, we may have
2808 //
2809 // %x1 = G_LOAD %addr (load N from @somewhere)
2810 // ...
2811 // call @foo
2812 // ...
2813 // %x2 = G_LOAD %addr (load N from @somewhere)
2814 // ...
2815 // %or = G_OR %x1, %x2
2816 //
2817 // It's possible that @foo will modify whatever lives at the address we're
2818 // loading from. To be safe, let's just assume that all loads and stores
2819 // are different (unless we have something which is guaranteed to not
2820 // change.)
2821 if (I1->mayLoadOrStore() && !I1->isDereferenceableInvariantLoad())
2822 return false;
2823
2824 // If both instructions are loads or stores, they are equal only if both
2825 // are dereferenceable invariant loads with the same number of bits.
2826 if (I1->mayLoadOrStore() && I2->mayLoadOrStore()) {
2827 GLoadStore *LS1 = dyn_cast<GLoadStore>(Val: I1);
2828 GLoadStore *LS2 = dyn_cast<GLoadStore>(Val: I2);
2829 if (!LS1 || !LS2)
2830 return false;
2831
2832 if (!I2->isDereferenceableInvariantLoad() ||
2833 (LS1->getMemSizeInBits() != LS2->getMemSizeInBits()))
2834 return false;
2835 }
2836
2837 // Check for physical registers on the instructions first to avoid cases
2838 // like this:
2839 //
2840 // %a = COPY $physreg
2841 // ...
2842 // SOMETHING implicit-def $physreg
2843 // ...
2844 // %b = COPY $physreg
2845 //
2846 // These copies are not equivalent.
2847 if (any_of(Range: I1->uses(), P: [](const MachineOperand &MO) {
2848 return MO.isReg() && MO.getReg().isPhysical();
2849 })) {
2850 // Check if we have a case like this:
2851 //
2852 // %a = COPY $physreg
2853 // %b = COPY %a
2854 //
2855 // In this case, I1 and I2 will both be equal to %a = COPY $physreg.
2856 // From that, we know that they must have the same value, since they must
2857 // have come from the same COPY.
2858 return I1->isIdenticalTo(Other: *I2);
2859 }
2860
2861 // We don't have any physical registers, so we don't necessarily need the
2862 // same vreg defs.
2863 //
2864 // On the off-chance that there's some target instruction feeding into the
2865 // instruction, let's use produceSameValue instead of isIdenticalTo.
2866 if (Builder.getTII().produceSameValue(MI0: *I1, MI1: *I2, MRI: &MRI)) {
2867 // Handle instructions with multiple defs that produce same values. Values
2868 // are same for operands with same index.
2869 // %0:_(s8), %1:_(s8), %2:_(s8), %3:_(s8) = G_UNMERGE_VALUES %4:_(<4 x s8>)
2870 // %5:_(s8), %6:_(s8), %7:_(s8), %8:_(s8) = G_UNMERGE_VALUES %4:_(<4 x s8>)
2871 // I1 and I2 are different instructions but produce same values,
2872 // %1 and %6 are same, %1 and %7 are not the same value.
2873 return I1->findRegisterDefOperandIdx(Reg: InstAndDef1->Reg, /*TRI=*/nullptr) ==
2874 I2->findRegisterDefOperandIdx(Reg: InstAndDef2->Reg, /*TRI=*/nullptr);
2875 }
2876 return false;
2877}
2878
2879bool CombinerHelper::matchConstantFPOp(const MachineOperand &MOP,
2880 double C) const {
2881 if (!MOP.isReg())
2882 return false;
2883 std::optional<FPValueAndVReg> MaybeCst;
2884 if (!mi_match(R: MOP.getReg(), MRI, P: m_GFCstOrSplat(FPValReg&: MaybeCst)))
2885 return false;
2886
2887 return MaybeCst->Value.isExactlyValue(V: C);
2888}
2889
2890void CombinerHelper::replaceSingleDefInstWithOperand(MachineInstr &MI,
2891 unsigned OpIdx) const {
2892 assert(MI.getNumExplicitDefs() == 1 && "Expected one explicit def?");
2893 Register OldReg = MI.getOperand(i: 0).getReg();
2894 Register Replacement = MI.getOperand(i: OpIdx).getReg();
2895 assert(canReplaceReg(OldReg, Replacement, MRI) && "Cannot replace register?");
2896 replaceRegWith(MRI, FromReg: OldReg, ToReg: Replacement);
2897 MI.eraseFromParent();
2898}
2899
2900void CombinerHelper::replaceSingleDefInstWithReg(MachineInstr &MI,
2901 Register Replacement) const {
2902 assert(MI.getNumExplicitDefs() == 1 && "Expected one explicit def?");
2903 Register OldReg = MI.getOperand(i: 0).getReg();
2904 assert(canReplaceReg(OldReg, Replacement, MRI) && "Cannot replace register?");
2905 replaceRegWith(MRI, FromReg: OldReg, ToReg: Replacement);
2906 MI.eraseFromParent();
2907}
2908
2909bool CombinerHelper::matchConstantLargerBitWidth(MachineInstr &MI,
2910 unsigned ConstIdx) const {
2911 Register ConstReg = MI.getOperand(i: ConstIdx).getReg();
2912 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
2913
2914 // Get the shift amount
2915 auto VRegAndVal = getIConstantVRegValWithLookThrough(VReg: ConstReg, MRI);
2916 if (!VRegAndVal)
2917 return false;
2918
2919 // Return true of shift amount >= Bitwidth
2920 return (VRegAndVal->Value.uge(RHS: DstTy.getSizeInBits()));
2921}
2922
2923void CombinerHelper::applyFunnelShiftConstantModulo(MachineInstr &MI) const {
2924 assert((MI.getOpcode() == TargetOpcode::G_FSHL ||
2925 MI.getOpcode() == TargetOpcode::G_FSHR) &&
2926 "This is not a funnel shift operation");
2927
2928 Register ConstReg = MI.getOperand(i: 3).getReg();
2929 LLT ConstTy = MRI.getType(Reg: ConstReg);
2930 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
2931
2932 auto VRegAndVal = getIConstantVRegValWithLookThrough(VReg: ConstReg, MRI);
2933 assert((VRegAndVal) && "Value is not a constant");
2934
2935 // Calculate the new Shift Amount = Old Shift Amount % BitWidth
2936 APInt NewConst = VRegAndVal->Value.urem(
2937 RHS: APInt(ConstTy.getSizeInBits(), DstTy.getScalarSizeInBits()));
2938
2939 auto NewConstInstr = Builder.buildConstant(Res: ConstTy, Val: NewConst.getZExtValue());
2940 Builder.buildInstr(
2941 Opc: MI.getOpcode(), DstOps: {MI.getOperand(i: 0)},
2942 SrcOps: {MI.getOperand(i: 1), MI.getOperand(i: 2), NewConstInstr.getReg(Idx: 0)});
2943
2944 MI.eraseFromParent();
2945}
2946
2947bool CombinerHelper::matchSelectSameVal(MachineInstr &MI) const {
2948 assert(MI.getOpcode() == TargetOpcode::G_SELECT);
2949 // Match (cond ? x : x)
2950 return matchEqualDefs(MOP1: MI.getOperand(i: 2), MOP2: MI.getOperand(i: 3)) &&
2951 canReplaceReg(DstReg: MI.getOperand(i: 0).getReg(), SrcReg: MI.getOperand(i: 2).getReg(),
2952 MRI);
2953}
2954
2955bool CombinerHelper::matchOperandIsKnownToBeAPowerOfTwo(
2956 const MachineOperand &MO, bool OrNegative) const {
2957 return isKnownToBeAPowerOfTwo(Val: MO.getReg(), MRI, ValueTracking: VT, OrNegative);
2958}
2959
2960void CombinerHelper::replaceInstWithFConstant(MachineInstr &MI,
2961 double C) const {
2962 assert(MI.getNumDefs() == 1 && "Expected only one def?");
2963 Builder.buildFConstant(Res: MI.getOperand(i: 0), Val: C);
2964 MI.eraseFromParent();
2965}
2966
2967void CombinerHelper::replaceInstWithConstant(MachineInstr &MI,
2968 int64_t C) const {
2969 assert(MI.getNumDefs() == 1 && "Expected only one def?");
2970 Builder.buildConstant(Res: MI.getOperand(i: 0), Val: C);
2971 MI.eraseFromParent();
2972}
2973
2974void CombinerHelper::replaceInstWithConstant(MachineInstr &MI, APInt C) const {
2975 assert(MI.getNumDefs() == 1 && "Expected only one def?");
2976 Builder.buildConstant(Res: MI.getOperand(i: 0), Val: C);
2977 MI.eraseFromParent();
2978}
2979
2980void CombinerHelper::replaceInstWithFConstant(MachineInstr &MI,
2981 ConstantFP *CFP) const {
2982 assert(MI.getNumDefs() == 1 && "Expected only one def?");
2983 Builder.buildFConstant(Res: MI.getOperand(i: 0), Val: CFP->getValueAPF());
2984 MI.eraseFromParent();
2985}
2986
2987void CombinerHelper::replaceInstWithUndef(MachineInstr &MI) const {
2988 assert(MI.getNumDefs() == 1 && "Expected only one def?");
2989 Builder.buildUndef(Res: MI.getOperand(i: 0));
2990 MI.eraseFromParent();
2991}
2992
2993bool CombinerHelper::matchCombineInsertVecElts(
2994 MachineInstr &MI, SmallVectorImpl<Register> &MatchInfo) const {
2995 assert(MI.getOpcode() == TargetOpcode::G_INSERT_VECTOR_ELT &&
2996 "Invalid opcode");
2997 Register DstReg = MI.getOperand(i: 0).getReg();
2998 LLT DstTy = MRI.getType(Reg: DstReg);
2999 assert(DstTy.isVector() && "Invalid G_INSERT_VECTOR_ELT?");
3000
3001 if (DstTy.isScalableVector())
3002 return false;
3003
3004 unsigned NumElts = DstTy.getNumElements();
3005 // If this MI is part of a sequence of insert_vec_elts, then
3006 // don't do the combine in the middle of the sequence.
3007 if (MRI.hasOneUse(RegNo: DstReg) && MRI.use_instr_begin(RegNo: DstReg)->getOpcode() ==
3008 TargetOpcode::G_INSERT_VECTOR_ELT)
3009 return false;
3010 MachineInstr *CurrInst = &MI;
3011 MachineInstr *TmpInst;
3012 int64_t IntImm;
3013 Register TmpReg;
3014 MatchInfo.resize(N: NumElts);
3015 while (mi_match(
3016 MI&: *CurrInst, MRI,
3017 P: m_GInsertVecElt(Src0: m_MInstr(MI&: TmpInst), Src1: m_Reg(R&: TmpReg), Src2: m_ICst(Cst&: IntImm)))) {
3018 if (IntImm >= NumElts || IntImm < 0)
3019 return false;
3020 if (!MatchInfo[IntImm])
3021 MatchInfo[IntImm] = TmpReg;
3022 CurrInst = TmpInst;
3023 }
3024 // Variable index.
3025 if (CurrInst->getOpcode() == TargetOpcode::G_INSERT_VECTOR_ELT)
3026 return false;
3027 if (TmpInst->getOpcode() == TargetOpcode::G_BUILD_VECTOR) {
3028 for (unsigned I = 1; I < TmpInst->getNumOperands(); ++I) {
3029 if (!MatchInfo[I - 1].isValid())
3030 MatchInfo[I - 1] = TmpInst->getOperand(i: I).getReg();
3031 }
3032 return true;
3033 }
3034 // If we didn't end in a G_IMPLICIT_DEF and the source is not fully
3035 // overwritten, bail out.
3036 return TmpInst->getOpcode() == TargetOpcode::G_IMPLICIT_DEF ||
3037 all_of(Range&: MatchInfo, P: [](Register Reg) { return !!Reg; });
3038}
3039
3040void CombinerHelper::applyCombineInsertVecElts(
3041 MachineInstr &MI, SmallVectorImpl<Register> &MatchInfo) const {
3042 Register UndefReg;
3043 auto GetUndef = [&]() {
3044 if (UndefReg)
3045 return UndefReg;
3046 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
3047 UndefReg = Builder.buildUndef(Res: DstTy.getScalarType()).getReg(Idx: 0);
3048 return UndefReg;
3049 };
3050 for (Register &Reg : MatchInfo) {
3051 if (!Reg)
3052 Reg = GetUndef();
3053 }
3054 Builder.buildBuildVector(Res: MI.getOperand(i: 0).getReg(), Ops: MatchInfo);
3055 MI.eraseFromParent();
3056}
3057
3058bool CombinerHelper::matchBinopWithNegInner(Register MInner, Register Other,
3059 unsigned RootOpc, Register Dst,
3060 LLT Ty,
3061 BuildFnTy &MatchInfo) const {
3062 /// Helper function for matchBinopWithNeg: tries to match one commuted form
3063 /// of `a bitwiseop (~b +/- c)` -> `a bitwiseop ~(b -/+ c)`.
3064 MachineInstr *InnerDef;
3065 if (!mi_match(R: MInner, MRI, P: m_MInstr(MI&: InnerDef)))
3066 return false;
3067
3068 unsigned InnerOpc = InnerDef->getOpcode();
3069 if (InnerOpc != TargetOpcode::G_ADD && InnerOpc != TargetOpcode::G_SUB)
3070 return false;
3071
3072 if (!MRI.hasOneNonDBGUse(RegNo: MInner))
3073 return false;
3074
3075 Register InnerLHS = InnerDef->getOperand(i: 1).getReg();
3076 Register InnerRHS = InnerDef->getOperand(i: 2).getReg();
3077 Register NotSrc;
3078 Register B, C;
3079
3080 // Check if either operand is ~b
3081 auto TryMatch = [&](Register MaybeNot, Register Other) {
3082 if (mi_match(R: MaybeNot, MRI, P: m_Not(Src: m_Reg(R&: NotSrc)))) {
3083 if (!MRI.hasOneNonDBGUse(RegNo: MaybeNot))
3084 return false;
3085 B = NotSrc;
3086 C = Other;
3087 return true;
3088 }
3089 return false;
3090 };
3091
3092 // For SUB, the not must be the LHS. For ADD, it can be either operand.
3093 if (!TryMatch(InnerLHS, InnerRHS) &&
3094 !(InnerOpc == TargetOpcode::G_ADD && TryMatch(InnerRHS, InnerLHS)))
3095 return false;
3096
3097 // Flip add/sub
3098 unsigned FlippedOpc = (InnerOpc == TargetOpcode::G_ADD) ? TargetOpcode::G_SUB
3099 : TargetOpcode::G_ADD;
3100
3101 Register A = Other;
3102 MatchInfo = [=](MachineIRBuilder &Builder) {
3103 auto NewInner = Builder.buildInstr(Opc: FlippedOpc, DstOps: {Ty}, SrcOps: {B, C});
3104 auto NewNot = Builder.buildNot(Dst: Ty, Src0: NewInner);
3105 Builder.buildInstr(Opc: RootOpc, DstOps: {Dst}, SrcOps: {A, NewNot});
3106 };
3107 return true;
3108}
3109
3110bool CombinerHelper::matchBinopWithNeg(MachineInstr &MI,
3111 BuildFnTy &MatchInfo) const {
3112 // Fold `a bitwiseop (~b +/- c)` -> `a bitwiseop ~(b -/+ c)`
3113 // Root MI is one of G_AND, G_OR, G_XOR.
3114 // We also look for commuted forms of operations. Pattern shouldn't apply
3115 // if there are multiple reasons of inner operations.
3116
3117 unsigned RootOpc = MI.getOpcode();
3118 Register Dst = MI.getOperand(i: 0).getReg();
3119 LLT Ty = MRI.getType(Reg: Dst);
3120
3121 Register LHS = MI.getOperand(i: 1).getReg();
3122 Register RHS = MI.getOperand(i: 2).getReg();
3123 // Check the commuted and uncommuted forms of the operation.
3124 return matchBinopWithNegInner(MInner: LHS, Other: RHS, RootOpc, Dst, Ty, MatchInfo) ||
3125 matchBinopWithNegInner(MInner: RHS, Other: LHS, RootOpc, Dst, Ty, MatchInfo);
3126}
3127
3128bool CombinerHelper::matchHoistLogicOpWithSameOpcodeHands(
3129 MachineInstr &MI, InstructionStepsMatchInfo &MatchInfo) const {
3130 // Matches: logic (hand x, ...), (hand y, ...) -> hand (logic x, y), ...
3131 //
3132 // Creates the new hand + logic instruction (but does not insert them.)
3133 //
3134 // On success, MatchInfo is populated with the new instructions. These are
3135 // inserted in applyHoistLogicOpWithSameOpcodeHands.
3136 unsigned LogicOpcode = MI.getOpcode();
3137 assert(LogicOpcode == TargetOpcode::G_AND ||
3138 LogicOpcode == TargetOpcode::G_OR ||
3139 LogicOpcode == TargetOpcode::G_XOR);
3140 MachineIRBuilder MIB(MI);
3141 Register Dst = MI.getOperand(i: 0).getReg();
3142 Register LHSReg = MI.getOperand(i: 1).getReg();
3143 Register RHSReg = MI.getOperand(i: 2).getReg();
3144
3145 // Don't recompute anything.
3146 if (!MRI.hasOneNonDBGUse(RegNo: LHSReg) || !MRI.hasOneNonDBGUse(RegNo: RHSReg))
3147 return false;
3148
3149 // Make sure we have (hand x, ...), (hand y, ...)
3150 MachineInstr *LeftHandInst = getDefIgnoringCopies(Reg: LHSReg, MRI);
3151 MachineInstr *RightHandInst = getDefIgnoringCopies(Reg: RHSReg, MRI);
3152 if (!LeftHandInst || !RightHandInst)
3153 return false;
3154 unsigned HandOpcode = LeftHandInst->getOpcode();
3155 if (HandOpcode != RightHandInst->getOpcode())
3156 return false;
3157 if (LeftHandInst->getNumOperands() < 2 ||
3158 !LeftHandInst->getOperand(i: 1).isReg() ||
3159 RightHandInst->getNumOperands() < 2 ||
3160 !RightHandInst->getOperand(i: 1).isReg())
3161 return false;
3162
3163 // Make sure the types match up, and if we're doing this post-legalization,
3164 // we end up with legal types.
3165 Register X = LeftHandInst->getOperand(i: 1).getReg();
3166 Register Y = RightHandInst->getOperand(i: 1).getReg();
3167 LLT XTy = MRI.getType(Reg: X);
3168 LLT YTy = MRI.getType(Reg: Y);
3169 if (!XTy.isValid() || XTy != YTy)
3170 return false;
3171
3172 // Optional extra source register.
3173 Register ExtraHandOpSrcReg;
3174 switch (HandOpcode) {
3175 default:
3176 return false;
3177 case TargetOpcode::G_ANYEXT:
3178 case TargetOpcode::G_SEXT:
3179 case TargetOpcode::G_ZEXT: {
3180 // Match: logic (ext X), (ext Y) --> ext (logic X, Y)
3181 break;
3182 }
3183 case TargetOpcode::G_TRUNC: {
3184 // Match: logic (trunc X), (trunc Y) -> trunc (logic X, Y)
3185 const MachineFunction *MF = MI.getMF();
3186 LLVMContext &Ctx = MF->getFunction().getContext();
3187
3188 LLT DstTy = MRI.getType(Reg: Dst);
3189 const TargetLowering &TLI = getTargetLowering();
3190
3191 // Be extra careful sinking truncate. If it's free, there's no benefit in
3192 // widening a binop.
3193 if (TLI.isZExtFree(FromTy: DstTy, ToTy: XTy, Ctx) && TLI.isTruncateFree(FromTy: XTy, ToTy: DstTy, Ctx))
3194 return false;
3195 break;
3196 }
3197 case TargetOpcode::G_AND:
3198 case TargetOpcode::G_ASHR:
3199 case TargetOpcode::G_LSHR:
3200 case TargetOpcode::G_SHL: {
3201 // Match: logic (binop x, z), (binop y, z) -> binop (logic x, y), z
3202 MachineOperand &ZOp = LeftHandInst->getOperand(i: 2);
3203 if (!matchEqualDefs(MOP1: ZOp, MOP2: RightHandInst->getOperand(i: 2)))
3204 return false;
3205 ExtraHandOpSrcReg = ZOp.getReg();
3206 break;
3207 }
3208 }
3209
3210 if (!isLegalOrBeforeLegalizer(Query: {LogicOpcode, {XTy, YTy}}))
3211 return false;
3212
3213 // Record the steps to build the new instructions.
3214 //
3215 // Steps to build (logic x, y)
3216 auto NewLogicDst = MRI.createGenericVirtualRegister(Ty: XTy);
3217 OperandBuildSteps LogicBuildSteps = {
3218 [=](MachineInstrBuilder &MIB) { MIB.addDef(RegNo: NewLogicDst); },
3219 [=](MachineInstrBuilder &MIB) { MIB.addReg(RegNo: X); },
3220 [=](MachineInstrBuilder &MIB) { MIB.addReg(RegNo: Y); }};
3221 InstructionBuildSteps LogicSteps(LogicOpcode, LogicBuildSteps);
3222
3223 // Steps to build hand (logic x, y), ...z
3224 OperandBuildSteps HandBuildSteps = {
3225 [=](MachineInstrBuilder &MIB) { MIB.addDef(RegNo: Dst); },
3226 [=](MachineInstrBuilder &MIB) { MIB.addReg(RegNo: NewLogicDst); }};
3227 if (ExtraHandOpSrcReg.isValid())
3228 HandBuildSteps.push_back(
3229 Elt: [=](MachineInstrBuilder &MIB) { MIB.addReg(RegNo: ExtraHandOpSrcReg); });
3230 InstructionBuildSteps HandSteps(HandOpcode, HandBuildSteps);
3231
3232 MatchInfo = InstructionStepsMatchInfo({LogicSteps, HandSteps});
3233 return true;
3234}
3235
3236void CombinerHelper::applyBuildInstructionSteps(
3237 MachineInstr &MI, InstructionStepsMatchInfo &MatchInfo) const {
3238 assert(MatchInfo.InstrsToBuild.size() &&
3239 "Expected at least one instr to build?");
3240 for (auto &InstrToBuild : MatchInfo.InstrsToBuild) {
3241 assert(InstrToBuild.Opcode && "Expected a valid opcode?");
3242 assert(InstrToBuild.OperandFns.size() && "Expected at least one operand?");
3243 MachineInstrBuilder Instr = Builder.buildInstr(Opcode: InstrToBuild.Opcode);
3244 for (auto &OperandFn : InstrToBuild.OperandFns)
3245 OperandFn(Instr);
3246 }
3247 MI.eraseFromParent();
3248}
3249
3250bool CombinerHelper::matchAshrShlToSextInreg(
3251 MachineInstr &MI, std::tuple<Register, int64_t> &MatchInfo) const {
3252 assert(MI.getOpcode() == TargetOpcode::G_ASHR);
3253 int64_t ShlCst, AshrCst;
3254 Register Src;
3255 if (!mi_match(R: MI.getOperand(i: 0).getReg(), MRI,
3256 P: m_GAShr(L: m_GShl(L: m_Reg(R&: Src), R: m_ICstOrSplat(Cst&: ShlCst)),
3257 R: m_ICstOrSplat(Cst&: AshrCst))))
3258 return false;
3259 if (ShlCst != AshrCst)
3260 return false;
3261 if (!isLegalOrBeforeLegalizer(
3262 Query: {TargetOpcode::G_SEXT_INREG,
3263 {MRI.getType(Reg: Src)},
3264 {},
3265 {MRI.getType(Reg: Src).getScalarSizeInBits() - ShlCst}}))
3266 return false;
3267 MatchInfo = std::make_tuple(args&: Src, args&: ShlCst);
3268 return true;
3269}
3270
3271void CombinerHelper::applyAshShlToSextInreg(
3272 MachineInstr &MI, std::tuple<Register, int64_t> &MatchInfo) const {
3273 assert(MI.getOpcode() == TargetOpcode::G_ASHR);
3274 Register Src;
3275 int64_t ShiftAmt;
3276 std::tie(args&: Src, args&: ShiftAmt) = MatchInfo;
3277 unsigned Size = MRI.getType(Reg: Src).getScalarSizeInBits();
3278 Builder.buildSExtInReg(Res: MI.getOperand(i: 0).getReg(), Op: Src, ImmOp: Size - ShiftAmt);
3279 MI.eraseFromParent();
3280}
3281
3282/// and(and(x, C1), C2) -> C1&C2 ? and(x, C1&C2) : 0
3283bool CombinerHelper::matchOverlappingAnd(
3284 MachineInstr &MI,
3285 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
3286 assert(MI.getOpcode() == TargetOpcode::G_AND);
3287
3288 Register Dst = MI.getOperand(i: 0).getReg();
3289 LLT Ty = MRI.getType(Reg: Dst);
3290
3291 Register R;
3292 int64_t C1;
3293 int64_t C2;
3294 if (!mi_match(
3295 R: Dst, MRI,
3296 P: m_GAnd(L: m_GAnd(L: m_Reg(R), R: m_ICst(Cst&: C1)), R: m_ICst(Cst&: C2))))
3297 return false;
3298
3299 MatchInfo = [=](MachineIRBuilder &B) {
3300 if (C1 & C2) {
3301 B.buildAnd(Dst, Src0: R, Src1: B.buildConstant(Res: Ty, Val: C1 & C2));
3302 return;
3303 }
3304 auto Zero = B.buildConstant(Res: Ty, Val: 0);
3305 replaceRegWith(MRI, FromReg: Dst, ToReg: Zero->getOperand(i: 0).getReg());
3306 };
3307 return true;
3308}
3309
3310bool CombinerHelper::matchRedundantAnd(MachineInstr &MI,
3311 Register &Replacement) const {
3312 // Given
3313 //
3314 // %y:_(sN) = G_SOMETHING
3315 // %x:_(sN) = G_SOMETHING
3316 // %res:_(sN) = G_AND %x, %y
3317 //
3318 // Eliminate the G_AND when it is known that x & y == x or x & y == y.
3319 //
3320 // Patterns like this can appear as a result of legalization. E.g.
3321 //
3322 // %cmp:_(s32) = G_ICMP intpred(pred), %x(s32), %y
3323 // %one:_(s32) = G_CONSTANT i32 1
3324 // %and:_(s32) = G_AND %cmp, %one
3325 //
3326 // In this case, G_ICMP only produces a single bit, so x & 1 == x.
3327 assert(MI.getOpcode() == TargetOpcode::G_AND);
3328 if (!VT)
3329 return false;
3330
3331 Register AndDst = MI.getOperand(i: 0).getReg();
3332 Register LHS = MI.getOperand(i: 1).getReg();
3333 Register RHS = MI.getOperand(i: 2).getReg();
3334
3335 // Check the RHS (maybe a constant) first, and if we have no KnownBits there,
3336 // we can't do anything. If we do, then it depends on whether we have
3337 // KnownBits on the LHS.
3338 KnownBits RHSBits = VT->getKnownBits(R: RHS);
3339 if (RHSBits.isUnknown())
3340 return false;
3341
3342 KnownBits LHSBits = VT->getKnownBits(R: LHS);
3343
3344 // Check that x & Mask == x.
3345 // x & 1 == x, always
3346 // x & 0 == x, only if x is also 0
3347 // Meaning Mask has no effect if every bit is either one in Mask or zero in x.
3348 //
3349 // Check if we can replace AndDst with the LHS of the G_AND
3350 if (canReplaceReg(DstReg: AndDst, SrcReg: LHS, MRI) &&
3351 (LHSBits.Zero | RHSBits.One).isAllOnes()) {
3352 Replacement = LHS;
3353 return true;
3354 }
3355
3356 // Check if we can replace AndDst with the RHS of the G_AND
3357 if (canReplaceReg(DstReg: AndDst, SrcReg: RHS, MRI) &&
3358 (LHSBits.One | RHSBits.Zero).isAllOnes()) {
3359 Replacement = RHS;
3360 return true;
3361 }
3362
3363 return false;
3364}
3365
3366bool CombinerHelper::matchRedundantOr(MachineInstr &MI,
3367 Register &Replacement) const {
3368 // Given
3369 //
3370 // %y:_(sN) = G_SOMETHING
3371 // %x:_(sN) = G_SOMETHING
3372 // %res:_(sN) = G_OR %x, %y
3373 //
3374 // Eliminate the G_OR when it is known that x | y == x or x | y == y.
3375 assert(MI.getOpcode() == TargetOpcode::G_OR);
3376 if (!VT)
3377 return false;
3378
3379 Register OrDst = MI.getOperand(i: 0).getReg();
3380 Register LHS = MI.getOperand(i: 1).getReg();
3381 Register RHS = MI.getOperand(i: 2).getReg();
3382
3383 KnownBits LHSBits = VT->getKnownBits(R: LHS);
3384 KnownBits RHSBits = VT->getKnownBits(R: RHS);
3385
3386 // Check that x | Mask == x.
3387 // x | 0 == x, always
3388 // x | 1 == x, only if x is also 1
3389 // Meaning Mask has no effect if every bit is either zero in Mask or one in x.
3390 //
3391 // Check if we can replace OrDst with the LHS of the G_OR
3392 if (canReplaceReg(DstReg: OrDst, SrcReg: LHS, MRI) &&
3393 (LHSBits.One | RHSBits.Zero).isAllOnes()) {
3394 Replacement = LHS;
3395 return true;
3396 }
3397
3398 // Check if we can replace OrDst with the RHS of the G_OR
3399 if (canReplaceReg(DstReg: OrDst, SrcReg: RHS, MRI) &&
3400 (LHSBits.Zero | RHSBits.One).isAllOnes()) {
3401 Replacement = RHS;
3402 return true;
3403 }
3404
3405 return false;
3406}
3407
3408bool CombinerHelper::matchRedundantSExtInReg(MachineInstr &MI) const {
3409 // If the input is already sign extended, just drop the extension.
3410 Register Src = MI.getOperand(i: 1).getReg();
3411 unsigned ExtBits = MI.getOperand(i: 2).getImm();
3412 unsigned TypeSize = MRI.getType(Reg: Src).getScalarSizeInBits();
3413 return VT->computeNumSignBits(R: Src) >= (TypeSize - ExtBits + 1);
3414}
3415
3416static bool isConstValidTrue(const TargetLowering &TLI, unsigned ScalarSizeBits,
3417 int64_t Cst, bool IsVector, bool IsFP) {
3418 // For i1, Cst will always be -1 regardless of boolean contents.
3419 return (ScalarSizeBits == 1 && Cst == -1) ||
3420 isConstTrueVal(TLI, Val: Cst, IsVector, IsFP);
3421}
3422
3423// This pattern aims to match the following shape to avoid extra mov
3424// instructions
3425// G_BUILD_VECTOR(
3426// G_UNMERGE_VALUES(src, 0)
3427// G_UNMERGE_VALUES(src, 1)
3428// G_IMPLICIT_DEF
3429// G_IMPLICIT_DEF
3430// )
3431// ->
3432// G_CONCAT_VECTORS(
3433// src,
3434// undef
3435// )
3436bool CombinerHelper::matchCombineBuildUnmerge(MachineInstr &MI,
3437 MachineRegisterInfo &MRI,
3438 Register &UnmergeSrc) const {
3439 auto &BV = cast<GBuildVector>(Val&: MI);
3440
3441 unsigned BuildUseCount = BV.getNumSources();
3442 if (BuildUseCount % 2 != 0)
3443 return false;
3444
3445 unsigned NumUnmerge = BuildUseCount / 2;
3446
3447 auto *Unmerge = getOpcodeDef<GUnmerge>(Reg: BV.getSourceReg(I: 0), MRI);
3448
3449 // Check the first operand is an unmerge and has the correct number of
3450 // operands
3451 if (!Unmerge || Unmerge->getNumDefs() != NumUnmerge)
3452 return false;
3453
3454 UnmergeSrc = Unmerge->getSourceReg();
3455
3456 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
3457 LLT UnmergeSrcTy = MRI.getType(Reg: UnmergeSrc);
3458
3459 if (!UnmergeSrcTy.isVector())
3460 return false;
3461
3462 // Ensure we only generate legal instructions post-legalizer
3463 if (!IsPreLegalize &&
3464 !isLegal(Query: {TargetOpcode::G_CONCAT_VECTORS, {DstTy, UnmergeSrcTy}}))
3465 return false;
3466
3467 // Check that all of the operands before the midpoint come from the same
3468 // unmerge and are in the same order as they are used in the build_vector
3469 for (unsigned I = 0; I < NumUnmerge; ++I) {
3470 auto MaybeUnmergeReg = BV.getSourceReg(I);
3471 auto *LoopUnmerge = getOpcodeDef<GUnmerge>(Reg: MaybeUnmergeReg, MRI);
3472
3473 if (!LoopUnmerge || LoopUnmerge != Unmerge)
3474 return false;
3475
3476 if (LoopUnmerge->getOperand(i: I).getReg() != MaybeUnmergeReg)
3477 return false;
3478 }
3479
3480 // Check that all of the unmerged values are used
3481 if (Unmerge->getNumDefs() != NumUnmerge)
3482 return false;
3483
3484 // Check that all of the operands after the mid point are undefs.
3485 for (unsigned I = NumUnmerge; I < BuildUseCount; ++I) {
3486 auto *Undef = getDefIgnoringCopies(Reg: BV.getSourceReg(I), MRI);
3487
3488 if (Undef->getOpcode() != TargetOpcode::G_IMPLICIT_DEF)
3489 return false;
3490 }
3491
3492 return true;
3493}
3494
3495void CombinerHelper::applyCombineBuildUnmerge(MachineInstr &MI,
3496 MachineRegisterInfo &MRI,
3497 MachineIRBuilder &B,
3498 Register &UnmergeSrc) const {
3499 assert(UnmergeSrc && "Expected there to be one matching G_UNMERGE_VALUES");
3500 B.setInstrAndDebugLoc(MI);
3501
3502 Register UndefVec = B.buildUndef(Res: MRI.getType(Reg: UnmergeSrc)).getReg(Idx: 0);
3503 B.buildConcatVectors(Res: MI.getOperand(i: 0), Ops: {UnmergeSrc, UndefVec});
3504
3505 MI.eraseFromParent();
3506}
3507
3508// This combine tries to reduce the number of scalarised G_TRUNC instructions by
3509// using vector truncates instead
3510//
3511// EXAMPLE:
3512// %a(i32), %b(i32) = G_UNMERGE_VALUES %src(<2 x i32>)
3513// %T_a(i16) = G_TRUNC %a(i32)
3514// %T_b(i16) = G_TRUNC %b(i32)
3515// %Undef(i16) = G_IMPLICIT_DEF(i16)
3516// %dst(v4i16) = G_BUILD_VECTORS %T_a(i16), %T_b(i16), %Undef(i16), %Undef(i16)
3517//
3518// ===>
3519// %Undef(<2 x i32>) = G_IMPLICIT_DEF(<2 x i32>)
3520// %Mid(<4 x s32>) = G_CONCAT_VECTORS %src(<2 x i32>), %Undef(<2 x i32>)
3521// %dst(<4 x s16>) = G_TRUNC %Mid(<4 x s32>)
3522//
3523// Only matches sources made up of G_TRUNCs followed by G_IMPLICIT_DEFs
3524bool CombinerHelper::matchUseVectorTruncate(MachineInstr &MI,
3525 Register &MatchInfo) const {
3526 auto BuildMI = cast<GBuildVector>(Val: &MI);
3527 unsigned NumOperands = BuildMI->getNumSources();
3528 LLT DstTy = MRI.getType(Reg: BuildMI->getReg(Idx: 0));
3529
3530 // Check the G_BUILD_VECTOR sources
3531 unsigned I;
3532 GUnmerge *UnmergeMI = nullptr;
3533
3534 // Check all source TRUNCs come from the same UNMERGE instruction
3535 // and that the element order matches (BUILD_VECTOR position I
3536 // corresponds to UNMERGE result I)
3537 for (I = 0; I < NumOperands; ++I) {
3538 // Check if the G_TRUNC instructions all come from the same MI
3539 Register TruncSrcReg;
3540 if (!mi_match(R: BuildMI->getSourceReg(I), MRI, P: m_GTrunc(Src: m_Reg(R&: TruncSrcReg))))
3541 break;
3542
3543 if (!UnmergeMI) {
3544 if (!mi_match(R: TruncSrcReg, MRI, P: m_GUnmerge(Inst&: UnmergeMI)))
3545 return false;
3546 } else {
3547 MachineInstr *UnmergeSrcMI;
3548 if (!mi_match(R: TruncSrcReg, MRI, P: m_MInstr(MI&: UnmergeSrcMI)) ||
3549 UnmergeMI != UnmergeSrcMI)
3550 return false;
3551 }
3552 // Element order must match: position I must use UNMERGE result I.
3553 if (UnmergeMI->getOperand(i: I).getReg() != TruncSrcReg)
3554 return false;
3555 }
3556 if (I < 2)
3557 return false;
3558
3559 // Check the remaining source elements are only G_IMPLICIT_DEF
3560 for (; I < NumOperands; ++I) {
3561 if (!mi_match(R: BuildMI->getSourceReg(I), MRI, P: m_GImplicitDef()))
3562 return false;
3563 }
3564
3565 // Check the size of unmerge source
3566 MatchInfo = UnmergeMI->getSourceReg();
3567 LLT UnmergeSrcTy = MRI.getType(Reg: MatchInfo);
3568 if (!DstTy.getElementCount().isKnownMultipleOf(RHS: UnmergeSrcTy.getNumElements()))
3569 return false;
3570
3571 // Check the unmerge source and destination element types match
3572 LLT UnmergeSrcEltTy = UnmergeSrcTy.getElementType();
3573 Register UnmergeDstReg = UnmergeMI->getOperand(i: 0).getReg();
3574 LLT UnmergeDstEltTy = MRI.getType(Reg: UnmergeDstReg);
3575 if (UnmergeSrcEltTy != UnmergeDstEltTy)
3576 return false;
3577
3578 // Only generate legal instructions post-legalizer
3579 if (!IsPreLegalize) {
3580 LLT MidTy = DstTy.changeElementType(NewEltTy: UnmergeSrcTy.getScalarType());
3581
3582 if (DstTy.getElementCount() != UnmergeSrcTy.getElementCount() &&
3583 !isLegal(Query: {TargetOpcode::G_CONCAT_VECTORS, {MidTy, UnmergeSrcTy}}))
3584 return false;
3585
3586 if (!isLegal(Query: {TargetOpcode::G_TRUNC, {DstTy, MidTy}}))
3587 return false;
3588 }
3589
3590 return true;
3591}
3592
3593void CombinerHelper::applyUseVectorTruncate(MachineInstr &MI,
3594 Register &MatchInfo) const {
3595 Register MidReg;
3596 auto BuildMI = cast<GBuildVector>(Val: &MI);
3597 Register DstReg = BuildMI->getReg(Idx: 0);
3598 LLT DstTy = MRI.getType(Reg: DstReg);
3599 LLT UnmergeSrcTy = MRI.getType(Reg: MatchInfo);
3600 unsigned DstTyNumElt = DstTy.getNumElements();
3601 unsigned UnmergeSrcTyNumElt = UnmergeSrcTy.getNumElements();
3602
3603 // No need to pad vector if only G_TRUNC is needed
3604 if (DstTyNumElt / UnmergeSrcTyNumElt == 1) {
3605 MidReg = MatchInfo;
3606 } else {
3607 Register UndefReg = Builder.buildUndef(Res: UnmergeSrcTy).getReg(Idx: 0);
3608 SmallVector<Register> ConcatRegs = {MatchInfo};
3609 for (unsigned I = 1; I < DstTyNumElt / UnmergeSrcTyNumElt; ++I)
3610 ConcatRegs.push_back(Elt: UndefReg);
3611
3612 auto MidTy = DstTy.changeElementType(NewEltTy: UnmergeSrcTy.getScalarType());
3613 MidReg = Builder.buildConcatVectors(Res: MidTy, Ops: ConcatRegs).getReg(Idx: 0);
3614 }
3615
3616 Builder.buildTrunc(Res: DstReg, Op: MidReg);
3617 MI.eraseFromParent();
3618}
3619
3620bool CombinerHelper::matchNotCmp(
3621 MachineInstr &MI, SmallVectorImpl<Register> &RegsToNegate) const {
3622 assert(MI.getOpcode() == TargetOpcode::G_XOR);
3623 LLT Ty = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
3624 const auto &TLI = *Builder.getMF().getSubtarget().getTargetLowering();
3625 Register XorSrc;
3626 Register CstReg;
3627 // We match xor(src, true) here.
3628 if (!mi_match(R: MI.getOperand(i: 0).getReg(), MRI,
3629 P: m_GXor(L: m_Reg(R&: XorSrc), R: m_Reg(R&: CstReg))))
3630 return false;
3631
3632 if (!MRI.hasOneNonDBGUse(RegNo: XorSrc))
3633 return false;
3634
3635 // Check that XorSrc is the root of a tree of comparisons combined with ANDs
3636 // and ORs. The suffix of RegsToNegate starting from index I is used a work
3637 // list of tree nodes to visit.
3638 RegsToNegate.push_back(Elt: XorSrc);
3639 // Remember whether the comparisons are all integer or all floating point.
3640 bool IsInt = false;
3641 bool IsFP = false;
3642 for (unsigned I = 0; I < RegsToNegate.size(); ++I) {
3643 Register Reg = RegsToNegate[I];
3644 if (!MRI.hasOneNonDBGUse(RegNo: Reg))
3645 return false;
3646 MachineInstr *Def;
3647 if (!mi_match(R: Reg, MRI, P: m_MInstr(MI&: Def)))
3648 return false;
3649 switch (Def->getOpcode()) {
3650 default:
3651 // Don't match if the tree contains anything other than ANDs, ORs and
3652 // comparisons.
3653 return false;
3654 case TargetOpcode::G_ICMP:
3655 if (IsFP)
3656 return false;
3657 IsInt = true;
3658 // When we apply the combine we will invert the predicate.
3659 break;
3660 case TargetOpcode::G_FCMP:
3661 if (IsInt)
3662 return false;
3663 IsFP = true;
3664 // When we apply the combine we will invert the predicate.
3665 break;
3666 case TargetOpcode::G_AND:
3667 case TargetOpcode::G_OR:
3668 // Implement De Morgan's laws:
3669 // ~(x & y) -> ~x | ~y
3670 // ~(x | y) -> ~x & ~y
3671 // When we apply the combine we will change the opcode and recursively
3672 // negate the operands.
3673 RegsToNegate.push_back(Elt: Def->getOperand(i: 1).getReg());
3674 RegsToNegate.push_back(Elt: Def->getOperand(i: 2).getReg());
3675 break;
3676 }
3677 }
3678
3679 // Now we know whether the comparisons are integer or floating point, check
3680 // the constant in the xor.
3681 int64_t Cst;
3682 if (Ty.isVector()) {
3683 int64_t SplatCst;
3684 if (!mi_match(R: CstReg, MRI, P: m_ICstOrSplat(Cst&: SplatCst)))
3685 return false;
3686 if (!isConstValidTrue(TLI, ScalarSizeBits: Ty.getScalarSizeInBits(), Cst: SplatCst, IsVector: true, IsFP))
3687 return false;
3688 } else {
3689 if (!mi_match(R: CstReg, MRI, P: m_ICst(Cst)))
3690 return false;
3691 if (!isConstValidTrue(TLI, ScalarSizeBits: Ty.getSizeInBits(), Cst, IsVector: false, IsFP))
3692 return false;
3693 }
3694
3695 return true;
3696}
3697
3698void CombinerHelper::applyNotCmp(
3699 MachineInstr &MI, SmallVectorImpl<Register> &RegsToNegate) const {
3700 for (Register Reg : RegsToNegate) {
3701 MachineInstr *Def = MRI.getVRegDef(Reg);
3702 Observer.changingInstr(MI&: *Def);
3703 // For each comparison, invert the opcode. For each AND and OR, change the
3704 // opcode.
3705 switch (Def->getOpcode()) {
3706 default:
3707 llvm_unreachable("Unexpected opcode");
3708 case TargetOpcode::G_ICMP:
3709 case TargetOpcode::G_FCMP: {
3710 MachineOperand &PredOp = Def->getOperand(i: 1);
3711 CmpInst::Predicate NewP = CmpInst::getInversePredicate(
3712 pred: (CmpInst::Predicate)PredOp.getPredicate());
3713 PredOp.setPredicate(NewP);
3714 break;
3715 }
3716 case TargetOpcode::G_AND:
3717 Def->setDesc(Builder.getTII().get(Opcode: TargetOpcode::G_OR));
3718 break;
3719 case TargetOpcode::G_OR:
3720 Def->setDesc(Builder.getTII().get(Opcode: TargetOpcode::G_AND));
3721 break;
3722 }
3723 Observer.changedInstr(MI&: *Def);
3724 }
3725
3726 replaceRegWith(MRI, FromReg: MI.getOperand(i: 0).getReg(), ToReg: MI.getOperand(i: 1).getReg());
3727 MI.eraseFromParent();
3728}
3729
3730bool CombinerHelper::matchXorOfAndWithSameReg(
3731 MachineInstr &MI, std::pair<Register, Register> &MatchInfo) const {
3732 // Match (xor (and x, y), y) (or any of its commuted cases)
3733 assert(MI.getOpcode() == TargetOpcode::G_XOR);
3734 Register &X = MatchInfo.first;
3735 Register &Y = MatchInfo.second;
3736 Register AndReg = MI.getOperand(i: 1).getReg();
3737 Register SharedReg = MI.getOperand(i: 2).getReg();
3738
3739 // Find a G_AND on either side of the G_XOR.
3740 // Look for one of
3741 //
3742 // (xor (and x, y), SharedReg)
3743 // (xor SharedReg, (and x, y))
3744 if (!mi_match(R: AndReg, MRI, P: m_GAnd(L: m_Reg(R&: X), R: m_Reg(R&: Y)))) {
3745 std::swap(a&: AndReg, b&: SharedReg);
3746 if (!mi_match(R: AndReg, MRI, P: m_GAnd(L: m_Reg(R&: X), R: m_Reg(R&: Y))))
3747 return false;
3748 }
3749
3750 // Only do this if we'll eliminate the G_AND.
3751 if (!MRI.hasOneNonDBGUse(RegNo: AndReg))
3752 return false;
3753
3754 // We can combine if SharedReg is the same as either the LHS or RHS of the
3755 // G_AND.
3756 if (Y != SharedReg)
3757 std::swap(a&: X, b&: Y);
3758 return Y == SharedReg;
3759}
3760
3761void CombinerHelper::applyXorOfAndWithSameReg(
3762 MachineInstr &MI, std::pair<Register, Register> &MatchInfo) const {
3763 // Fold (xor (and x, y), y) -> (and (not x), y)
3764 Register X, Y;
3765 std::tie(args&: X, args&: Y) = MatchInfo;
3766 auto Not = Builder.buildNot(Dst: MRI.getType(Reg: X), Src0: X);
3767 Observer.changingInstr(MI);
3768 MI.setDesc(Builder.getTII().get(Opcode: TargetOpcode::G_AND));
3769 MI.getOperand(i: 1).setReg(Not->getOperand(i: 0).getReg());
3770 MI.getOperand(i: 2).setReg(Y);
3771 Observer.changedInstr(MI);
3772}
3773
3774bool CombinerHelper::matchPtrAddZero(MachineInstr &MI) const {
3775 auto &PtrAdd = cast<GPtrAdd>(Val&: MI);
3776 Register DstReg = PtrAdd.getReg(Idx: 0);
3777 LLT Ty = MRI.getType(Reg: DstReg);
3778 const DataLayout &DL = Builder.getMF().getDataLayout();
3779
3780 if (DL.isNonIntegralAddressSpace(AddrSpace: Ty.getScalarType().getAddressSpace()))
3781 return false;
3782
3783 if (Ty.isPointer()) {
3784 auto ConstVal = getIConstantVRegVal(VReg: PtrAdd.getBaseReg(), MRI);
3785 return ConstVal && *ConstVal == 0;
3786 }
3787
3788 assert(Ty.isVector() && "Expecting a vector type");
3789 const MachineInstr *VecMI;
3790 if (!mi_match(R: PtrAdd.getBaseReg(), MRI, P: m_MInstr(MI&: VecMI)))
3791 return false;
3792 return isBuildVectorAllZeros(MI: *VecMI, MRI);
3793}
3794
3795/// The second source operand is known to be a power of 2.
3796void CombinerHelper::applySimplifyURemByPow2(MachineInstr &MI) const {
3797 Register DstReg = MI.getOperand(i: 0).getReg();
3798 Register Src0 = MI.getOperand(i: 1).getReg();
3799 Register Pow2Src1 = MI.getOperand(i: 2).getReg();
3800 LLT Ty = MRI.getType(Reg: DstReg);
3801
3802 // Fold (urem x, pow2) -> (and x, pow2-1)
3803 auto NegOne = Builder.buildConstant(Res: Ty, Val: -1);
3804 auto Add = Builder.buildAdd(Dst: Ty, Src0: Pow2Src1, Src1: NegOne);
3805 Builder.buildAnd(Dst: DstReg, Src0, Src1: Add);
3806 MI.eraseFromParent();
3807}
3808
3809bool CombinerHelper::matchFoldBinOpIntoSelect(MachineInstr &MI,
3810 unsigned &SelectOpNo) const {
3811 Register LHS = MI.getOperand(i: 1).getReg();
3812 Register RHS = MI.getOperand(i: 2).getReg();
3813
3814 Register OtherOperandReg = RHS;
3815 SelectOpNo = 1;
3816 Register SelectTrue, SelectFalse;
3817
3818 // Don't do this unless the old select is going away. We want to eliminate the
3819 // binary operator, not replace a binop with a select.
3820 if (!mi_match(R: LHS, MRI,
3821 P: m_GISelect(Src0: m_Reg(), Src1: m_Reg(R&: SelectTrue), Src2: m_Reg(R&: SelectFalse))) ||
3822 !MRI.hasOneNonDBGUse(RegNo: LHS)) {
3823 OtherOperandReg = LHS;
3824 SelectOpNo = 2;
3825 if (!mi_match(R: RHS, MRI,
3826 P: m_GISelect(Src0: m_Reg(), Src1: m_Reg(R&: SelectTrue), Src2: m_Reg(R&: SelectFalse))) ||
3827 !MRI.hasOneNonDBGUse(RegNo: RHS))
3828 return false;
3829 }
3830
3831 MachineInstr *SelectLHS, *SelectRHS;
3832 if (!mi_match(R: SelectTrue, MRI, P: m_MInstr(MI&: SelectLHS)) ||
3833 !mi_match(R: SelectFalse, MRI, P: m_MInstr(MI&: SelectRHS)))
3834 return false;
3835
3836 if (!isConstantOrConstantVector(MI: *SelectLHS, MRI,
3837 /*AllowFP*/ true,
3838 /*AllowOpaqueConstants*/ false))
3839 return false;
3840 if (!isConstantOrConstantVector(MI: *SelectRHS, MRI,
3841 /*AllowFP*/ true,
3842 /*AllowOpaqueConstants*/ false))
3843 return false;
3844
3845 unsigned BinOpcode = MI.getOpcode();
3846
3847 // We know that one of the operands is a select of constants. Now verify that
3848 // the other binary operator operand is either a constant, or we can handle a
3849 // variable.
3850 bool CanFoldNonConst =
3851 (BinOpcode == TargetOpcode::G_AND || BinOpcode == TargetOpcode::G_OR) &&
3852 (isNullOrNullSplat(MI: *SelectLHS, MRI) ||
3853 isAllOnesOrAllOnesSplat(MI: *SelectLHS, MRI)) &&
3854 (isNullOrNullSplat(MI: *SelectRHS, MRI) ||
3855 isAllOnesOrAllOnesSplat(MI: *SelectRHS, MRI));
3856 if (CanFoldNonConst)
3857 return true;
3858
3859 MachineInstr *OtherOperandDef;
3860 if (!mi_match(R: OtherOperandReg, MRI, P: m_MInstr(MI&: OtherOperandDef)))
3861 return false;
3862 return isConstantOrConstantVector(MI: *OtherOperandDef, MRI,
3863 /*AllowFP*/ true,
3864 /*AllowOpaqueConstants*/ false);
3865}
3866
3867/// \p SelectOperand is the operand in binary operator \p MI that is the select
3868/// to fold.
3869void CombinerHelper::applyFoldBinOpIntoSelect(
3870 MachineInstr &MI, const unsigned &SelectOperand) const {
3871 Register Dst = MI.getOperand(i: 0).getReg();
3872 Register LHS = MI.getOperand(i: 1).getReg();
3873 Register RHS = MI.getOperand(i: 2).getReg();
3874 GSelect *Select =
3875 cast<GSelect>(Val: MRI.getVRegDef(Reg: MI.getOperand(i: SelectOperand).getReg()));
3876
3877 Register SelectCond = Select->getCondReg();
3878 Register SelectTrue = Select->getTrueReg();
3879 Register SelectFalse = Select->getFalseReg();
3880
3881 LLT Ty = MRI.getType(Reg: Dst);
3882 unsigned BinOpcode = MI.getOpcode();
3883
3884 Register FoldTrue, FoldFalse;
3885
3886 // We have a select-of-constants followed by a binary operator with a
3887 // constant. Eliminate the binop by pulling the constant math into the select.
3888 // Example: add (select Cond, CT, CF), CBO --> select Cond, CT + CBO, CF + CBO
3889 if (SelectOperand == 1) {
3890 // TODO: SelectionDAG verifies this actually constant folds before
3891 // committing to the combine.
3892
3893 FoldTrue = Builder.buildInstr(Opc: BinOpcode, DstOps: {Ty}, SrcOps: {SelectTrue, RHS}).getReg(Idx: 0);
3894 FoldFalse =
3895 Builder.buildInstr(Opc: BinOpcode, DstOps: {Ty}, SrcOps: {SelectFalse, RHS}).getReg(Idx: 0);
3896 } else {
3897 FoldTrue = Builder.buildInstr(Opc: BinOpcode, DstOps: {Ty}, SrcOps: {LHS, SelectTrue}).getReg(Idx: 0);
3898 FoldFalse =
3899 Builder.buildInstr(Opc: BinOpcode, DstOps: {Ty}, SrcOps: {LHS, SelectFalse}).getReg(Idx: 0);
3900 }
3901
3902 Builder.buildSelect(Res: Dst, Tst: SelectCond, Op0: FoldTrue, Op1: FoldFalse, Flags: MI.getFlags());
3903 MI.eraseFromParent();
3904}
3905
3906std::optional<SmallVector<Register, 8>>
3907CombinerHelper::findCandidatesForLoadOrCombine(const MachineInstr *Root) const {
3908 assert(Root->getOpcode() == TargetOpcode::G_OR && "Expected G_OR only!");
3909 // We want to detect if Root is part of a tree which represents a bunch
3910 // of loads being merged into a larger load. We'll try to recognize patterns
3911 // like, for example:
3912 //
3913 // Reg Reg
3914 // \ /
3915 // OR_1 Reg
3916 // \ /
3917 // OR_2
3918 // \ Reg
3919 // .. /
3920 // Root
3921 //
3922 // Reg Reg Reg Reg
3923 // \ / \ /
3924 // OR_1 OR_2
3925 // \ /
3926 // \ /
3927 // ...
3928 // Root
3929 //
3930 // Each "Reg" may have been produced by a load + some arithmetic. This
3931 // function will save each of them.
3932 SmallVector<Register, 8> RegsToVisit;
3933 SmallVector<const MachineInstr *, 7> Ors = {Root};
3934
3935 // In the "worst" case, we're dealing with a load for each byte. So, there
3936 // are at most #bytes - 1 ORs.
3937 const unsigned MaxIter =
3938 MRI.getType(Reg: Root->getOperand(i: 0).getReg()).getSizeInBytes() - 1;
3939 for (unsigned Iter = 0; Iter < MaxIter; ++Iter) {
3940 if (Ors.empty())
3941 break;
3942 const MachineInstr *Curr = Ors.pop_back_val();
3943 Register OrLHS = Curr->getOperand(i: 1).getReg();
3944 Register OrRHS = Curr->getOperand(i: 2).getReg();
3945
3946 // In the combine, we want to elimate the entire tree.
3947 if (!MRI.hasOneNonDBGUse(RegNo: OrLHS) || !MRI.hasOneNonDBGUse(RegNo: OrRHS))
3948 return std::nullopt;
3949
3950 // If it's a G_OR, save it and continue to walk. If it's not, then it's
3951 // something that may be a load + arithmetic.
3952 if (const MachineInstr *Or = getOpcodeDef(Opcode: TargetOpcode::G_OR, Reg: OrLHS, MRI))
3953 Ors.push_back(Elt: Or);
3954 else
3955 RegsToVisit.push_back(Elt: OrLHS);
3956 if (const MachineInstr *Or = getOpcodeDef(Opcode: TargetOpcode::G_OR, Reg: OrRHS, MRI))
3957 Ors.push_back(Elt: Or);
3958 else
3959 RegsToVisit.push_back(Elt: OrRHS);
3960 }
3961
3962 // We're going to try and merge each register into a wider power-of-2 type,
3963 // so we ought to have an even number of registers.
3964 if (RegsToVisit.empty() || RegsToVisit.size() % 2 != 0)
3965 return std::nullopt;
3966 return RegsToVisit;
3967}
3968
3969/// Helper function for findLoadOffsetsForLoadOrCombine.
3970///
3971/// Check if \p Reg is the result of loading a \p MemSizeInBits wide value,
3972/// and then moving that value into a specific byte offset.
3973///
3974/// e.g. x[i] << 24
3975///
3976/// \returns The load instruction and the byte offset it is moved into.
3977static std::optional<std::pair<GZExtLoad *, int64_t>>
3978matchLoadAndBytePosition(Register Reg, unsigned MemSizeInBits,
3979 const MachineRegisterInfo &MRI) {
3980 assert(MRI.hasOneNonDBGUse(Reg) &&
3981 "Expected Reg to only have one non-debug use?");
3982 Register MaybeLoad;
3983 int64_t Shift;
3984 if (!mi_match(R: Reg, MRI,
3985 P: m_OneNonDBGUse(SP: m_GShl(L: m_Reg(R&: MaybeLoad), R: m_ICst(Cst&: Shift))))) {
3986 Shift = 0;
3987 MaybeLoad = Reg;
3988 }
3989
3990 if (Shift % MemSizeInBits != 0)
3991 return std::nullopt;
3992
3993 // TODO: Handle other types of loads.
3994 auto *Load = getOpcodeDef<GZExtLoad>(Reg: MaybeLoad, MRI);
3995 if (!Load)
3996 return std::nullopt;
3997
3998 if (!Load->isUnordered() || Load->getMemSizeInBits() != MemSizeInBits)
3999 return std::nullopt;
4000
4001 return std::make_pair(x&: Load, y: Shift / MemSizeInBits);
4002}
4003
4004std::optional<std::tuple<GZExtLoad *, int64_t, GZExtLoad *>>
4005CombinerHelper::findLoadOffsetsForLoadOrCombine(
4006 SmallDenseMap<int64_t, int64_t, 8> &MemOffset2Idx,
4007 const SmallVector<Register, 8> &RegsToVisit,
4008 const unsigned MemSizeInBits) const {
4009
4010 // Each load found for the pattern. There should be one for each RegsToVisit.
4011 SmallSetVector<const MachineInstr *, 8> Loads;
4012
4013 // The lowest index used in any load. (The lowest "i" for each x[i].)
4014 int64_t LowestIdx = INT64_MAX;
4015
4016 // The load which uses the lowest index.
4017 GZExtLoad *LowestIdxLoad = nullptr;
4018
4019 // Keeps track of the load indices we see. We shouldn't see any indices twice.
4020 SmallSet<int64_t, 8> SeenIdx;
4021
4022 // Ensure each load is in the same MBB.
4023 // TODO: Support multiple MachineBasicBlocks.
4024 MachineBasicBlock *MBB = nullptr;
4025 const MachineMemOperand *MMO = nullptr;
4026
4027 // Earliest instruction-order load in the pattern.
4028 GZExtLoad *EarliestLoad = nullptr;
4029
4030 // Latest instruction-order load in the pattern.
4031 GZExtLoad *LatestLoad = nullptr;
4032
4033 // Base pointer which every load should share.
4034 Register BasePtr;
4035
4036 // We want to find a load for each register. Each load should have some
4037 // appropriate bit twiddling arithmetic. During this loop, we will also keep
4038 // track of the load which uses the lowest index. Later, we will check if we
4039 // can use its pointer in the final, combined load.
4040 for (auto Reg : RegsToVisit) {
4041 // Find the load, and find the position that it will end up in (e.g. a
4042 // shifted) value.
4043 auto LoadAndPos = matchLoadAndBytePosition(Reg, MemSizeInBits, MRI);
4044 if (!LoadAndPos)
4045 return std::nullopt;
4046 GZExtLoad *Load;
4047 int64_t DstPos;
4048 std::tie(args&: Load, args&: DstPos) = *LoadAndPos;
4049
4050 // TODO: Handle multiple MachineBasicBlocks. Currently not handled because
4051 // it is difficult to check for stores/calls/etc between loads.
4052 MachineBasicBlock *LoadMBB = Load->getParent();
4053 if (!MBB)
4054 MBB = LoadMBB;
4055 if (LoadMBB != MBB)
4056 return std::nullopt;
4057
4058 // Make sure that the MachineMemOperands of every seen load are compatible.
4059 auto &LoadMMO = Load->getMMO();
4060 if (!MMO)
4061 MMO = &LoadMMO;
4062 if (MMO->getAddrSpace() != LoadMMO.getAddrSpace())
4063 return std::nullopt;
4064
4065 // Find out what the base pointer and index for the load is.
4066 Register LoadPtr;
4067 int64_t Idx;
4068 if (!mi_match(R: Load->getOperand(i: 1).getReg(), MRI,
4069 P: m_GPtrAdd(L: m_Reg(R&: LoadPtr), R: m_ICst(Cst&: Idx)))) {
4070 LoadPtr = Load->getOperand(i: 1).getReg();
4071 Idx = 0;
4072 }
4073
4074 // Don't combine things like a[i], a[i] -> a bigger load.
4075 if (!SeenIdx.insert(V: Idx).second)
4076 return std::nullopt;
4077
4078 // Every load must share the same base pointer; don't combine things like:
4079 //
4080 // a[i], b[i + 1] -> a bigger load.
4081 if (!BasePtr.isValid())
4082 BasePtr = LoadPtr;
4083 if (BasePtr != LoadPtr)
4084 return std::nullopt;
4085
4086 if (Idx < LowestIdx) {
4087 LowestIdx = Idx;
4088 LowestIdxLoad = Load;
4089 }
4090
4091 // Keep track of the byte offset that this load ends up at. If we have seen
4092 // the byte offset, then stop here. We do not want to combine:
4093 //
4094 // a[i] << 16, a[i + k] << 16 -> a bigger load.
4095 if (!MemOffset2Idx.try_emplace(Key: DstPos, Args&: Idx).second)
4096 return std::nullopt;
4097 Loads.insert(X: Load);
4098
4099 // Keep track of the position of the earliest/latest loads in the pattern.
4100 // We will check that there are no load fold barriers between them later
4101 // on.
4102 //
4103 // FIXME: Is there a better way to check for load fold barriers?
4104 if (!EarliestLoad || dominates(DefMI: *Load, UseMI: *EarliestLoad))
4105 EarliestLoad = Load;
4106 if (!LatestLoad || dominates(DefMI: *LatestLoad, UseMI: *Load))
4107 LatestLoad = Load;
4108 }
4109
4110 // We found a load for each register. Let's check if each load satisfies the
4111 // pattern.
4112 assert(Loads.size() == RegsToVisit.size() &&
4113 "Expected to find a load for each register?");
4114 assert(EarliestLoad != LatestLoad && EarliestLoad &&
4115 LatestLoad && "Expected at least two loads?");
4116
4117 // Check if there are any stores, calls, etc. between any of the loads. If
4118 // there are, then we can't safely perform the combine.
4119 //
4120 // MaxIter is chosen based off the (worst case) number of iterations it
4121 // typically takes to succeed in the LLVM test suite plus some padding.
4122 //
4123 // FIXME: Is there a better way to check for load fold barriers?
4124 const unsigned MaxIter = 20;
4125 unsigned Iter = 0;
4126 for (const auto &MI : instructionsWithoutDebug(It: EarliestLoad->getIterator(),
4127 End: LatestLoad->getIterator())) {
4128 if (Loads.count(key: &MI))
4129 continue;
4130 if (MI.isLoadFoldBarrier())
4131 return std::nullopt;
4132 if (Iter++ == MaxIter)
4133 return std::nullopt;
4134 }
4135
4136 return std::make_tuple(args&: LowestIdxLoad, args&: LowestIdx, args&: LatestLoad);
4137}
4138
4139bool CombinerHelper::matchLoadOrCombine(
4140 MachineInstr &MI,
4141 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
4142 assert(MI.getOpcode() == TargetOpcode::G_OR);
4143 MachineFunction &MF = *MI.getMF();
4144 // Assuming a little-endian target, transform:
4145 // s8 *a = ...
4146 // s32 val = a[0] | (a[1] << 8) | (a[2] << 16) | (a[3] << 24)
4147 // =>
4148 // s32 val = *((i32)a)
4149 //
4150 // s8 *a = ...
4151 // s32 val = (a[0] << 24) | (a[1] << 16) | (a[2] << 8) | a[3]
4152 // =>
4153 // s32 val = BSWAP(*((s32)a))
4154 Register Dst = MI.getOperand(i: 0).getReg();
4155 LLT Ty = MRI.getType(Reg: Dst);
4156 if (Ty.isVector())
4157 return false;
4158
4159 // We need to combine at least two loads into this type. Since the smallest
4160 // possible load is into a byte, we need at least a 16-bit wide type.
4161 const unsigned WideMemSizeInBits = Ty.getSizeInBits();
4162 if (WideMemSizeInBits < 16 || WideMemSizeInBits % 8 != 0)
4163 return false;
4164
4165 // Match a collection of non-OR instructions in the pattern.
4166 auto RegsToVisit = findCandidatesForLoadOrCombine(Root: &MI);
4167 if (!RegsToVisit)
4168 return false;
4169
4170 // We have a collection of non-OR instructions. Figure out how wide each of
4171 // the small loads should be based off of the number of potential loads we
4172 // found.
4173 const unsigned NarrowMemSizeInBits = WideMemSizeInBits / RegsToVisit->size();
4174 if (NarrowMemSizeInBits % 8 != 0)
4175 return false;
4176
4177 // Check if each register feeding into each OR is a load from the same
4178 // base pointer + some arithmetic.
4179 //
4180 // e.g. a[0], a[1] << 8, a[2] << 16, etc.
4181 //
4182 // Also verify that each of these ends up putting a[i] into the same memory
4183 // offset as a load into a wide type would.
4184 SmallDenseMap<int64_t, int64_t, 8> MemOffset2Idx;
4185 GZExtLoad *LowestIdxLoad, *LatestLoad;
4186 int64_t LowestIdx;
4187 auto MaybeLoadInfo = findLoadOffsetsForLoadOrCombine(
4188 MemOffset2Idx, RegsToVisit: *RegsToVisit, MemSizeInBits: NarrowMemSizeInBits);
4189 if (!MaybeLoadInfo)
4190 return false;
4191 std::tie(args&: LowestIdxLoad, args&: LowestIdx, args&: LatestLoad) = *MaybeLoadInfo;
4192
4193 // We have a bunch of loads being OR'd together. Using the addresses + offsets
4194 // we found before, check if this corresponds to a big or little endian byte
4195 // pattern. If it does, then we can represent it using a load + possibly a
4196 // BSWAP.
4197 bool IsBigEndianTarget = MF.getDataLayout().isBigEndian();
4198 std::optional<bool> IsBigEndian = isBigEndian(MemOffset2Idx, LowestIdx);
4199 if (!IsBigEndian)
4200 return false;
4201 bool NeedsBSwap = IsBigEndianTarget != *IsBigEndian;
4202 if (NeedsBSwap && !isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_BSWAP, {Ty}}))
4203 return false;
4204
4205 // Make sure that the load from the lowest index produces offset 0 in the
4206 // final value.
4207 //
4208 // This ensures that we won't combine something like this:
4209 //
4210 // load x[i] -> byte 2
4211 // load x[i+1] -> byte 0 ---> wide_load x[i]
4212 // load x[i+2] -> byte 1
4213 const unsigned NumLoadsInTy = WideMemSizeInBits / NarrowMemSizeInBits;
4214 const unsigned ZeroByteOffset =
4215 *IsBigEndian
4216 ? bigEndianByteAt(ByteWidth: NumLoadsInTy, I: 0)
4217 : littleEndianByteAt(ByteWidth: NumLoadsInTy, I: 0);
4218 auto ZeroOffsetIdx = MemOffset2Idx.find(Val: ZeroByteOffset);
4219 if (ZeroOffsetIdx == MemOffset2Idx.end() ||
4220 ZeroOffsetIdx->second != LowestIdx)
4221 return false;
4222
4223 // We wil reuse the pointer from the load which ends up at byte offset 0. It
4224 // may not use index 0.
4225 Register Ptr = LowestIdxLoad->getPointerReg();
4226 const MachineMemOperand &MMO = LowestIdxLoad->getMMO();
4227 LegalityQuery::MemDesc MMDesc(MMO);
4228 MMDesc.MemoryTy = Ty;
4229 if (!isLegalOrBeforeLegalizer(
4230 Query: {TargetOpcode::G_LOAD, {Ty, MRI.getType(Reg: Ptr)}, {MMDesc}}))
4231 return false;
4232 auto PtrInfo = MMO.getPointerInfo();
4233 auto *NewMMO = MF.getMachineMemOperand(MMO: &MMO, PtrInfo, Size: WideMemSizeInBits / 8);
4234
4235 // Load must be allowed and fast on the target.
4236 LLVMContext &C = MF.getFunction().getContext();
4237 auto &DL = MF.getDataLayout();
4238 unsigned Fast = 0;
4239 if (!getTargetLowering().allowsMemoryAccess(Context&: C, DL, Ty, MMO: *NewMMO, Fast: &Fast) ||
4240 !Fast)
4241 return false;
4242
4243 MatchInfo = [=](MachineIRBuilder &MIB) {
4244 MIB.setInstrAndDebugLoc(*LatestLoad);
4245 Register LoadDst = NeedsBSwap ? MRI.cloneVirtualRegister(VReg: Dst) : Dst;
4246 MIB.buildLoad(Res: LoadDst, Addr: Ptr, MMO&: *NewMMO);
4247 if (NeedsBSwap)
4248 MIB.buildBSwap(Dst, Src0: LoadDst);
4249 };
4250 return true;
4251}
4252
4253bool CombinerHelper::matchExtendThroughPhis(MachineInstr &MI,
4254 MachineInstr *&ExtMI) const {
4255 auto &PHI = cast<GPhi>(Val&: MI);
4256 Register DstReg = PHI.getReg(Idx: 0);
4257
4258 // TODO: Extending a vector may be expensive, don't do this until heuristics
4259 // are better.
4260 if (MRI.getType(Reg: DstReg).isVector())
4261 return false;
4262
4263 // Try to match a phi, whose only use is an extend.
4264 if (!MRI.hasOneNonDBGUse(RegNo: DstReg))
4265 return false;
4266 ExtMI = &*MRI.use_instr_nodbg_begin(RegNo: DstReg);
4267 switch (ExtMI->getOpcode()) {
4268 case TargetOpcode::G_ANYEXT:
4269 return true; // G_ANYEXT is usually free.
4270 case TargetOpcode::G_ZEXT:
4271 case TargetOpcode::G_SEXT:
4272 break;
4273 default:
4274 return false;
4275 }
4276
4277 // If the target is likely to fold this extend away, don't propagate.
4278 if (Builder.getTII().isExtendLikelyToBeFolded(ExtMI&: *ExtMI, MRI))
4279 return false;
4280
4281 // We don't want to propagate the extends unless there's a good chance that
4282 // they'll be optimized in some way.
4283 // Collect the unique incoming values.
4284 SmallPtrSet<MachineInstr *, 4> InSrcs;
4285 for (unsigned I = 0; I < PHI.getNumIncomingValues(); ++I) {
4286 auto *DefMI = getDefIgnoringCopies(Reg: PHI.getIncomingValue(I), MRI);
4287 switch (DefMI->getOpcode()) {
4288 case TargetOpcode::G_LOAD:
4289 case TargetOpcode::G_TRUNC:
4290 case TargetOpcode::G_SEXT:
4291 case TargetOpcode::G_ZEXT:
4292 case TargetOpcode::G_ANYEXT:
4293 case TargetOpcode::G_CONSTANT:
4294 InSrcs.insert(Ptr: DefMI);
4295 // Don't try to propagate if there are too many places to create new
4296 // extends, chances are it'll increase code size.
4297 if (InSrcs.size() > 2)
4298 return false;
4299 break;
4300 default:
4301 return false;
4302 }
4303 }
4304 return true;
4305}
4306
4307void CombinerHelper::applyExtendThroughPhis(MachineInstr &MI,
4308 MachineInstr *&ExtMI) const {
4309 auto &PHI = cast<GPhi>(Val&: MI);
4310 Register DstReg = ExtMI->getOperand(i: 0).getReg();
4311 LLT ExtTy = MRI.getType(Reg: DstReg);
4312
4313 // Propagate the extension into the block of each incoming reg's block.
4314 // Use a SetVector here because PHIs can have duplicate edges, and we want
4315 // deterministic iteration order.
4316 SmallSetVector<MachineInstr *, 8> SrcMIs;
4317 SmallDenseMap<MachineInstr *, MachineInstr *, 8> OldToNewSrcMap;
4318 for (unsigned I = 0; I < PHI.getNumIncomingValues(); ++I) {
4319 auto SrcReg = PHI.getIncomingValue(I);
4320 MachineInstr *SrcMI;
4321 if (!mi_match(R: SrcReg, MRI, P: m_MInstr(MI&: SrcMI)))
4322 continue;
4323 if (!SrcMIs.insert(X: SrcMI))
4324 continue;
4325
4326 // Build an extend after each src inst.
4327 auto *MBB = SrcMI->getParent();
4328 MachineBasicBlock::iterator InsertPt = ++SrcMI->getIterator();
4329 if (InsertPt != MBB->end() && InsertPt->isPHI())
4330 InsertPt = MBB->getFirstNonPHI();
4331
4332 Builder.setInsertPt(MBB&: *SrcMI->getParent(), II: InsertPt);
4333 Builder.setDebugLoc(MI.getDebugLoc());
4334 auto NewExt = Builder.buildExtOrTrunc(ExtOpc: ExtMI->getOpcode(), Res: ExtTy, Op: SrcReg);
4335 OldToNewSrcMap[SrcMI] = NewExt;
4336 }
4337
4338 // Create a new phi with the extended inputs.
4339 Builder.setInstrAndDebugLoc(MI);
4340 auto NewPhi = Builder.buildInstrNoInsert(Opcode: TargetOpcode::G_PHI);
4341 NewPhi.addDef(RegNo: DstReg);
4342 for (const MachineOperand &MO : llvm::drop_begin(RangeOrContainer: MI.operands())) {
4343 if (!MO.isReg()) {
4344 NewPhi.addMBB(MBB: MO.getMBB());
4345 continue;
4346 }
4347 auto *NewSrc = OldToNewSrcMap[MRI.getVRegDef(Reg: MO.getReg())];
4348 NewPhi.addUse(RegNo: NewSrc->getOperand(i: 0).getReg());
4349 }
4350 Builder.insertInstr(MIB: NewPhi);
4351 ExtMI->eraseFromParent();
4352}
4353
4354bool CombinerHelper::matchExtractVecEltBuildVec(MachineInstr &MI,
4355 Register &Reg) const {
4356 assert(MI.getOpcode() == TargetOpcode::G_EXTRACT_VECTOR_ELT);
4357 // If we have a constant index, look for a G_BUILD_VECTOR source
4358 // and find the source register that the index maps to.
4359 Register SrcVec = MI.getOperand(i: 1).getReg();
4360 LLT SrcTy = MRI.getType(Reg: SrcVec);
4361 if (SrcTy.isScalableVector())
4362 return false;
4363
4364 auto Cst = getIConstantVRegValWithLookThrough(VReg: MI.getOperand(i: 2).getReg(), MRI);
4365 if (!Cst || Cst->Value.getZExtValue() >= SrcTy.getNumElements())
4366 return false;
4367
4368 unsigned VecIdx = Cst->Value.getZExtValue();
4369
4370 // Check if we have a build_vector or build_vector_trunc with an optional
4371 // trunc in front.
4372 MachineInstr *SrcVecMI;
4373 Register TruncSrc;
4374 if (mi_match(R: SrcVec, MRI, P: m_GTrunc(Src: m_Reg(R&: TruncSrc)))) {
4375 if (!mi_match(R: TruncSrc, MRI, P: m_MInstr(MI&: SrcVecMI)))
4376 return false;
4377 } else if (!mi_match(R: SrcVec, MRI, P: m_MInstr(MI&: SrcVecMI)))
4378 return false;
4379
4380 if (SrcVecMI->getOpcode() != TargetOpcode::G_BUILD_VECTOR &&
4381 SrcVecMI->getOpcode() != TargetOpcode::G_BUILD_VECTOR_TRUNC)
4382 return false;
4383
4384 EVT Ty(getMVTForLLT(Ty: SrcTy));
4385 if (!MRI.hasOneNonDBGUse(RegNo: SrcVec) &&
4386 !getTargetLowering().aggressivelyPreferBuildVectorSources(VecVT: Ty))
4387 return false;
4388
4389 Reg = SrcVecMI->getOperand(i: VecIdx + 1).getReg();
4390 return true;
4391}
4392
4393void CombinerHelper::applyExtractVecEltBuildVec(MachineInstr &MI,
4394 Register &Reg) const {
4395 // Check the type of the register, since it may have come from a
4396 // G_BUILD_VECTOR_TRUNC.
4397 LLT ScalarTy = MRI.getType(Reg);
4398 Register DstReg = MI.getOperand(i: 0).getReg();
4399 LLT DstTy = MRI.getType(Reg: DstReg);
4400
4401 if (ScalarTy != DstTy) {
4402 assert(ScalarTy.getSizeInBits() > DstTy.getSizeInBits());
4403 Builder.buildTrunc(Res: DstReg, Op: Reg);
4404 MI.eraseFromParent();
4405 return;
4406 }
4407 replaceSingleDefInstWithReg(MI, Replacement: Reg);
4408}
4409
4410bool CombinerHelper::matchExtractAllEltsFromBuildVector(
4411 MachineInstr &MI,
4412 SmallVectorImpl<std::pair<Register, MachineInstr *>> &SrcDstPairs) const {
4413 assert(MI.getOpcode() == TargetOpcode::G_BUILD_VECTOR);
4414 // This combine tries to find build_vector's which have every source element
4415 // extracted using G_EXTRACT_VECTOR_ELT. This can happen when transforms like
4416 // the masked load scalarization is run late in the pipeline. There's already
4417 // a combine for a similar pattern starting from the extract, but that
4418 // doesn't attempt to do it if there are multiple uses of the build_vector,
4419 // which in this case is true. Starting the combine from the build_vector
4420 // feels more natural than trying to find sibling nodes of extracts.
4421 // E.g.
4422 // %vec(<4 x s32>) = G_BUILD_VECTOR %s1(s32), %s2, %s3, %s4
4423 // %ext1 = G_EXTRACT_VECTOR_ELT %vec, 0
4424 // %ext2 = G_EXTRACT_VECTOR_ELT %vec, 1
4425 // %ext3 = G_EXTRACT_VECTOR_ELT %vec, 2
4426 // %ext4 = G_EXTRACT_VECTOR_ELT %vec, 3
4427 // ==>
4428 // replace ext{1,2,3,4} with %s{1,2,3,4}
4429
4430 Register DstReg = MI.getOperand(i: 0).getReg();
4431 LLT DstTy = MRI.getType(Reg: DstReg);
4432 unsigned NumElts = DstTy.getNumElements();
4433
4434 SmallBitVector ExtractedElts(NumElts);
4435 for (MachineInstr &II : MRI.use_nodbg_instructions(Reg: DstReg)) {
4436 if (II.getOpcode() != TargetOpcode::G_EXTRACT_VECTOR_ELT)
4437 return false;
4438 auto Cst = getIConstantVRegVal(VReg: II.getOperand(i: 2).getReg(), MRI);
4439 if (!Cst)
4440 return false;
4441 unsigned Idx = Cst->getZExtValue();
4442 if (Idx >= NumElts)
4443 return false; // Out of range.
4444 ExtractedElts.set(Idx);
4445 SrcDstPairs.emplace_back(
4446 Args: std::make_pair(x: MI.getOperand(i: Idx + 1).getReg(), y: &II));
4447 }
4448 // Match if every element was extracted.
4449 return ExtractedElts.all();
4450}
4451
4452void CombinerHelper::applyExtractAllEltsFromBuildVector(
4453 MachineInstr &MI,
4454 SmallVectorImpl<std::pair<Register, MachineInstr *>> &SrcDstPairs) const {
4455 assert(MI.getOpcode() == TargetOpcode::G_BUILD_VECTOR);
4456 for (auto &Pair : SrcDstPairs) {
4457 auto *ExtMI = Pair.second;
4458 replaceRegWith(MRI, FromReg: ExtMI->getOperand(i: 0).getReg(), ToReg: Pair.first);
4459 ExtMI->eraseFromParent();
4460 }
4461 MI.eraseFromParent();
4462}
4463
4464void CombinerHelper::applyBuildFn(
4465 MachineInstr &MI,
4466 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
4467 applyBuildFnNoErase(MI, MatchInfo);
4468 MI.eraseFromParent();
4469}
4470
4471void CombinerHelper::applyBuildFnNoErase(
4472 MachineInstr &MI,
4473 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
4474 MatchInfo(Builder);
4475}
4476
4477bool CombinerHelper::matchOrShiftToFunnelShift(MachineInstr &MI,
4478 bool AllowScalarConstants,
4479 BuildFnTy &MatchInfo) const {
4480 assert(MI.getOpcode() == TargetOpcode::G_OR);
4481
4482 Register Dst = MI.getOperand(i: 0).getReg();
4483 LLT Ty = MRI.getType(Reg: Dst);
4484 unsigned BitWidth = Ty.getScalarSizeInBits();
4485
4486 Register ShlSrc, ShlAmt, LShrSrc, LShrAmt, Amt;
4487 unsigned FshOpc = 0;
4488
4489 // Match (or (shl ...), (lshr ...)).
4490 if (!mi_match(R: Dst, MRI,
4491 // m_GOr() handles the commuted version as well.
4492 P: m_GOr(L: m_GShl(L: m_Reg(R&: ShlSrc), R: m_Reg(R&: ShlAmt)),
4493 R: m_GLShr(L: m_Reg(R&: LShrSrc), R: m_Reg(R&: LShrAmt)))))
4494 return false;
4495
4496 // Given constants C0 and C1 such that C0 + C1 is bit-width:
4497 // (or (shl x, C0), (lshr y, C1)) -> (fshl x, y, C0) or (fshr x, y, C1)
4498 int64_t CstShlAmt = 0, CstLShrAmt;
4499 if (mi_match(R: ShlAmt, MRI, P: m_ICstOrSplat(Cst&: CstShlAmt)) &&
4500 mi_match(R: LShrAmt, MRI, P: m_ICstOrSplat(Cst&: CstLShrAmt)) &&
4501 CstShlAmt + CstLShrAmt == BitWidth) {
4502 FshOpc = TargetOpcode::G_FSHR;
4503 Amt = LShrAmt;
4504 } else if (mi_match(R: LShrAmt, MRI,
4505 P: m_GSub(L: m_SpecificICstOrSplat(RequestedValue: BitWidth), R: m_Reg(R&: Amt))) &&
4506 ShlAmt == Amt) {
4507 // (or (shl x, amt), (lshr y, (sub bw, amt))) -> (fshl x, y, amt)
4508 FshOpc = TargetOpcode::G_FSHL;
4509 } else if (mi_match(R: ShlAmt, MRI,
4510 P: m_GSub(L: m_SpecificICstOrSplat(RequestedValue: BitWidth), R: m_Reg(R&: Amt))) &&
4511 LShrAmt == Amt) {
4512 // (or (shl x, (sub bw, amt)), (lshr y, amt)) -> (fshr x, y, amt)
4513 FshOpc = TargetOpcode::G_FSHR;
4514 } else {
4515 return false;
4516 }
4517
4518 LLT AmtTy = MRI.getType(Reg: Amt);
4519 if (!isLegalOrBeforeLegalizer(Query: {FshOpc, {Ty, AmtTy}}) &&
4520 (!AllowScalarConstants || CstShlAmt == 0 || !Ty.isScalar()))
4521 return false;
4522
4523 MatchInfo = [=](MachineIRBuilder &B) {
4524 B.buildInstr(Opc: FshOpc, DstOps: {Dst}, SrcOps: {ShlSrc, LShrSrc, Amt});
4525 };
4526 return true;
4527}
4528
4529/// Match an FSHL or FSHR that can be combined to a ROTR or ROTL rotate.
4530bool CombinerHelper::matchFunnelShiftToRotate(MachineInstr &MI) const {
4531 unsigned Opc = MI.getOpcode();
4532 assert(Opc == TargetOpcode::G_FSHL || Opc == TargetOpcode::G_FSHR);
4533 Register X = MI.getOperand(i: 1).getReg();
4534 Register Y = MI.getOperand(i: 2).getReg();
4535 if (X != Y)
4536 return false;
4537 unsigned RotateOpc =
4538 Opc == TargetOpcode::G_FSHL ? TargetOpcode::G_ROTL : TargetOpcode::G_ROTR;
4539 return isLegalOrBeforeLegalizer(Query: {RotateOpc, {MRI.getType(Reg: X), MRI.getType(Reg: Y)}});
4540}
4541
4542void CombinerHelper::applyFunnelShiftToRotate(MachineInstr &MI) const {
4543 unsigned Opc = MI.getOpcode();
4544 assert(Opc == TargetOpcode::G_FSHL || Opc == TargetOpcode::G_FSHR);
4545 bool IsFSHL = Opc == TargetOpcode::G_FSHL;
4546 Observer.changingInstr(MI);
4547 MI.setDesc(Builder.getTII().get(Opcode: IsFSHL ? TargetOpcode::G_ROTL
4548 : TargetOpcode::G_ROTR));
4549 MI.removeOperand(OpNo: 2);
4550 Observer.changedInstr(MI);
4551}
4552
4553// Fold (rot x, c) -> (rot x, c % BitSize)
4554bool CombinerHelper::matchRotateOutOfRange(MachineInstr &MI) const {
4555 assert(MI.getOpcode() == TargetOpcode::G_ROTL ||
4556 MI.getOpcode() == TargetOpcode::G_ROTR);
4557 unsigned Bitsize =
4558 MRI.getType(Reg: MI.getOperand(i: 0).getReg()).getScalarSizeInBits();
4559 Register AmtReg = MI.getOperand(i: 2).getReg();
4560 bool OutOfRange = false;
4561 auto MatchOutOfRange = [Bitsize, &OutOfRange](const Constant *C) {
4562 if (auto *CI = dyn_cast<ConstantInt>(Val: C))
4563 OutOfRange |= CI->getValue().uge(RHS: Bitsize);
4564 return true;
4565 };
4566 return matchUnaryPredicate(MRI, Reg: AmtReg, Match: MatchOutOfRange) && OutOfRange;
4567}
4568
4569void CombinerHelper::applyRotateOutOfRange(MachineInstr &MI) const {
4570 assert(MI.getOpcode() == TargetOpcode::G_ROTL ||
4571 MI.getOpcode() == TargetOpcode::G_ROTR);
4572 unsigned Bitsize =
4573 MRI.getType(Reg: MI.getOperand(i: 0).getReg()).getScalarSizeInBits();
4574 Register Amt = MI.getOperand(i: 2).getReg();
4575 LLT AmtTy = MRI.getType(Reg: Amt);
4576 auto Bits = Builder.buildConstant(Res: AmtTy, Val: Bitsize);
4577 Amt = Builder.buildURem(Dst: AmtTy, Src0: MI.getOperand(i: 2).getReg(), Src1: Bits).getReg(Idx: 0);
4578 Observer.changingInstr(MI);
4579 MI.getOperand(i: 2).setReg(Amt);
4580 Observer.changedInstr(MI);
4581}
4582
4583bool CombinerHelper::matchICmpToTrueFalseKnownBits(MachineInstr &MI,
4584 int64_t &MatchInfo) const {
4585 assert(MI.getOpcode() == TargetOpcode::G_ICMP);
4586 auto Pred = static_cast<CmpInst::Predicate>(MI.getOperand(i: 1).getPredicate());
4587
4588 // We want to avoid calling KnownBits on the LHS if possible, as this combine
4589 // has no filter and runs on every G_ICMP instruction. We can avoid calling
4590 // KnownBits on the LHS in two cases:
4591 //
4592 // - The RHS is unknown: Constants are always on RHS. If the RHS is unknown
4593 // we cannot do any transforms so we can safely bail out early.
4594 // - The RHS is zero: we don't need to know the LHS to do unsigned <0 and
4595 // >=0.
4596 auto KnownRHS = VT->getKnownBits(R: MI.getOperand(i: 3).getReg());
4597 if (KnownRHS.isUnknown())
4598 return false;
4599
4600 std::optional<bool> KnownVal;
4601 if (KnownRHS.isZero()) {
4602 // ? uge 0 -> always true
4603 // ? ult 0 -> always false
4604 if (Pred == CmpInst::ICMP_UGE)
4605 KnownVal = true;
4606 else if (Pred == CmpInst::ICMP_ULT)
4607 KnownVal = false;
4608 }
4609
4610 if (!KnownVal) {
4611 auto KnownLHS = VT->getKnownBits(R: MI.getOperand(i: 2).getReg());
4612 KnownVal = ICmpInst::compare(LHS: KnownLHS, RHS: KnownRHS, Pred);
4613 }
4614
4615 if (!KnownVal)
4616 return false;
4617 MatchInfo =
4618 *KnownVal
4619 ? getICmpTrueVal(TLI: getTargetLowering(),
4620 /*IsVector = */
4621 MRI.getType(Reg: MI.getOperand(i: 0).getReg()).isVector(),
4622 /* IsFP = */ false)
4623 : 0;
4624 return true;
4625}
4626
4627bool CombinerHelper::matchICmpToLHSKnownBits(
4628 MachineInstr &MI,
4629 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
4630 assert(MI.getOpcode() == TargetOpcode::G_ICMP);
4631 // Given:
4632 //
4633 // %x = G_WHATEVER (... x is known to be 0 or 1 ...)
4634 // %cmp = G_ICMP ne %x, 0
4635 //
4636 // Or:
4637 //
4638 // %x = G_WHATEVER (... x is known to be 0 or 1 ...)
4639 // %cmp = G_ICMP eq %x, 1
4640 //
4641 // We can replace %cmp with %x assuming true is 1 on the target.
4642 auto Pred = static_cast<CmpInst::Predicate>(MI.getOperand(i: 1).getPredicate());
4643 if (!CmpInst::isEquality(pred: Pred))
4644 return false;
4645 Register Dst = MI.getOperand(i: 0).getReg();
4646 LLT DstTy = MRI.getType(Reg: Dst);
4647 if (getICmpTrueVal(TLI: getTargetLowering(), IsVector: DstTy.isVector(),
4648 /* IsFP = */ false) != 1)
4649 return false;
4650 int64_t OneOrZero = Pred == CmpInst::ICMP_EQ;
4651 if (!mi_match(R: MI.getOperand(i: 3).getReg(), MRI, P: m_SpecificICst(RequestedValue: OneOrZero)))
4652 return false;
4653 Register LHS = MI.getOperand(i: 2).getReg();
4654 auto KnownLHS = VT->getKnownBits(R: LHS);
4655 if (KnownLHS.getMinValue() != 0 || KnownLHS.getMaxValue() != 1)
4656 return false;
4657 // Make sure replacing Dst with the LHS is a legal operation.
4658 LLT LHSTy = MRI.getType(Reg: LHS);
4659 unsigned LHSSize = LHSTy.getSizeInBits();
4660 unsigned DstSize = DstTy.getSizeInBits();
4661 unsigned Op = TargetOpcode::COPY;
4662 if (DstSize != LHSSize)
4663 Op = DstSize < LHSSize ? TargetOpcode::G_TRUNC : TargetOpcode::G_ZEXT;
4664 if (!isLegalOrBeforeLegalizer(Query: {Op, {DstTy, LHSTy}}))
4665 return false;
4666 MatchInfo = [=](MachineIRBuilder &B) { B.buildInstr(Opc: Op, DstOps: {Dst}, SrcOps: {LHS}); };
4667 return true;
4668}
4669
4670// Replace (and (or x, c1), c2) with (and x, c2) iff c1 & c2 == 0
4671bool CombinerHelper::matchAndOrDisjointMask(
4672 MachineInstr &MI,
4673 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
4674 assert(MI.getOpcode() == TargetOpcode::G_AND);
4675
4676 // Ignore vector types to simplify matching the two constants.
4677 // TODO: do this for vectors and scalars via a demanded bits analysis.
4678 LLT Ty = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
4679 if (Ty.isVector())
4680 return false;
4681
4682 Register Src;
4683 Register AndMaskReg;
4684 int64_t AndMaskBits;
4685 int64_t OrMaskBits;
4686 if (!mi_match(MI, MRI,
4687 P: m_GAnd(L: m_GOr(L: m_Reg(R&: Src), R: m_ICst(Cst&: OrMaskBits)),
4688 R: m_all_of(preds: m_ICst(Cst&: AndMaskBits), preds: m_Reg(R&: AndMaskReg)))))
4689 return false;
4690
4691 // Check if OrMask could turn on any bits in Src.
4692 if (AndMaskBits & OrMaskBits)
4693 return false;
4694
4695 MatchInfo = [=, &MI](MachineIRBuilder &B) {
4696 Observer.changingInstr(MI);
4697 // Canonicalize the result to have the constant on the RHS.
4698 if (MI.getOperand(i: 1).getReg() == AndMaskReg)
4699 MI.getOperand(i: 2).setReg(AndMaskReg);
4700 MI.getOperand(i: 1).setReg(Src);
4701 Observer.changedInstr(MI);
4702 };
4703 return true;
4704}
4705
4706/// Form a G_SBFX from a G_SEXT_INREG fed by a right shift.
4707bool CombinerHelper::matchBitfieldExtractFromSExtInReg(
4708 MachineInstr &MI,
4709 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
4710 assert(MI.getOpcode() == TargetOpcode::G_SEXT_INREG);
4711 Register Dst = MI.getOperand(i: 0).getReg();
4712 Register Src = MI.getOperand(i: 1).getReg();
4713 LLT Ty = MRI.getType(Reg: Src);
4714 LLT ExtractTy = getTargetLowering().getPreferredShiftAmountTy(ShiftValueTy: Ty);
4715 if (!LI || !LI->isLegalOrCustom(Query: {TargetOpcode::G_SBFX, {Ty, ExtractTy}}))
4716 return false;
4717 int64_t Width = MI.getOperand(i: 2).getImm();
4718 Register ShiftSrc;
4719 int64_t ShiftImm;
4720 if (!mi_match(
4721 R: Src, MRI,
4722 P: m_OneNonDBGUse(SP: m_any_of(preds: m_GAShr(L: m_Reg(R&: ShiftSrc), R: m_ICst(Cst&: ShiftImm)),
4723 preds: m_GLShr(L: m_Reg(R&: ShiftSrc), R: m_ICst(Cst&: ShiftImm))))))
4724 return false;
4725 if (ShiftImm < 0 || ShiftImm + Width > Ty.getScalarSizeInBits())
4726 return false;
4727
4728 MatchInfo = [=](MachineIRBuilder &B) {
4729 auto Cst1 = B.buildConstant(Res: ExtractTy, Val: ShiftImm);
4730 auto Cst2 = B.buildConstant(Res: ExtractTy, Val: Width);
4731 B.buildSbfx(Dst, Src: ShiftSrc, LSB: Cst1, Width: Cst2);
4732 };
4733 return true;
4734}
4735
4736/// Form a G_UBFX from "(a srl b) & mask", where b and mask are constants.
4737bool CombinerHelper::matchBitfieldExtractFromAnd(MachineInstr &MI,
4738 BuildFnTy &MatchInfo) const {
4739 GAnd *And = cast<GAnd>(Val: &MI);
4740 Register Dst = And->getReg(Idx: 0);
4741 LLT Ty = MRI.getType(Reg: Dst);
4742 LLT ExtractTy = getTargetLowering().getPreferredShiftAmountTy(ShiftValueTy: Ty);
4743 // Note that isLegalOrBeforeLegalizer is stricter and does not take custom
4744 // into account.
4745 if (LI && !LI->isLegalOrCustom(Query: {TargetOpcode::G_UBFX, {Ty, ExtractTy}}))
4746 return false;
4747
4748 int64_t AndImm, LSBImm;
4749 Register ShiftSrc;
4750 const unsigned Size = Ty.getScalarSizeInBits();
4751 if (!mi_match(R: And->getReg(Idx: 0), MRI,
4752 P: m_GAnd(L: m_OneNonDBGUse(SP: m_GLShr(L: m_Reg(R&: ShiftSrc), R: m_ICst(Cst&: LSBImm))),
4753 R: m_ICst(Cst&: AndImm))))
4754 return false;
4755
4756 // AndImm is sign-extended to 64 bits by m_ICst; restrict it to the operand
4757 // width so an all-ones mask (a redundant AND) is not misread as a wider mask.
4758 uint64_t MaybeMask = static_cast<uint64_t>(AndImm);
4759 if (Size < 64)
4760 MaybeMask &= maskTrailingOnes<uint64_t>(N: Size);
4761
4762 // The mask is a mask of the low bits iff imm & (imm+1) == 0.
4763 if (MaybeMask & (MaybeMask + 1))
4764 return false;
4765
4766 // LSB must fit within the register.
4767 if (static_cast<uint64_t>(LSBImm) >= Size)
4768 return false;
4769
4770 uint64_t Width = APInt(Size, MaybeMask).countr_one();
4771 // The extracted field [LSB, LSB+Width) must fit within the register.
4772 // Otherwise this is a redundant AND (e.g. an all-ones mask combined with a
4773 // non-zero shift) that is better handled by other combines, and would form
4774 // an out-of-range bitfield extract.
4775 if (static_cast<uint64_t>(LSBImm) + Width > Size)
4776 return false;
4777
4778 MatchInfo = [=](MachineIRBuilder &B) {
4779 auto WidthCst = B.buildConstant(Res: ExtractTy, Val: Width);
4780 auto LSBCst = B.buildConstant(Res: ExtractTy, Val: LSBImm);
4781 B.buildInstr(Opc: TargetOpcode::G_UBFX, DstOps: {Dst}, SrcOps: {ShiftSrc, LSBCst, WidthCst});
4782 };
4783 return true;
4784}
4785
4786bool CombinerHelper::matchBitfieldExtractFromShr(
4787 MachineInstr &MI,
4788 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
4789 const unsigned Opcode = MI.getOpcode();
4790 assert(Opcode == TargetOpcode::G_ASHR || Opcode == TargetOpcode::G_LSHR);
4791
4792 const Register Dst = MI.getOperand(i: 0).getReg();
4793
4794 const unsigned ExtrOpcode = Opcode == TargetOpcode::G_ASHR
4795 ? TargetOpcode::G_SBFX
4796 : TargetOpcode::G_UBFX;
4797
4798 // Check if the type we would use for the extract is legal
4799 LLT Ty = MRI.getType(Reg: Dst);
4800 LLT ExtractTy = getTargetLowering().getPreferredShiftAmountTy(ShiftValueTy: Ty);
4801 if (!LI || !LI->isLegalOrCustom(Query: {ExtrOpcode, {Ty, ExtractTy}}))
4802 return false;
4803
4804 Register ShlSrc;
4805 int64_t ShrAmt;
4806 int64_t ShlAmt;
4807 const unsigned Size = Ty.getScalarSizeInBits();
4808
4809 // Try to match shr (shl x, c1), c2
4810 if (!mi_match(R: Dst, MRI,
4811 P: m_BinOp(Opcode,
4812 L: m_OneNonDBGUse(SP: m_GShl(L: m_Reg(R&: ShlSrc), R: m_ICst(Cst&: ShlAmt))),
4813 R: m_ICst(Cst&: ShrAmt))))
4814 return false;
4815
4816 // Make sure that the shift sizes can fit a bitfield extract
4817 if (ShlAmt < 0 || ShlAmt > ShrAmt || ShrAmt >= Size)
4818 return false;
4819
4820 // Skip this combine if the G_SEXT_INREG combine could handle it
4821 if (Opcode == TargetOpcode::G_ASHR && ShlAmt == ShrAmt)
4822 return false;
4823
4824 // Calculate start position and width of the extract
4825 const int64_t Pos = ShrAmt - ShlAmt;
4826 const int64_t Width = Size - ShrAmt;
4827
4828 MatchInfo = [=](MachineIRBuilder &B) {
4829 auto WidthCst = B.buildConstant(Res: ExtractTy, Val: Width);
4830 auto PosCst = B.buildConstant(Res: ExtractTy, Val: Pos);
4831 B.buildInstr(Opc: ExtrOpcode, DstOps: {Dst}, SrcOps: {ShlSrc, PosCst, WidthCst});
4832 };
4833 return true;
4834}
4835
4836bool CombinerHelper::matchBitfieldExtractFromShrAnd(
4837 MachineInstr &MI,
4838 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
4839 const unsigned Opcode = MI.getOpcode();
4840 assert(Opcode == TargetOpcode::G_LSHR || Opcode == TargetOpcode::G_ASHR);
4841
4842 const Register Dst = MI.getOperand(i: 0).getReg();
4843 LLT Ty = MRI.getType(Reg: Dst);
4844 LLT ExtractTy = getTargetLowering().getPreferredShiftAmountTy(ShiftValueTy: Ty);
4845 if (LI && !LI->isLegalOrCustom(Query: {TargetOpcode::G_UBFX, {Ty, ExtractTy}}))
4846 return false;
4847
4848 // Try to match shr (and x, c1), c2
4849 Register AndSrc;
4850 int64_t ShrAmt;
4851 int64_t SMask;
4852 if (!mi_match(R: Dst, MRI,
4853 P: m_BinOp(Opcode,
4854 L: m_OneNonDBGUse(SP: m_GAnd(L: m_Reg(R&: AndSrc), R: m_ICst(Cst&: SMask))),
4855 R: m_ICst(Cst&: ShrAmt))))
4856 return false;
4857
4858 const unsigned Size = Ty.getScalarSizeInBits();
4859 if (ShrAmt < 0 || ShrAmt >= Size)
4860 return false;
4861
4862 // If the shift subsumes the mask, emit the 0 directly.
4863 if (0 == (SMask >> ShrAmt)) {
4864 MatchInfo = [=](MachineIRBuilder &B) {
4865 B.buildConstant(Res: Dst, Val: 0);
4866 };
4867 return true;
4868 }
4869
4870 // Check that ubfx can do the extraction, with no holes in the mask.
4871 uint64_t UMask = SMask;
4872 UMask |= maskTrailingOnes<uint64_t>(N: ShrAmt);
4873 UMask &= maskTrailingOnes<uint64_t>(N: Size);
4874 if (!isMask_64(Value: UMask))
4875 return false;
4876
4877 // Calculate start position and width of the extract.
4878 const int64_t Pos = ShrAmt;
4879 const int64_t Width = llvm::countr_one(Value: UMask) - ShrAmt;
4880
4881 // It's preferable to keep the shift, rather than form G_SBFX.
4882 // TODO: remove the G_AND via demanded bits analysis.
4883 if (Opcode == TargetOpcode::G_ASHR && Width + ShrAmt == Size)
4884 return false;
4885
4886 MatchInfo = [=](MachineIRBuilder &B) {
4887 auto WidthCst = B.buildConstant(Res: ExtractTy, Val: Width);
4888 auto PosCst = B.buildConstant(Res: ExtractTy, Val: Pos);
4889 B.buildInstr(Opc: TargetOpcode::G_UBFX, DstOps: {Dst}, SrcOps: {AndSrc, PosCst, WidthCst});
4890 };
4891 return true;
4892}
4893
4894bool CombinerHelper::reassociationCanBreakAddressingModePattern(
4895 MachineInstr &MI) const {
4896 auto &PtrAdd = cast<GPtrAdd>(Val&: MI);
4897
4898 Register Src1Reg = PtrAdd.getBaseReg();
4899 auto *Src1Def = getOpcodeDef<GPtrAdd>(Reg: Src1Reg, MRI);
4900 if (!Src1Def)
4901 return false;
4902
4903 Register Src2Reg = PtrAdd.getOffsetReg();
4904
4905 if (MRI.hasOneNonDBGUse(RegNo: Src1Reg))
4906 return false;
4907
4908 auto C1 = getIConstantVRegVal(VReg: Src1Def->getOffsetReg(), MRI);
4909 if (!C1)
4910 return false;
4911 auto C2 = getIConstantVRegVal(VReg: Src2Reg, MRI);
4912 if (!C2)
4913 return false;
4914
4915 const APInt &C1APIntVal = *C1;
4916 const APInt &C2APIntVal = *C2;
4917 const int64_t CombinedValue = (C1APIntVal + C2APIntVal).getSExtValue();
4918
4919 for (auto &UseMI : MRI.use_nodbg_instructions(Reg: PtrAdd.getReg(Idx: 0))) {
4920 // This combine may end up running before ptrtoint/inttoptr combines
4921 // manage to eliminate redundant conversions, so try to look through them.
4922 MachineInstr *ConvUseMI = &UseMI;
4923 unsigned ConvUseOpc = ConvUseMI->getOpcode();
4924 while (ConvUseOpc == TargetOpcode::G_INTTOPTR ||
4925 ConvUseOpc == TargetOpcode::G_PTRTOINT) {
4926 Register DefReg = ConvUseMI->getOperand(i: 0).getReg();
4927 if (!MRI.hasOneNonDBGUse(RegNo: DefReg))
4928 break;
4929 ConvUseMI = &*MRI.use_instr_nodbg_begin(RegNo: DefReg);
4930 ConvUseOpc = ConvUseMI->getOpcode();
4931 }
4932 auto *LdStMI = dyn_cast<GLoadStore>(Val: ConvUseMI);
4933 if (!LdStMI)
4934 continue;
4935 // Is x[offset2] already not a legal addressing mode? If so then
4936 // reassociating the constants breaks nothing (we test offset2 because
4937 // that's the one we hope to fold into the load or store).
4938 TargetLoweringBase::AddrMode AM;
4939 AM.HasBaseReg = true;
4940 AM.BaseOffs = C2APIntVal.getSExtValue();
4941 unsigned AS = MRI.getType(Reg: LdStMI->getPointerReg()).getAddressSpace();
4942 Type *AccessTy = getTypeForLLT(Ty: LdStMI->getMMO().getMemoryType(),
4943 C&: PtrAdd.getMF()->getFunction().getContext());
4944 const auto &TLI = *PtrAdd.getMF()->getSubtarget().getTargetLowering();
4945 if (!TLI.isLegalAddressingMode(DL: PtrAdd.getMF()->getDataLayout(), AM,
4946 Ty: AccessTy, AddrSpace: AS))
4947 continue;
4948
4949 // Would x[offset1+offset2] still be a legal addressing mode?
4950 AM.BaseOffs = CombinedValue;
4951 if (!TLI.isLegalAddressingMode(DL: PtrAdd.getMF()->getDataLayout(), AM,
4952 Ty: AccessTy, AddrSpace: AS))
4953 return true;
4954 }
4955
4956 return false;
4957}
4958
4959bool CombinerHelper::matchReassocConstantInnerRHS(GPtrAdd &MI,
4960 MachineInstr *RHS,
4961 BuildFnTy &MatchInfo) const {
4962 // G_PTR_ADD(BASE, G_ADD(X, C)) -> G_PTR_ADD(G_PTR_ADD(BASE, X), C)
4963 Register Src1Reg = MI.getOperand(i: 1).getReg();
4964 if (RHS->getOpcode() != TargetOpcode::G_ADD)
4965 return false;
4966 auto C2 = getIConstantVRegVal(VReg: RHS->getOperand(i: 2).getReg(), MRI);
4967 if (!C2)
4968 return false;
4969
4970 // If both additions are nuw, the reassociated additions are also nuw.
4971 // If the original G_PTR_ADD is additionally nusw, X and C are both not
4972 // negative, so BASE+X is between BASE and BASE+(X+C). The new G_PTR_ADDs are
4973 // therefore also nusw.
4974 // If the original G_PTR_ADD is additionally inbounds (which implies nusw),
4975 // the new G_PTR_ADDs are then also inbounds.
4976 unsigned PtrAddFlags = MI.getFlags();
4977 unsigned AddFlags = RHS->getFlags();
4978 bool IsNoUWrap = PtrAddFlags & AddFlags & MachineInstr::MIFlag::NoUWrap;
4979 bool IsNoUSWrap = IsNoUWrap && (PtrAddFlags & MachineInstr::MIFlag::NoUSWrap);
4980 bool IsInBounds = IsNoUWrap && (PtrAddFlags & MachineInstr::MIFlag::InBounds);
4981 unsigned Flags = 0;
4982 if (IsNoUWrap)
4983 Flags |= MachineInstr::MIFlag::NoUWrap;
4984 if (IsNoUSWrap)
4985 Flags |= MachineInstr::MIFlag::NoUSWrap;
4986 if (IsInBounds)
4987 Flags |= MachineInstr::MIFlag::InBounds;
4988
4989 MatchInfo = [=, &MI](MachineIRBuilder &B) {
4990 LLT PtrTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
4991
4992 auto NewBase =
4993 Builder.buildPtrAdd(Res: PtrTy, Op0: Src1Reg, Op1: RHS->getOperand(i: 1).getReg(), Flags);
4994 Observer.changingInstr(MI);
4995 MI.getOperand(i: 1).setReg(NewBase.getReg(Idx: 0));
4996 MI.getOperand(i: 2).setReg(RHS->getOperand(i: 2).getReg());
4997 MI.setFlags(Flags);
4998 Observer.changedInstr(MI);
4999 };
5000 return !reassociationCanBreakAddressingModePattern(MI);
5001}
5002
5003bool CombinerHelper::matchReassocConstantInnerLHS(GPtrAdd &MI,
5004 MachineInstr *LHS,
5005 MachineInstr *RHS,
5006 BuildFnTy &MatchInfo) const {
5007 // G_PTR_ADD (G_PTR_ADD X, C), Y) -> (G_PTR_ADD (G_PTR_ADD(X, Y), C)
5008 // if and only if (G_PTR_ADD X, C) has one use.
5009 Register LHSBase;
5010 std::optional<ValueAndVReg> LHSCstOff;
5011 if (!mi_match(R: MI.getBaseReg(), MRI,
5012 P: m_OneNonDBGUse(SP: m_GPtrAdd(L: m_Reg(R&: LHSBase), R: m_GCst(ValReg&: LHSCstOff)))))
5013 return false;
5014
5015 auto *LHSPtrAdd = cast<GPtrAdd>(Val: LHS);
5016
5017 // Reassociating nuw additions preserves nuw. If both original G_PTR_ADDs are
5018 // nuw and inbounds (which implies nusw), the offsets are both non-negative,
5019 // so the new G_PTR_ADDs are also inbounds.
5020 unsigned PtrAddFlags = MI.getFlags();
5021 unsigned LHSPtrAddFlags = LHSPtrAdd->getFlags();
5022 bool IsNoUWrap = PtrAddFlags & LHSPtrAddFlags & MachineInstr::MIFlag::NoUWrap;
5023 bool IsNoUSWrap = IsNoUWrap && (PtrAddFlags & LHSPtrAddFlags &
5024 MachineInstr::MIFlag::NoUSWrap);
5025 bool IsInBounds = IsNoUWrap && (PtrAddFlags & LHSPtrAddFlags &
5026 MachineInstr::MIFlag::InBounds);
5027 unsigned Flags = 0;
5028 if (IsNoUWrap)
5029 Flags |= MachineInstr::MIFlag::NoUWrap;
5030 if (IsNoUSWrap)
5031 Flags |= MachineInstr::MIFlag::NoUSWrap;
5032 if (IsInBounds)
5033 Flags |= MachineInstr::MIFlag::InBounds;
5034
5035 MatchInfo = [=, &MI](MachineIRBuilder &B) {
5036 // When we change LHSPtrAdd's offset register we might cause it to use a reg
5037 // before its def. Sink the instruction so the outer PTR_ADD to ensure this
5038 // doesn't happen.
5039 LHSPtrAdd->moveBefore(MovePos: &MI);
5040 Register RHSReg = MI.getOffsetReg();
5041 // set VReg will cause type mismatch if it comes from extend/trunc
5042 auto NewCst = B.buildConstant(Res: MRI.getType(Reg: RHSReg), Val: LHSCstOff->Value);
5043 Observer.changingInstr(MI);
5044 MI.getOperand(i: 2).setReg(NewCst.getReg(Idx: 0));
5045 MI.setFlags(Flags);
5046 Observer.changedInstr(MI);
5047 Observer.changingInstr(MI&: *LHSPtrAdd);
5048 LHSPtrAdd->getOperand(i: 2).setReg(RHSReg);
5049 LHSPtrAdd->setFlags(Flags);
5050 Observer.changedInstr(MI&: *LHSPtrAdd);
5051 };
5052 return !reassociationCanBreakAddressingModePattern(MI);
5053}
5054
5055bool CombinerHelper::matchReassocFoldConstantsInSubTree(
5056 GPtrAdd &MI, MachineInstr *LHS, MachineInstr *RHS,
5057 BuildFnTy &MatchInfo) const {
5058 // G_PTR_ADD(G_PTR_ADD(BASE, C1), C2) -> G_PTR_ADD(BASE, C1+C2)
5059 auto *LHSPtrAdd = dyn_cast<GPtrAdd>(Val: LHS);
5060 if (!LHSPtrAdd)
5061 return false;
5062
5063 Register Src2Reg = MI.getOperand(i: 2).getReg();
5064 Register LHSSrc1 = LHSPtrAdd->getBaseReg();
5065 Register LHSSrc2 = LHSPtrAdd->getOffsetReg();
5066 auto C1 = getIConstantVRegVal(VReg: LHSSrc2, MRI);
5067 if (!C1)
5068 return false;
5069 auto C2 = getIConstantVRegVal(VReg: Src2Reg, MRI);
5070 if (!C2)
5071 return false;
5072
5073 // Reassociating nuw additions preserves nuw. If both original G_PTR_ADDs are
5074 // inbounds, reaching the same result in one G_PTR_ADD is also inbounds.
5075 // The nusw constraints are satisfied because imm1+imm2 cannot exceed the
5076 // largest signed integer that fits into the index type, which is the maximum
5077 // size of allocated objects according to the IR Language Reference.
5078 unsigned PtrAddFlags = MI.getFlags();
5079 unsigned LHSPtrAddFlags = LHSPtrAdd->getFlags();
5080 bool IsNoUWrap = PtrAddFlags & LHSPtrAddFlags & MachineInstr::MIFlag::NoUWrap;
5081 bool IsInBounds =
5082 PtrAddFlags & LHSPtrAddFlags & MachineInstr::MIFlag::InBounds;
5083 unsigned Flags = 0;
5084 if (IsNoUWrap)
5085 Flags |= MachineInstr::MIFlag::NoUWrap;
5086 if (IsInBounds) {
5087 Flags |= MachineInstr::MIFlag::InBounds;
5088 Flags |= MachineInstr::MIFlag::NoUSWrap;
5089 }
5090
5091 MatchInfo = [=, &MI](MachineIRBuilder &B) {
5092 auto NewCst = B.buildConstant(Res: MRI.getType(Reg: Src2Reg), Val: *C1 + *C2);
5093 Observer.changingInstr(MI);
5094 MI.getOperand(i: 1).setReg(LHSSrc1);
5095 MI.getOperand(i: 2).setReg(NewCst.getReg(Idx: 0));
5096 MI.setFlags(Flags);
5097 Observer.changedInstr(MI);
5098 };
5099 return !reassociationCanBreakAddressingModePattern(MI);
5100}
5101
5102bool CombinerHelper::matchReassocPtrAdd(MachineInstr &MI,
5103 BuildFnTy &MatchInfo) const {
5104 auto &PtrAdd = cast<GPtrAdd>(Val&: MI);
5105 // We're trying to match a few pointer computation patterns here for
5106 // re-association opportunities.
5107 // 1) Isolating a constant operand to be on the RHS, e.g.:
5108 // G_PTR_ADD(BASE, G_ADD(X, C)) -> G_PTR_ADD(G_PTR_ADD(BASE, X), C)
5109 //
5110 // 2) Folding two constants in each sub-tree as long as such folding
5111 // doesn't break a legal addressing mode.
5112 // G_PTR_ADD(G_PTR_ADD(BASE, C1), C2) -> G_PTR_ADD(BASE, C1+C2)
5113 //
5114 // 3) Move a constant from the LHS of an inner op to the RHS of the outer.
5115 // G_PTR_ADD (G_PTR_ADD X, C), Y) -> G_PTR_ADD (G_PTR_ADD(X, Y), C)
5116 // iif (G_PTR_ADD X, C) has one use.
5117 MachineInstr *LHS, *RHS;
5118 if (!mi_match(R: PtrAdd.getBaseReg(), MRI, P: m_MInstr(MI&: LHS)) ||
5119 !mi_match(R: PtrAdd.getOffsetReg(), MRI, P: m_MInstr(MI&: RHS)))
5120 return false;
5121
5122 // Try to match example 2.
5123 if (matchReassocFoldConstantsInSubTree(MI&: PtrAdd, LHS, RHS, MatchInfo))
5124 return true;
5125
5126 // Try to match example 3.
5127 if (matchReassocConstantInnerLHS(MI&: PtrAdd, LHS, RHS, MatchInfo))
5128 return true;
5129
5130 // Try to match example 1.
5131 if (matchReassocConstantInnerRHS(MI&: PtrAdd, RHS, MatchInfo))
5132 return true;
5133
5134 return false;
5135}
5136bool CombinerHelper::tryReassocBinOp(unsigned Opc, Register DstReg,
5137 Register OpLHS, Register OpRHS,
5138 BuildFnTy &MatchInfo) const {
5139 LLT OpRHSTy = MRI.getType(Reg: OpRHS);
5140 MachineInstr *OpLHSDef;
5141 if (!mi_match(R: OpLHS, MRI, P: m_MInstr(MI&: OpLHSDef)) || OpLHSDef->getOpcode() != Opc)
5142 return false;
5143
5144 Register OpLHSLHS = OpLHSDef->getOperand(i: 1).getReg();
5145 Register OpLHSRHS = OpLHSDef->getOperand(i: 2).getReg();
5146
5147 // If the inner op is (X op C), pull the constant out so it can be folded with
5148 // other constants in the expression tree. Folding is not guaranteed so we
5149 // might have (C1 op C2). In that case do not pull a constant out because it
5150 // won't help and can lead to infinite loops.
5151 if (isConstantOrConstantSplatVector(Def: OpLHSRHS, MRI) &&
5152 !isConstantOrConstantSplatVector(Def: OpLHSLHS, MRI)) {
5153 if (isConstantOrConstantSplatVector(Def: OpRHS, MRI)) {
5154 // (Opc (Opc X, C1), C2) -> (Opc X, (Opc C1, C2))
5155 MatchInfo = [=](MachineIRBuilder &B) {
5156 auto NewCst = B.buildInstr(Opc, DstOps: {OpRHSTy}, SrcOps: {OpLHSRHS, OpRHS});
5157 B.buildInstr(Opc, DstOps: {DstReg}, SrcOps: {OpLHSLHS, NewCst});
5158 };
5159 return true;
5160 }
5161 if (getTargetLowering().isReassocProfitable(MRI, N0: OpLHS, N1: OpRHS)) {
5162 // Reassociate: (op (op x, c1), y) -> (op (op x, y), c1)
5163 // iff (op x, c1) has one use
5164 MatchInfo = [=](MachineIRBuilder &B) {
5165 auto NewLHSLHS = B.buildInstr(Opc, DstOps: {OpRHSTy}, SrcOps: {OpLHSLHS, OpRHS});
5166 B.buildInstr(Opc, DstOps: {DstReg}, SrcOps: {NewLHSLHS, OpLHSRHS});
5167 };
5168 return true;
5169 }
5170 }
5171
5172 return false;
5173}
5174
5175bool CombinerHelper::matchReassocCommBinOp(MachineInstr &MI,
5176 BuildFnTy &MatchInfo) const {
5177 // We don't check if the reassociation will break a legal addressing mode
5178 // here since pointer arithmetic is handled by G_PTR_ADD.
5179 unsigned Opc = MI.getOpcode();
5180 Register DstReg = MI.getOperand(i: 0).getReg();
5181 Register LHSReg = MI.getOperand(i: 1).getReg();
5182 Register RHSReg = MI.getOperand(i: 2).getReg();
5183
5184 if (tryReassocBinOp(Opc, DstReg, OpLHS: LHSReg, OpRHS: RHSReg, MatchInfo))
5185 return true;
5186 if (tryReassocBinOp(Opc, DstReg, OpLHS: RHSReg, OpRHS: LHSReg, MatchInfo))
5187 return true;
5188 return false;
5189}
5190
5191bool CombinerHelper::matchConstantFoldCastOp(MachineInstr &MI,
5192 APInt &MatchInfo) const {
5193 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
5194 Register SrcOp = MI.getOperand(i: 1).getReg();
5195
5196 if (auto MaybeCst = ConstantFoldCastOp(Opcode: MI.getOpcode(), DstTy, Op0: SrcOp, MRI)) {
5197 MatchInfo = *MaybeCst;
5198 return true;
5199 }
5200
5201 return false;
5202}
5203
5204bool CombinerHelper::matchConstantFoldUnaryIntOp(MachineInstr &MI,
5205 BuildFnTy &MatchInfo) const {
5206 Register Dst = MI.getOperand(i: 0).getReg();
5207 auto Csts = ConstantFoldUnaryIntOp(Opcode: MI.getOpcode(), DstTy: MRI.getType(Reg: Dst),
5208 Src: MI.getOperand(i: 1).getReg(), MRI);
5209 if (Csts.empty())
5210 return false;
5211
5212 MatchInfo = [Dst, Csts = std::move(Csts)](MachineIRBuilder &B) {
5213 if (Csts.size() == 1)
5214 B.buildConstant(Res: Dst, Val: Csts[0]);
5215 else
5216 B.buildBuildVectorConstant(Res: Dst, Ops: Csts);
5217 };
5218 return true;
5219}
5220
5221bool CombinerHelper::matchConstantFoldBinOp(MachineInstr &MI,
5222 APInt &MatchInfo) const {
5223 Register Op1 = MI.getOperand(i: 1).getReg();
5224 Register Op2 = MI.getOperand(i: 2).getReg();
5225 auto MaybeCst = ConstantFoldBinOp(Opcode: MI.getOpcode(), Op1, Op2, MRI);
5226 if (!MaybeCst)
5227 return false;
5228 MatchInfo = *MaybeCst;
5229 return true;
5230}
5231
5232bool CombinerHelper::matchConstantFoldFPBinOp(MachineInstr &MI,
5233 ConstantFP *&MatchInfo) const {
5234 Register Op1 = MI.getOperand(i: 1).getReg();
5235 Register Op2 = MI.getOperand(i: 2).getReg();
5236 auto MaybeCst = ConstantFoldFPBinOp(Opcode: MI.getOpcode(), Op1, Op2, MRI);
5237 if (!MaybeCst)
5238 return false;
5239 MatchInfo =
5240 ConstantFP::get(Context&: MI.getMF()->getFunction().getContext(), V: *MaybeCst);
5241 return true;
5242}
5243
5244bool CombinerHelper::matchConstantFoldFMA(MachineInstr &MI,
5245 ConstantFP *&MatchInfo) const {
5246 assert(MI.getOpcode() == TargetOpcode::G_FMA ||
5247 MI.getOpcode() == TargetOpcode::G_FMAD);
5248 auto [_, Op1, Op2, Op3] = MI.getFirst4Regs();
5249
5250 const ConstantFP *Op3Cst = getConstantFPVRegVal(VReg: Op3, MRI);
5251 if (!Op3Cst)
5252 return false;
5253
5254 const ConstantFP *Op2Cst = getConstantFPVRegVal(VReg: Op2, MRI);
5255 if (!Op2Cst)
5256 return false;
5257
5258 const ConstantFP *Op1Cst = getConstantFPVRegVal(VReg: Op1, MRI);
5259 if (!Op1Cst)
5260 return false;
5261
5262 APFloat Op1F = Op1Cst->getValueAPF();
5263 Op1F.fusedMultiplyAdd(Multiplicand: Op2Cst->getValueAPF(), Addend: Op3Cst->getValueAPF(),
5264 RM: APFloat::rmNearestTiesToEven);
5265 MatchInfo = ConstantFP::get(Context&: MI.getMF()->getFunction().getContext(), V: Op1F);
5266 return true;
5267}
5268
5269bool CombinerHelper::matchNarrowBinopFeedingAnd(
5270 MachineInstr &MI,
5271 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
5272 // Look for a binop feeding into an AND with a mask:
5273 //
5274 // %add = G_ADD %lhs, %rhs
5275 // %and = G_AND %add, 000...11111111
5276 //
5277 // Check if it's possible to perform the binop at a narrower width and zext
5278 // back to the original width like so:
5279 //
5280 // %narrow_lhs = G_TRUNC %lhs
5281 // %narrow_rhs = G_TRUNC %rhs
5282 // %narrow_add = G_ADD %narrow_lhs, %narrow_rhs
5283 // %new_add = G_ZEXT %narrow_add
5284 // %and = G_AND %new_add, 000...11111111
5285 //
5286 // This can allow later combines to eliminate the G_AND if it turns out
5287 // that the mask is irrelevant.
5288 assert(MI.getOpcode() == TargetOpcode::G_AND);
5289 Register Dst = MI.getOperand(i: 0).getReg();
5290 Register AndLHS = MI.getOperand(i: 1).getReg();
5291 Register AndRHS = MI.getOperand(i: 2).getReg();
5292 LLT WideTy = MRI.getType(Reg: Dst);
5293
5294 // If the potential binop has more than one use, then it's possible that one
5295 // of those uses will need its full width.
5296 if (!WideTy.isScalar() || !MRI.hasOneNonDBGUse(RegNo: AndLHS))
5297 return false;
5298
5299 // Check if the LHS feeding the AND is impacted by the high bits that we're
5300 // masking out.
5301 //
5302 // e.g. for 64-bit x, y:
5303 //
5304 // add_64(x, y) & 65535 == zext(add_16(trunc(x), trunc(y))) & 65535
5305 MachineInstr *LHSInst = getDefIgnoringCopies(Reg: AndLHS, MRI);
5306 if (!LHSInst)
5307 return false;
5308 unsigned LHSOpc = LHSInst->getOpcode();
5309 switch (LHSOpc) {
5310 default:
5311 return false;
5312 case TargetOpcode::G_ADD:
5313 case TargetOpcode::G_SUB:
5314 case TargetOpcode::G_MUL:
5315 case TargetOpcode::G_AND:
5316 case TargetOpcode::G_OR:
5317 case TargetOpcode::G_XOR:
5318 break;
5319 }
5320
5321 // Find the mask on the RHS.
5322 auto Cst = getIConstantVRegValWithLookThrough(VReg: AndRHS, MRI);
5323 if (!Cst)
5324 return false;
5325 auto Mask = Cst->Value;
5326 if (!Mask.isMask())
5327 return false;
5328
5329 // No point in combining if there's nothing to truncate.
5330 unsigned NarrowWidth = Mask.countr_one();
5331 if (NarrowWidth == WideTy.getSizeInBits())
5332 return false;
5333 LLT NarrowTy = LLT::integer(SizeInBits: NarrowWidth);
5334
5335 // Check if adding the zext + truncates could be harmful.
5336 auto &MF = *MI.getMF();
5337 const auto &TLI = getTargetLowering();
5338 LLVMContext &Ctx = MF.getFunction().getContext();
5339 if (!TLI.isTruncateFree(FromTy: WideTy, ToTy: NarrowTy, Ctx) ||
5340 !TLI.isZExtFree(FromTy: NarrowTy, ToTy: WideTy, Ctx))
5341 return false;
5342 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_TRUNC, {NarrowTy, WideTy}}) ||
5343 !isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_ZEXT, {WideTy, NarrowTy}}))
5344 return false;
5345 Register BinOpLHS = LHSInst->getOperand(i: 1).getReg();
5346 Register BinOpRHS = LHSInst->getOperand(i: 2).getReg();
5347 MatchInfo = [=, &MI](MachineIRBuilder &B) {
5348 auto NarrowLHS = Builder.buildTrunc(Res: NarrowTy, Op: BinOpLHS);
5349 auto NarrowRHS = Builder.buildTrunc(Res: NarrowTy, Op: BinOpRHS);
5350 auto NarrowBinOp =
5351 Builder.buildInstr(Opc: LHSOpc, DstOps: {NarrowTy}, SrcOps: {NarrowLHS, NarrowRHS});
5352 auto Ext = Builder.buildZExt(Res: WideTy, Op: NarrowBinOp);
5353 Observer.changingInstr(MI);
5354 MI.getOperand(i: 1).setReg(Ext.getReg(Idx: 0));
5355 Observer.changedInstr(MI);
5356 };
5357 return true;
5358}
5359
5360bool CombinerHelper::matchMulOBy2(MachineInstr &MI,
5361 BuildFnTy &MatchInfo) const {
5362 unsigned Opc = MI.getOpcode();
5363 assert(Opc == TargetOpcode::G_UMULO || Opc == TargetOpcode::G_SMULO);
5364
5365 if (!mi_match(R: MI.getOperand(i: 3).getReg(), MRI, P: m_SpecificICstOrSplat(RequestedValue: 2)))
5366 return false;
5367
5368 MatchInfo = [=, &MI](MachineIRBuilder &B) {
5369 Observer.changingInstr(MI);
5370 unsigned NewOpc = Opc == TargetOpcode::G_UMULO ? TargetOpcode::G_UADDO
5371 : TargetOpcode::G_SADDO;
5372 MI.setDesc(Builder.getTII().get(Opcode: NewOpc));
5373 MI.getOperand(i: 3).setReg(MI.getOperand(i: 2).getReg());
5374 Observer.changedInstr(MI);
5375 };
5376 return true;
5377}
5378
5379bool CombinerHelper::matchMulOBy0(MachineInstr &MI,
5380 BuildFnTy &MatchInfo) const {
5381 // (G_*MULO x, 0) -> 0 + no carry out
5382 assert(MI.getOpcode() == TargetOpcode::G_UMULO ||
5383 MI.getOpcode() == TargetOpcode::G_SMULO);
5384 if (!mi_match(R: MI.getOperand(i: 3).getReg(), MRI, P: m_SpecificICstOrSplat(RequestedValue: 0)))
5385 return false;
5386 Register Dst = MI.getOperand(i: 0).getReg();
5387 Register Carry = MI.getOperand(i: 1).getReg();
5388 if (!isConstantLegalOrBeforeLegalizer(Ty: MRI.getType(Reg: Dst)) ||
5389 !isConstantLegalOrBeforeLegalizer(Ty: MRI.getType(Reg: Carry)))
5390 return false;
5391 MatchInfo = [=](MachineIRBuilder &B) {
5392 B.buildConstant(Res: Dst, Val: 0);
5393 B.buildConstant(Res: Carry, Val: 0);
5394 };
5395 return true;
5396}
5397
5398bool CombinerHelper::matchAddEToAddO(MachineInstr &MI,
5399 BuildFnTy &MatchInfo) const {
5400 // (G_*ADDE x, y, 0) -> (G_*ADDO x, y)
5401 // (G_*SUBE x, y, 0) -> (G_*SUBO x, y)
5402 assert(MI.getOpcode() == TargetOpcode::G_UADDE ||
5403 MI.getOpcode() == TargetOpcode::G_SADDE ||
5404 MI.getOpcode() == TargetOpcode::G_USUBE ||
5405 MI.getOpcode() == TargetOpcode::G_SSUBE);
5406 if (!mi_match(R: MI.getOperand(i: 4).getReg(), MRI, P: m_SpecificICstOrSplat(RequestedValue: 0)))
5407 return false;
5408 MatchInfo = [&](MachineIRBuilder &B) {
5409 unsigned NewOpcode;
5410 switch (MI.getOpcode()) {
5411 case TargetOpcode::G_UADDE:
5412 NewOpcode = TargetOpcode::G_UADDO;
5413 break;
5414 case TargetOpcode::G_SADDE:
5415 NewOpcode = TargetOpcode::G_SADDO;
5416 break;
5417 case TargetOpcode::G_USUBE:
5418 NewOpcode = TargetOpcode::G_USUBO;
5419 break;
5420 case TargetOpcode::G_SSUBE:
5421 NewOpcode = TargetOpcode::G_SSUBO;
5422 break;
5423 }
5424 Observer.changingInstr(MI);
5425 MI.setDesc(B.getTII().get(Opcode: NewOpcode));
5426 MI.removeOperand(OpNo: 4);
5427 Observer.changedInstr(MI);
5428 };
5429 return true;
5430}
5431
5432bool CombinerHelper::matchSubAddSameReg(MachineInstr &MI,
5433 BuildFnTy &MatchInfo) const {
5434 assert(MI.getOpcode() == TargetOpcode::G_SUB);
5435 Register Dst = MI.getOperand(i: 0).getReg();
5436 // (x + y) - z -> x (if y == z)
5437 // (x + y) - z -> y (if x == z)
5438 Register X, Y, Z;
5439 if (mi_match(R: Dst, MRI, P: m_GSub(L: m_GAdd(L: m_Reg(R&: X), R: m_Reg(R&: Y)), R: m_Reg(R&: Z)))) {
5440 Register ReplaceReg;
5441 int64_t CstX, CstY;
5442 if (Y == Z || (mi_match(R: Y, MRI, P: m_ICstOrSplat(Cst&: CstY)) &&
5443 mi_match(R: Z, MRI, P: m_SpecificICstOrSplat(RequestedValue: CstY))))
5444 ReplaceReg = X;
5445 else if (X == Z || (mi_match(R: X, MRI, P: m_ICstOrSplat(Cst&: CstX)) &&
5446 mi_match(R: Z, MRI, P: m_SpecificICstOrSplat(RequestedValue: CstX))))
5447 ReplaceReg = Y;
5448 if (ReplaceReg) {
5449 MatchInfo = [=](MachineIRBuilder &B) { B.buildCopy(Res: Dst, Op: ReplaceReg); };
5450 return true;
5451 }
5452 }
5453
5454 // x - (y + z) -> 0 - y (if x == z)
5455 // x - (y + z) -> 0 - z (if x == y)
5456 if (mi_match(R: Dst, MRI, P: m_GSub(L: m_Reg(R&: X), R: m_GAdd(L: m_Reg(R&: Y), R: m_Reg(R&: Z))))) {
5457 Register ReplaceReg;
5458 int64_t CstX;
5459 if (X == Z || (mi_match(R: X, MRI, P: m_ICstOrSplat(Cst&: CstX)) &&
5460 mi_match(R: Z, MRI, P: m_SpecificICstOrSplat(RequestedValue: CstX))))
5461 ReplaceReg = Y;
5462 else if (X == Y || (mi_match(R: X, MRI, P: m_ICstOrSplat(Cst&: CstX)) &&
5463 mi_match(R: Y, MRI, P: m_SpecificICstOrSplat(RequestedValue: CstX))))
5464 ReplaceReg = Z;
5465 if (ReplaceReg) {
5466 MatchInfo = [=](MachineIRBuilder &B) {
5467 auto Zero = B.buildConstant(Res: MRI.getType(Reg: Dst), Val: 0);
5468 B.buildSub(Dst, Src0: Zero, Src1: ReplaceReg);
5469 };
5470 return true;
5471 }
5472 }
5473 return false;
5474}
5475
5476MachineInstr *CombinerHelper::buildUDivOrURemUsingMul(MachineInstr &MI) const {
5477 unsigned Opcode = MI.getOpcode();
5478 assert(Opcode == TargetOpcode::G_UDIV || Opcode == TargetOpcode::G_UREM);
5479 auto &UDivorRem = cast<GenericMachineInstr>(Val&: MI);
5480 Register Dst = UDivorRem.getReg(Idx: 0);
5481 Register LHS = UDivorRem.getReg(Idx: 1);
5482 Register RHS = UDivorRem.getReg(Idx: 2);
5483 LLT Ty = MRI.getType(Reg: Dst);
5484 LLT ScalarTy = Ty.getScalarType();
5485 const unsigned EltBits = ScalarTy.getScalarSizeInBits();
5486 LLT ShiftAmtTy = getTargetLowering().getPreferredShiftAmountTy(ShiftValueTy: Ty);
5487 LLT ScalarShiftAmtTy = ShiftAmtTy.getScalarType();
5488
5489 auto &MIB = Builder;
5490
5491 bool UseSRL = false;
5492 SmallVector<Register, 16> Shifts, Factors;
5493 auto *RHSDefInstr = cast<GenericMachineInstr>(Val: getDefIgnoringCopies(Reg: RHS, MRI));
5494 bool IsSplat = getIConstantSplatVal(MI: *RHSDefInstr, MRI).has_value();
5495
5496 auto BuildExactUDIVPattern = [&](const Constant *C) {
5497 // Don't recompute inverses for each splat element.
5498 if (IsSplat && !Factors.empty()) {
5499 Shifts.push_back(Elt: Shifts[0]);
5500 Factors.push_back(Elt: Factors[0]);
5501 return true;
5502 }
5503
5504 auto *CI = cast<ConstantInt>(Val: C);
5505 APInt Divisor = CI->getValue();
5506 unsigned Shift = Divisor.countr_zero();
5507 if (Shift) {
5508 Divisor.lshrInPlace(ShiftAmt: Shift);
5509 UseSRL = true;
5510 }
5511
5512 // Calculate the multiplicative inverse modulo BW.
5513 APInt Factor = Divisor.multiplicativeInverse();
5514 Shifts.push_back(Elt: MIB.buildConstant(Res: ScalarShiftAmtTy, Val: Shift).getReg(Idx: 0));
5515 Factors.push_back(Elt: MIB.buildConstant(Res: ScalarTy, Val: Factor).getReg(Idx: 0));
5516 return true;
5517 };
5518
5519 if (MI.getFlag(Flag: MachineInstr::MIFlag::IsExact)) {
5520 // Collect all magic values from the build vector.
5521 if (!matchUnaryPredicate(MRI, Reg: RHS, Match: BuildExactUDIVPattern))
5522 llvm_unreachable("Expected unary predicate match to succeed");
5523
5524 Register Shift, Factor;
5525 if (Ty.isVector()) {
5526 Shift = MIB.buildBuildVector(Res: ShiftAmtTy, Ops: Shifts).getReg(Idx: 0);
5527 Factor = MIB.buildBuildVector(Res: Ty, Ops: Factors).getReg(Idx: 0);
5528 } else {
5529 Shift = Shifts[0];
5530 Factor = Factors[0];
5531 }
5532
5533 Register Res = LHS;
5534
5535 if (UseSRL)
5536 Res = MIB.buildLShr(Dst: Ty, Src0: Res, Src1: Shift, Flags: MachineInstr::IsExact).getReg(Idx: 0);
5537
5538 return MIB.buildMul(Dst: Ty, Src0: Res, Src1: Factor);
5539 }
5540
5541 unsigned KnownLeadingZeros =
5542 VT ? VT->getKnownBits(R: LHS).countMinLeadingZeros() : 0;
5543
5544 bool UseNPQ = false;
5545 SmallVector<Register, 16> PreShifts, PostShifts, MagicFactors, NPQFactors;
5546 auto BuildUDIVPattern = [&](const Constant *C) {
5547 auto *CI = cast<ConstantInt>(Val: C);
5548 const APInt &Divisor = CI->getValue();
5549
5550 bool SelNPQ = false;
5551 APInt Magic(Divisor.getBitWidth(), 0);
5552 unsigned PreShift = 0, PostShift = 0;
5553
5554 // Magic algorithm doesn't work for division by 1. We need to emit a select
5555 // at the end.
5556 // TODO: Use undef values for divisor of 1.
5557 if (!Divisor.isOne()) {
5558
5559 // UnsignedDivisionByConstantInfo doesn't work correctly if leading zeros
5560 // in the dividend exceeds the leading zeros for the divisor.
5561 UnsignedDivisionByConstantInfo magics =
5562 UnsignedDivisionByConstantInfo::get(
5563 D: Divisor, LeadingZeros: std::min(a: KnownLeadingZeros, b: Divisor.countl_zero()));
5564
5565 Magic = std::move(magics.Magic);
5566
5567 assert(magics.PreShift < Divisor.getBitWidth() &&
5568 "We shouldn't generate an undefined shift!");
5569 assert(magics.PostShift < Divisor.getBitWidth() &&
5570 "We shouldn't generate an undefined shift!");
5571 assert((!magics.IsAdd || magics.PreShift == 0) && "Unexpected pre-shift");
5572 PreShift = magics.PreShift;
5573 PostShift = magics.PostShift;
5574 SelNPQ = magics.IsAdd;
5575 }
5576
5577 PreShifts.push_back(
5578 Elt: MIB.buildConstant(Res: ScalarShiftAmtTy, Val: PreShift).getReg(Idx: 0));
5579 MagicFactors.push_back(Elt: MIB.buildConstant(Res: ScalarTy, Val: Magic).getReg(Idx: 0));
5580 NPQFactors.push_back(
5581 Elt: MIB.buildConstant(Res: ScalarTy,
5582 Val: SelNPQ ? APInt::getOneBitSet(numBits: EltBits, BitNo: EltBits - 1)
5583 : APInt::getZero(numBits: EltBits))
5584 .getReg(Idx: 0));
5585 PostShifts.push_back(
5586 Elt: MIB.buildConstant(Res: ScalarShiftAmtTy, Val: PostShift).getReg(Idx: 0));
5587 UseNPQ |= SelNPQ;
5588 return true;
5589 };
5590
5591 // Collect the shifts/magic values from each element.
5592 bool Matched = matchUnaryPredicate(MRI, Reg: RHS, Match: BuildUDIVPattern);
5593 (void)Matched;
5594 assert(Matched && "Expected unary predicate match to succeed");
5595
5596 Register PreShift, PostShift, MagicFactor, NPQFactor;
5597 auto *RHSDef = getOpcodeDef<GBuildVector>(Reg: RHS, MRI);
5598 if (RHSDef) {
5599 PreShift = MIB.buildBuildVector(Res: ShiftAmtTy, Ops: PreShifts).getReg(Idx: 0);
5600 MagicFactor = MIB.buildBuildVector(Res: Ty, Ops: MagicFactors).getReg(Idx: 0);
5601 NPQFactor = MIB.buildBuildVector(Res: Ty, Ops: NPQFactors).getReg(Idx: 0);
5602 PostShift = MIB.buildBuildVector(Res: ShiftAmtTy, Ops: PostShifts).getReg(Idx: 0);
5603 } else {
5604 assert(MRI.getType(RHS).isScalar() &&
5605 "Non-build_vector operation should have been a scalar");
5606 PreShift = PreShifts[0];
5607 MagicFactor = MagicFactors[0];
5608 PostShift = PostShifts[0];
5609 }
5610
5611 Register Q = LHS;
5612 Q = MIB.buildLShr(Dst: Ty, Src0: Q, Src1: PreShift).getReg(Idx: 0);
5613
5614 // Multiply the numerator (operand 0) by the magic value.
5615 Q = MIB.buildUMulH(Dst: Ty, Src0: Q, Src1: MagicFactor).getReg(Idx: 0);
5616
5617 if (UseNPQ) {
5618 Register NPQ = MIB.buildSub(Dst: Ty, Src0: LHS, Src1: Q).getReg(Idx: 0);
5619
5620 // For vectors we might have a mix of non-NPQ/NPQ paths, so use
5621 // G_UMULH to act as a SRL-by-1 for NPQ, else multiply by zero.
5622 if (Ty.isVector())
5623 NPQ = MIB.buildUMulH(Dst: Ty, Src0: NPQ, Src1: NPQFactor).getReg(Idx: 0);
5624 else
5625 NPQ = MIB.buildLShr(Dst: Ty, Src0: NPQ, Src1: MIB.buildConstant(Res: ShiftAmtTy, Val: 1)).getReg(Idx: 0);
5626
5627 Q = MIB.buildAdd(Dst: Ty, Src0: NPQ, Src1: Q).getReg(Idx: 0);
5628 }
5629
5630 Q = MIB.buildLShr(Dst: Ty, Src0: Q, Src1: PostShift).getReg(Idx: 0);
5631 auto One = MIB.buildConstant(Res: Ty, Val: 1);
5632 auto IsOne = MIB.buildICmp(
5633 Pred: CmpInst::Predicate::ICMP_EQ,
5634 Res: Ty.isScalar() ? LLT::integer(SizeInBits: 1) : Ty.changeElementType(NewEltTy: LLT::integer(SizeInBits: 1)),
5635 Op0: RHS, Op1: One);
5636 auto ret = MIB.buildSelect(Res: Ty, Tst: IsOne, Op0: LHS, Op1: Q);
5637
5638 if (Opcode == TargetOpcode::G_UREM) {
5639 auto Prod = MIB.buildMul(Dst: Ty, Src0: ret, Src1: RHS);
5640 return MIB.buildSub(Dst: Ty, Src0: LHS, Src1: Prod);
5641 }
5642 return ret;
5643}
5644
5645bool CombinerHelper::matchUDivOrURemByConst(MachineInstr &MI) const {
5646 unsigned Opcode = MI.getOpcode();
5647 assert(Opcode == TargetOpcode::G_UDIV || Opcode == TargetOpcode::G_UREM);
5648 Register Dst = MI.getOperand(i: 0).getReg();
5649 Register RHS = MI.getOperand(i: 2).getReg();
5650 LLT DstTy = MRI.getType(Reg: Dst);
5651
5652 auto &MF = *MI.getMF();
5653 AttributeList Attr = MF.getFunction().getAttributes();
5654 const auto &TLI = getTargetLowering();
5655 LLVMContext &Ctx = MF.getFunction().getContext();
5656 if (DstTy.getScalarSizeInBits() == 1 ||
5657 TLI.isIntDivCheap(VT: getApproximateEVTForLLT(Ty: DstTy, Ctx), Attr))
5658 return false;
5659
5660 // Don't do this for minsize because the instruction sequence is usually
5661 // larger.
5662 if (MF.getFunction().hasMinSize())
5663 return false;
5664
5665 if (Opcode == TargetOpcode::G_UDIV &&
5666 MI.getFlag(Flag: MachineInstr::MIFlag::IsExact)) {
5667 return matchUnaryPredicate(
5668 MRI, Reg: RHS, Match: [](const Constant *C) { return C && !C->isNullValue(); });
5669 }
5670
5671 MachineInstr *RHSDef;
5672 if (!mi_match(R: RHS, MRI, P: m_MInstr(MI&: RHSDef)) ||
5673 !isConstantOrConstantVector(MI&: *RHSDef, MRI))
5674 return false;
5675
5676 // Don't do this if the types are not going to be legal.
5677 if (LI) {
5678 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_MUL, {DstTy, DstTy}}))
5679 return false;
5680 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_UMULH, {DstTy}}))
5681 return false;
5682 if (!isLegalOrBeforeLegalizer(
5683 Query: {TargetOpcode::G_ICMP,
5684 {DstTy.isVector() ? DstTy.changeElementSize(NewEltSize: 1) : LLT::scalar(SizeInBits: 1),
5685 DstTy}}))
5686 return false;
5687 if (Opcode == TargetOpcode::G_UREM &&
5688 !isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_SUB, {DstTy, DstTy}}))
5689 return false;
5690 }
5691
5692 return matchUnaryPredicate(
5693 MRI, Reg: RHS, Match: [](const Constant *C) { return C && !C->isNullValue(); });
5694}
5695
5696void CombinerHelper::applyUDivOrURemByConst(MachineInstr &MI) const {
5697 auto *NewMI = buildUDivOrURemUsingMul(MI);
5698 replaceSingleDefInstWithReg(MI, Replacement: NewMI->getOperand(i: 0).getReg());
5699}
5700
5701bool CombinerHelper::matchSDivOrSRemByConst(MachineInstr &MI) const {
5702 unsigned Opcode = MI.getOpcode();
5703 assert(Opcode == TargetOpcode::G_SDIV || Opcode == TargetOpcode::G_SREM);
5704 Register Dst = MI.getOperand(i: 0).getReg();
5705 Register RHS = MI.getOperand(i: 2).getReg();
5706 LLT DstTy = MRI.getType(Reg: Dst);
5707 auto SizeInBits = DstTy.getScalarSizeInBits();
5708 LLT WideTy = DstTy.changeElementSize(NewEltSize: SizeInBits * 2);
5709
5710 auto &MF = *MI.getMF();
5711 AttributeList Attr = MF.getFunction().getAttributes();
5712 const auto &TLI = getTargetLowering();
5713 LLVMContext &Ctx = MF.getFunction().getContext();
5714 if (DstTy.getScalarSizeInBits() < 3 ||
5715 TLI.isIntDivCheap(VT: getApproximateEVTForLLT(Ty: DstTy, Ctx), Attr))
5716 return false;
5717
5718 // Don't do this for minsize because the instruction sequence is usually
5719 // larger.
5720 if (MF.getFunction().hasMinSize())
5721 return false;
5722
5723 // If the sdiv has an 'exact' flag we can use a simpler lowering.
5724 if (Opcode == TargetOpcode::G_SDIV &&
5725 MI.getFlag(Flag: MachineInstr::MIFlag::IsExact)) {
5726 return matchUnaryPredicate(
5727 MRI, Reg: RHS, Match: [](const Constant *C) { return C && !C->isNullValue(); });
5728 }
5729
5730 MachineInstr *RHSDef;
5731 if (!mi_match(R: RHS, MRI, P: m_MInstr(MI&: RHSDef)) ||
5732 !isConstantOrConstantVector(MI&: *RHSDef, MRI))
5733 return false;
5734
5735 // Don't do this if the types are not going to be legal.
5736 if (LI) {
5737 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_MUL, {DstTy, DstTy}}))
5738 return false;
5739 if (!isLegal(Query: {TargetOpcode::G_SMULH, {DstTy}}) &&
5740 !isLegalOrHasWidenScalar(Query: {TargetOpcode::G_MUL, {WideTy, WideTy}}))
5741 return false;
5742 if (Opcode == TargetOpcode::G_SREM &&
5743 !isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_SUB, {DstTy, DstTy}}))
5744 return false;
5745 }
5746
5747 return matchUnaryPredicate(
5748 MRI, Reg: RHS, Match: [](const Constant *C) { return C && !C->isNullValue(); });
5749}
5750
5751void CombinerHelper::applySDivOrSRemByConst(MachineInstr &MI) const {
5752 auto *NewMI = buildSDivOrSRemUsingMul(MI);
5753 replaceSingleDefInstWithReg(MI, Replacement: NewMI->getOperand(i: 0).getReg());
5754}
5755
5756MachineInstr *CombinerHelper::buildSDivOrSRemUsingMul(MachineInstr &MI) const {
5757 unsigned Opcode = MI.getOpcode();
5758 assert(MI.getOpcode() == TargetOpcode::G_SDIV ||
5759 Opcode == TargetOpcode::G_SREM);
5760 auto &SDivorRem = cast<GenericMachineInstr>(Val&: MI);
5761 Register Dst = SDivorRem.getReg(Idx: 0);
5762 Register LHS = SDivorRem.getReg(Idx: 1);
5763 Register RHS = SDivorRem.getReg(Idx: 2);
5764 LLT Ty = MRI.getType(Reg: Dst);
5765 LLT ScalarTy = Ty.getScalarType();
5766 const unsigned EltBits = ScalarTy.getScalarSizeInBits();
5767 LLT ShiftAmtTy = getTargetLowering().getPreferredShiftAmountTy(ShiftValueTy: Ty);
5768 LLT ScalarShiftAmtTy = ShiftAmtTy.getScalarType();
5769 auto &MIB = Builder;
5770
5771 bool UseSRA = false;
5772 SmallVector<Register, 16> ExactShifts, ExactFactors;
5773
5774 auto *RHSDefInstr = cast<GenericMachineInstr>(Val: getDefIgnoringCopies(Reg: RHS, MRI));
5775 bool IsSplat = getIConstantSplatVal(MI: *RHSDefInstr, MRI).has_value();
5776
5777 auto BuildExactSDIVPattern = [&](const Constant *C) {
5778 // Don't recompute inverses for each splat element.
5779 if (IsSplat && !ExactFactors.empty()) {
5780 ExactShifts.push_back(Elt: ExactShifts[0]);
5781 ExactFactors.push_back(Elt: ExactFactors[0]);
5782 return true;
5783 }
5784
5785 auto *CI = cast<ConstantInt>(Val: C);
5786 APInt Divisor = CI->getValue();
5787 unsigned Shift = Divisor.countr_zero();
5788 if (Shift) {
5789 Divisor.ashrInPlace(ShiftAmt: Shift);
5790 UseSRA = true;
5791 }
5792
5793 // Calculate the multiplicative inverse modulo BW.
5794 // 2^W requires W + 1 bits, so we have to extend and then truncate.
5795 APInt Factor = Divisor.multiplicativeInverse();
5796 ExactShifts.push_back(Elt: MIB.buildConstant(Res: ScalarShiftAmtTy, Val: Shift).getReg(Idx: 0));
5797 ExactFactors.push_back(Elt: MIB.buildConstant(Res: ScalarTy, Val: Factor).getReg(Idx: 0));
5798 return true;
5799 };
5800
5801 if (MI.getFlag(Flag: MachineInstr::MIFlag::IsExact)) {
5802 // Collect all magic values from the build vector.
5803 bool Matched = matchUnaryPredicate(MRI, Reg: RHS, Match: BuildExactSDIVPattern);
5804 (void)Matched;
5805 assert(Matched && "Expected unary predicate match to succeed");
5806
5807 Register Shift, Factor;
5808 if (Ty.isVector()) {
5809 Shift = MIB.buildBuildVector(Res: ShiftAmtTy, Ops: ExactShifts).getReg(Idx: 0);
5810 Factor = MIB.buildBuildVector(Res: Ty, Ops: ExactFactors).getReg(Idx: 0);
5811 } else {
5812 Shift = ExactShifts[0];
5813 Factor = ExactFactors[0];
5814 }
5815
5816 Register Res = LHS;
5817
5818 if (UseSRA)
5819 Res = MIB.buildAShr(Dst: Ty, Src0: Res, Src1: Shift, Flags: MachineInstr::IsExact).getReg(Idx: 0);
5820
5821 return MIB.buildMul(Dst: Ty, Src0: Res, Src1: Factor);
5822 }
5823
5824 SmallVector<Register, 16> MagicFactors, Factors, Shifts, ShiftMasks;
5825
5826 auto BuildSDIVPattern = [&](const Constant *C) {
5827 auto *CI = cast<ConstantInt>(Val: C);
5828 const APInt &Divisor = CI->getValue();
5829
5830 SignedDivisionByConstantInfo Magics =
5831 SignedDivisionByConstantInfo::get(D: Divisor);
5832 int NumeratorFactor = 0;
5833 int ShiftMask = -1;
5834
5835 if (Divisor.isOne() || Divisor.isAllOnes()) {
5836 // If d is +1/-1, we just multiply the numerator by +1/-1.
5837 NumeratorFactor = Divisor.getSExtValue();
5838 Magics.Magic = 0;
5839 Magics.ShiftAmount = 0;
5840 ShiftMask = 0;
5841 } else if (Divisor.isStrictlyPositive() && Magics.Magic.isNegative()) {
5842 // If d > 0 and m < 0, add the numerator.
5843 NumeratorFactor = 1;
5844 } else if (Divisor.isNegative() && Magics.Magic.isStrictlyPositive()) {
5845 // If d < 0 and m > 0, subtract the numerator.
5846 NumeratorFactor = -1;
5847 }
5848
5849 MagicFactors.push_back(Elt: MIB.buildConstant(Res: ScalarTy, Val: Magics.Magic).getReg(Idx: 0));
5850 Factors.push_back(Elt: MIB.buildConstant(Res: ScalarTy, Val: NumeratorFactor).getReg(Idx: 0));
5851 Shifts.push_back(
5852 Elt: MIB.buildConstant(Res: ScalarShiftAmtTy, Val: Magics.ShiftAmount).getReg(Idx: 0));
5853 ShiftMasks.push_back(Elt: MIB.buildConstant(Res: ScalarTy, Val: ShiftMask).getReg(Idx: 0));
5854
5855 return true;
5856 };
5857
5858 // Collect the shifts/magic values from each element.
5859 bool Matched = matchUnaryPredicate(MRI, Reg: RHS, Match: BuildSDIVPattern);
5860 (void)Matched;
5861 assert(Matched && "Expected unary predicate match to succeed");
5862
5863 Register MagicFactor, Factor, Shift, ShiftMask;
5864 auto *RHSDef = getOpcodeDef<GBuildVector>(Reg: RHS, MRI);
5865 if (RHSDef) {
5866 MagicFactor = MIB.buildBuildVector(Res: Ty, Ops: MagicFactors).getReg(Idx: 0);
5867 Factor = MIB.buildBuildVector(Res: Ty, Ops: Factors).getReg(Idx: 0);
5868 Shift = MIB.buildBuildVector(Res: ShiftAmtTy, Ops: Shifts).getReg(Idx: 0);
5869 ShiftMask = MIB.buildBuildVector(Res: Ty, Ops: ShiftMasks).getReg(Idx: 0);
5870 } else {
5871 assert(MRI.getType(RHS).isScalar() &&
5872 "Non-build_vector operation should have been a scalar");
5873 MagicFactor = MagicFactors[0];
5874 Factor = Factors[0];
5875 Shift = Shifts[0];
5876 ShiftMask = ShiftMasks[0];
5877 }
5878
5879 Register Q = LHS;
5880 Q = MIB.buildSMulH(Dst: Ty, Src0: LHS, Src1: MagicFactor).getReg(Idx: 0);
5881
5882 // (Optionally) Add/subtract the numerator using Factor.
5883 Factor = MIB.buildMul(Dst: Ty, Src0: LHS, Src1: Factor).getReg(Idx: 0);
5884 Q = MIB.buildAdd(Dst: Ty, Src0: Q, Src1: Factor).getReg(Idx: 0);
5885
5886 // Shift right algebraic by shift value.
5887 Q = MIB.buildAShr(Dst: Ty, Src0: Q, Src1: Shift).getReg(Idx: 0);
5888
5889 // Extract the sign bit, mask it and add it to the quotient.
5890 auto SignShift = MIB.buildConstant(Res: ShiftAmtTy, Val: EltBits - 1);
5891 auto T = MIB.buildLShr(Dst: Ty, Src0: Q, Src1: SignShift);
5892 T = MIB.buildAnd(Dst: Ty, Src0: T, Src1: ShiftMask);
5893 auto ret = MIB.buildAdd(Dst: Ty, Src0: Q, Src1: T);
5894
5895 if (Opcode == TargetOpcode::G_SREM) {
5896 auto Prod = MIB.buildMul(Dst: Ty, Src0: ret, Src1: RHS);
5897 return MIB.buildSub(Dst: Ty, Src0: LHS, Src1: Prod);
5898 }
5899 return ret;
5900}
5901
5902bool CombinerHelper::matchDivByPow2(MachineInstr &MI, bool IsSigned) const {
5903 assert((MI.getOpcode() == TargetOpcode::G_SDIV ||
5904 MI.getOpcode() == TargetOpcode::G_UDIV) &&
5905 "Expected SDIV or UDIV");
5906 auto &Div = cast<GenericMachineInstr>(Val&: MI);
5907 Register RHS = Div.getReg(Idx: 2);
5908 auto MatchPow2 = [&](const Constant *C) {
5909 auto *CI = dyn_cast<ConstantInt>(Val: C);
5910 return CI && (CI->getValue().isPowerOf2() ||
5911 (IsSigned && CI->getValue().isNegatedPowerOf2()));
5912 };
5913 return matchUnaryPredicate(MRI, Reg: RHS, Match: MatchPow2, /*AllowUndefs=*/false);
5914}
5915
5916void CombinerHelper::applySDivByPow2(MachineInstr &MI) const {
5917 assert(MI.getOpcode() == TargetOpcode::G_SDIV && "Expected SDIV");
5918 auto &SDiv = cast<GenericMachineInstr>(Val&: MI);
5919 Register Dst = SDiv.getReg(Idx: 0);
5920 Register LHS = SDiv.getReg(Idx: 1);
5921 Register RHS = SDiv.getReg(Idx: 2);
5922 LLT Ty = MRI.getType(Reg: Dst);
5923 LLT ShiftAmtTy = getTargetLowering().getPreferredShiftAmountTy(ShiftValueTy: Ty);
5924 LLT CCVT = Ty.isVector() ? LLT::vector(EC: Ty.getElementCount(), ScalarTy: LLT::integer(SizeInBits: 1))
5925 : LLT::integer(SizeInBits: 1);
5926
5927 // Effectively we want to lower G_SDIV %lhs, %rhs, where %rhs is a power of 2,
5928 // to the following version:
5929 //
5930 // %c1 = G_CTTZ %rhs
5931 // %inexact = G_SUB $bitwidth, %c1
5932 // %sign = %G_ASHR %lhs, $(bitwidth - 1)
5933 // %lshr = G_LSHR %sign, %inexact
5934 // %add = G_ADD %lhs, %lshr
5935 // %ashr = G_ASHR %add, %c1
5936 // %ashr = G_SELECT, %isoneorallones, %lhs, %ashr
5937 // %zero = G_CONSTANT $0
5938 // %neg = G_NEG %ashr
5939 // %isneg = G_ICMP SLT %rhs, %zero
5940 // %res = G_SELECT %isneg, %neg, %ashr
5941
5942 unsigned BitWidth = Ty.getScalarSizeInBits();
5943 auto Zero = Builder.buildConstant(Res: Ty, Val: 0);
5944
5945 auto Bits = Builder.buildConstant(Res: ShiftAmtTy, Val: BitWidth);
5946 auto C1 = Builder.buildCTTZ(Dst: ShiftAmtTy, Src0: RHS);
5947 auto Inexact = Builder.buildSub(Dst: ShiftAmtTy, Src0: Bits, Src1: C1);
5948 // Splat the sign bit into the register
5949 auto Sign = Builder.buildAShr(
5950 Dst: Ty, Src0: LHS, Src1: Builder.buildConstant(Res: ShiftAmtTy, Val: BitWidth - 1));
5951
5952 // Add (LHS < 0) ? abs2 - 1 : 0;
5953 auto LSrl = Builder.buildLShr(Dst: Ty, Src0: Sign, Src1: Inexact);
5954 auto Add = Builder.buildAdd(Dst: Ty, Src0: LHS, Src1: LSrl);
5955 auto AShr = Builder.buildAShr(Dst: Ty, Src0: Add, Src1: C1);
5956
5957 // Special case: (sdiv X, 1) -> X
5958 // Special Case: (sdiv X, -1) -> 0-X
5959 auto One = Builder.buildConstant(Res: Ty, Val: 1);
5960 auto MinusOne = Builder.buildConstant(Res: Ty, Val: -1);
5961 auto IsOne = Builder.buildICmp(Pred: CmpInst::Predicate::ICMP_EQ, Res: CCVT, Op0: RHS, Op1: One);
5962 auto IsMinusOne =
5963 Builder.buildICmp(Pred: CmpInst::Predicate::ICMP_EQ, Res: CCVT, Op0: RHS, Op1: MinusOne);
5964 auto IsOneOrMinusOne = Builder.buildOr(Dst: CCVT, Src0: IsOne, Src1: IsMinusOne);
5965 AShr = Builder.buildSelect(Res: Ty, Tst: IsOneOrMinusOne, Op0: LHS, Op1: AShr);
5966
5967 // If divided by a positive value, we're done. Otherwise, the result must be
5968 // negated.
5969 auto Neg = Builder.buildNeg(Dst: Ty, Src0: AShr);
5970 auto IsNeg = Builder.buildICmp(Pred: CmpInst::Predicate::ICMP_SLT, Res: CCVT, Op0: RHS, Op1: Zero);
5971 Builder.buildSelect(Res: MI.getOperand(i: 0).getReg(), Tst: IsNeg, Op0: Neg, Op1: AShr);
5972 MI.eraseFromParent();
5973}
5974
5975void CombinerHelper::applyUDivByPow2(MachineInstr &MI) const {
5976 assert(MI.getOpcode() == TargetOpcode::G_UDIV && "Expected UDIV");
5977 auto &UDiv = cast<GenericMachineInstr>(Val&: MI);
5978 Register Dst = UDiv.getReg(Idx: 0);
5979 Register LHS = UDiv.getReg(Idx: 1);
5980 Register RHS = UDiv.getReg(Idx: 2);
5981 LLT Ty = MRI.getType(Reg: Dst);
5982 LLT ShiftAmtTy = getTargetLowering().getPreferredShiftAmountTy(ShiftValueTy: Ty);
5983
5984 auto C1 = Builder.buildCTTZ(Dst: ShiftAmtTy, Src0: RHS);
5985 Builder.buildLShr(Dst: MI.getOperand(i: 0).getReg(), Src0: LHS, Src1: C1);
5986 MI.eraseFromParent();
5987}
5988
5989void CombinerHelper::applySimplifySRemByPow2(MachineInstr &MI) const {
5990 assert(MI.getOpcode() == TargetOpcode::G_SREM && "Expected SREM");
5991 auto &SRem = cast<GBinOp>(Val&: MI);
5992 Register Dst = SRem.getReg(Idx: 0);
5993 Register LHS = SRem.getLHSReg();
5994 Register RHS = SRem.getRHSReg();
5995 LLT Ty = MRI.getType(Reg: Dst);
5996 LLT ShiftAmtTy = getTargetLowering().getPreferredShiftAmountTy(ShiftValueTy: Ty);
5997
5998 // Effectively we want to lower G_SREM %lhs, %rhs, where %rhs is +/- a power
5999 // of 2, to the following branch-free bias-and-mask version:
6000 //
6001 // %abs = G_ABS %rhs
6002 // %mask = G_SUB %abs, 1
6003 // %sign = G_ASHR %lhs, $(bitwidth - 1)
6004 // %bias = G_AND %sign, %mask
6005 // %biased = G_ADD %lhs, %bias
6006 // %masked = G_AND %biased, %mask
6007 // %res = G_SUB %masked, %bias
6008 //
6009 // The bias adds (|%rhs| - 1) for negative %lhs, correcting rounding towards
6010 // zero (instead of towards -inf that a plain mask would give). Constant
6011 // divisors collapse %mask to a single G_CONSTANT via the CSEMIRBuilder folds
6012 // for G_ABS and G_SUB.
6013
6014 unsigned BitWidth = Ty.getScalarSizeInBits();
6015 auto AbsRHS = Builder.buildAbs(Dst: Ty, Src: RHS);
6016 auto Mask = Builder.buildSub(Dst: Ty, Src0: AbsRHS, Src1: Builder.buildConstant(Res: Ty, Val: 1));
6017 auto BWMinusOne = Builder.buildConstant(Res: ShiftAmtTy, Val: BitWidth - 1);
6018 auto Sign = Builder.buildAShr(Dst: Ty, Src0: LHS, Src1: BWMinusOne);
6019 auto Bias = Builder.buildAnd(Dst: Ty, Src0: Sign, Src1: Mask);
6020 auto Biased = Builder.buildAdd(Dst: Ty, Src0: LHS, Src1: Bias);
6021 auto Masked = Builder.buildAnd(Dst: Ty, Src0: Biased, Src1: Mask);
6022 Builder.buildSub(Dst, Src0: Masked, Src1: Bias);
6023 MI.eraseFromParent();
6024}
6025
6026bool CombinerHelper::matchUMulHToLShr(MachineInstr &MI) const {
6027 assert(MI.getOpcode() == TargetOpcode::G_UMULH);
6028 Register RHS = MI.getOperand(i: 2).getReg();
6029 Register Dst = MI.getOperand(i: 0).getReg();
6030 LLT Ty = MRI.getType(Reg: Dst);
6031 LLT RHSTy = MRI.getType(Reg: RHS);
6032 LLT ShiftAmtTy = getTargetLowering().getPreferredShiftAmountTy(ShiftValueTy: Ty);
6033 auto MatchPow2ExceptOne = [&](const Constant *C) {
6034 if (auto *CI = dyn_cast<ConstantInt>(Val: C))
6035 return CI->getValue().isPowerOf2() && !CI->getValue().isOne();
6036 return false;
6037 };
6038 if (!matchUnaryPredicate(MRI, Reg: RHS, Match: MatchPow2ExceptOne, AllowUndefs: false))
6039 return false;
6040 // We need to check both G_LSHR and G_CTLZ because the combine uses G_CTLZ to
6041 // get log base 2, and it is not always legal for on a target.
6042 return isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_LSHR, {Ty, ShiftAmtTy}}) &&
6043 isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_CTLZ, {RHSTy, RHSTy}});
6044}
6045
6046void CombinerHelper::applyUMulHToLShr(MachineInstr &MI) const {
6047 Register LHS = MI.getOperand(i: 1).getReg();
6048 Register RHS = MI.getOperand(i: 2).getReg();
6049 Register Dst = MI.getOperand(i: 0).getReg();
6050 LLT Ty = MRI.getType(Reg: Dst);
6051 LLT ShiftAmtTy = getTargetLowering().getPreferredShiftAmountTy(ShiftValueTy: Ty);
6052 unsigned NumEltBits = Ty.getScalarSizeInBits();
6053
6054 auto LogBase2 = buildLogBase2(V: RHS, MIB&: Builder);
6055 auto ShiftAmt =
6056 Builder.buildSub(Dst: Ty, Src0: Builder.buildConstant(Res: Ty, Val: NumEltBits), Src1: LogBase2);
6057 auto Trunc = Builder.buildZExtOrTrunc(Res: ShiftAmtTy, Op: ShiftAmt);
6058 Builder.buildLShr(Dst, Src0: LHS, Src1: Trunc);
6059 MI.eraseFromParent();
6060}
6061
6062bool CombinerHelper::matchTruncSSatS(MachineInstr &MI,
6063 Register &MatchInfo) const {
6064 Register Dst = MI.getOperand(i: 0).getReg();
6065 Register Src = MI.getOperand(i: 1).getReg();
6066 LLT DstTy = MRI.getType(Reg: Dst);
6067 LLT SrcTy = MRI.getType(Reg: Src);
6068 unsigned NumDstBits = DstTy.getScalarSizeInBits();
6069 unsigned NumSrcBits = SrcTy.getScalarSizeInBits();
6070 assert(NumSrcBits > NumDstBits && "Unexpected types for truncate operation");
6071
6072 if (!LI || !isLegalOrHasFewerElements(
6073 Query: {TargetOpcode::G_TRUNC_SSAT_S, {DstTy, SrcTy}}))
6074 return false;
6075
6076 APInt SignedMax = APInt::getSignedMaxValue(numBits: NumDstBits).sext(width: NumSrcBits);
6077 APInt SignedMin = APInt::getSignedMinValue(numBits: NumDstBits).sext(width: NumSrcBits);
6078 if (mi_match(
6079 R: Src, MRI,
6080 P: m_GSMin(L: m_GSMax(L: m_Reg(R&: MatchInfo), R: m_SpecificICstOrSplat(RequestedValue: SignedMin)),
6081 R: m_SpecificICstOrSplat(RequestedValue: SignedMax))))
6082 return true;
6083 if (mi_match(
6084 R: Src, MRI,
6085 P: m_GSMax(L: m_GSMin(L: m_Reg(R&: MatchInfo), R: m_SpecificICstOrSplat(RequestedValue: SignedMax)),
6086 R: m_SpecificICstOrSplat(RequestedValue: SignedMin))))
6087 return true;
6088
6089 // CVP in the midend will often transform trunc(smin(smax(..)) into
6090 // trunc nsw(smin(..)) as the smax against INT_MIN never saturates.
6091 if (MI.getFlag(Flag: MachineInstr::MIFlag::NoSWrap) &&
6092 mi_match(R: Src, MRI,
6093 P: m_GSMin(L: m_Reg(R&: MatchInfo), R: m_SpecificICstOrSplat(RequestedValue: SignedMax))))
6094 return true;
6095
6096 return false;
6097}
6098
6099void CombinerHelper::applyTruncSSatS(MachineInstr &MI,
6100 Register &MatchInfo) const {
6101 Register Dst = MI.getOperand(i: 0).getReg();
6102 Builder.buildTruncSSatS(Res: Dst, Op: MatchInfo);
6103 MI.eraseFromParent();
6104}
6105
6106bool CombinerHelper::matchTruncSSatU(MachineInstr &MI,
6107 Register &MatchInfo) const {
6108 Register Dst = MI.getOperand(i: 0).getReg();
6109 Register Src = MI.getOperand(i: 1).getReg();
6110 LLT DstTy = MRI.getType(Reg: Dst);
6111 LLT SrcTy = MRI.getType(Reg: Src);
6112 unsigned NumDstBits = DstTy.getScalarSizeInBits();
6113 unsigned NumSrcBits = SrcTy.getScalarSizeInBits();
6114 assert(NumSrcBits > NumDstBits && "Unexpected types for truncate operation");
6115
6116 if (!LI || !isLegalOrHasFewerElements(
6117 Query: {TargetOpcode::G_TRUNC_SSAT_U, {DstTy, SrcTy}}))
6118 return false;
6119 APInt UnsignedMax = APInt::getMaxValue(numBits: NumDstBits).zext(width: NumSrcBits);
6120 return mi_match(R: Src, MRI,
6121 P: m_GSMin(L: m_GSMax(L: m_Reg(R&: MatchInfo), R: m_SpecificICstOrSplat(RequestedValue: 0)),
6122 R: m_SpecificICstOrSplat(RequestedValue: UnsignedMax))) ||
6123 mi_match(R: Src, MRI,
6124 P: m_GSMax(L: m_GSMin(L: m_Reg(R&: MatchInfo),
6125 R: m_SpecificICstOrSplat(RequestedValue: UnsignedMax)),
6126 R: m_SpecificICstOrSplat(RequestedValue: 0))) ||
6127 mi_match(R: Src, MRI,
6128 P: m_GUMin(L: m_GSMax(L: m_Reg(R&: MatchInfo), R: m_SpecificICstOrSplat(RequestedValue: 0)),
6129 R: m_SpecificICstOrSplat(RequestedValue: UnsignedMax)));
6130}
6131
6132void CombinerHelper::applyTruncSSatU(MachineInstr &MI,
6133 Register &MatchInfo) const {
6134 Register Dst = MI.getOperand(i: 0).getReg();
6135 Builder.buildTruncSSatU(Res: Dst, Op: MatchInfo);
6136 MI.eraseFromParent();
6137}
6138
6139bool CombinerHelper::matchTruncUSatU(MachineInstr &MI,
6140 MachineInstr &MinMI) const {
6141 Register Min = MinMI.getOperand(i: 2).getReg();
6142 Register Val = MinMI.getOperand(i: 1).getReg();
6143 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6144 LLT SrcTy = MRI.getType(Reg: Val);
6145 unsigned NumDstBits = DstTy.getScalarSizeInBits();
6146 unsigned NumSrcBits = SrcTy.getScalarSizeInBits();
6147 assert(NumSrcBits > NumDstBits && "Unexpected types for truncate operation");
6148
6149 if (!LI || !isLegalOrHasFewerElements(
6150 Query: {TargetOpcode::G_TRUNC_SSAT_U, {DstTy, SrcTy}}))
6151 return false;
6152 APInt UnsignedMax = APInt::getMaxValue(numBits: NumDstBits).zext(width: NumSrcBits);
6153 return mi_match(R: Min, MRI, P: m_SpecificICstOrSplat(RequestedValue: UnsignedMax)) &&
6154 !mi_match(R: Val, MRI, P: m_GSMax(L: m_Reg(), R: m_Reg()));
6155}
6156
6157bool CombinerHelper::matchTruncUSatUToFPTOUISat(MachineInstr &MI,
6158 MachineInstr &SrcMI) const {
6159 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6160 LLT SrcTy = MRI.getType(Reg: SrcMI.getOperand(i: 1).getReg());
6161
6162 return LI &&
6163 isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_FPTOUI_SAT, {DstTy, SrcTy}});
6164}
6165
6166bool CombinerHelper::matchRedundantNegOperands(MachineInstr &MI,
6167 BuildFnTy &MatchInfo) const {
6168 unsigned Opc = MI.getOpcode();
6169 assert(Opc == TargetOpcode::G_FADD || Opc == TargetOpcode::G_FSUB);
6170
6171 Register Dst = MI.getOperand(i: 0).getReg();
6172 Register X = MI.getOperand(i: 1).getReg();
6173 Register Y = MI.getOperand(i: 2).getReg();
6174 LLT Type = MRI.getType(Reg: Dst);
6175
6176 // fold (fadd x, fneg(y)) -> (fsub x, y)
6177 // fold (fadd fneg(y), x) -> (fsub x, y)
6178 // G_ADD is commutative so both cases are checked by m_GFAdd
6179 if (mi_match(R: Dst, MRI, P: m_GFAdd(L: m_Reg(R&: X), R: m_GFNeg(Src: m_Reg(R&: Y)))) &&
6180 isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_FSUB, {Type}})) {
6181 Opc = TargetOpcode::G_FSUB;
6182 }
6183 /// fold (fsub x, fneg(y)) -> (fadd x, y)
6184 else if (mi_match(R: Dst, MRI, P: m_GFSub(L: m_Reg(R&: X), R: m_GFNeg(Src: m_Reg(R&: Y)))) &&
6185 isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_FADD, {Type}})) {
6186 Opc = TargetOpcode::G_FADD;
6187 } else
6188 return false;
6189
6190 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6191 Observer.changingInstr(MI);
6192 MI.setDesc(B.getTII().get(Opcode: Opc));
6193 MI.getOperand(i: 1).setReg(X);
6194 MI.getOperand(i: 2).setReg(Y);
6195 Observer.changedInstr(MI);
6196 };
6197 return true;
6198}
6199
6200bool CombinerHelper::matchFsubToFneg(MachineInstr &MI,
6201 Register &MatchInfo) const {
6202 assert(MI.getOpcode() == TargetOpcode::G_FSUB);
6203
6204 Register LHS = MI.getOperand(i: 1).getReg();
6205 MatchInfo = MI.getOperand(i: 2).getReg();
6206 LLT Ty = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6207
6208 const auto LHSCst = Ty.isVector()
6209 ? getFConstantSplat(VReg: LHS, MRI, /* allowUndef */ AllowUndef: true)
6210 : getFConstantVRegValWithLookThrough(VReg: LHS, MRI);
6211 if (!LHSCst)
6212 return false;
6213
6214 // -0.0 is always allowed
6215 if (LHSCst->Value.isNegZero())
6216 return true;
6217
6218 // +0.0 is only allowed if nsz is set.
6219 if (LHSCst->Value.isPosZero())
6220 return MI.getFlag(Flag: MachineInstr::FmNsz);
6221
6222 return false;
6223}
6224
6225void CombinerHelper::applyFsubToFneg(MachineInstr &MI,
6226 Register &MatchInfo) const {
6227 Register Dst = MI.getOperand(i: 0).getReg();
6228 Builder.buildFNeg(
6229 Dst, Src0: Builder.buildFCanonicalize(Dst: MRI.getType(Reg: Dst), Src0: MatchInfo).getReg(Idx: 0));
6230 eraseInst(MI);
6231}
6232
6233/// Checks if \p MI is TargetOpcode::G_FMUL and contractable either
6234/// due to global flags or MachineInstr flags.
6235static bool isContractableFMul(MachineInstr &MI, bool AllowFusionGlobally) {
6236 if (MI.getOpcode() != TargetOpcode::G_FMUL)
6237 return false;
6238 return AllowFusionGlobally || MI.getFlag(Flag: MachineInstr::MIFlag::FmContract);
6239}
6240
6241static bool hasMoreUses(const MachineInstr &MI0, const MachineInstr &MI1,
6242 const MachineRegisterInfo &MRI) {
6243 return std::distance(first: MRI.use_instr_nodbg_begin(RegNo: MI0.getOperand(i: 0).getReg()),
6244 last: MRI.use_instr_nodbg_end()) >
6245 std::distance(first: MRI.use_instr_nodbg_begin(RegNo: MI1.getOperand(i: 0).getReg()),
6246 last: MRI.use_instr_nodbg_end());
6247}
6248
6249bool CombinerHelper::canCombineFMadOrFMA(MachineInstr &MI,
6250 bool &AllowFusionGlobally,
6251 bool &HasFMAD, bool &Aggressive,
6252 bool CanReassociate) const {
6253
6254 auto *MF = MI.getMF();
6255 const auto &TLI = *MF->getSubtarget().getTargetLowering();
6256 LLT DstType = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6257
6258 if (CanReassociate && !MI.getFlag(Flag: MachineInstr::MIFlag::FmReassoc))
6259 return false;
6260
6261 // Floating-point multiply-add with intermediate rounding.
6262 HasFMAD = (!isPreLegalize() && TLI.isFMADLegal(MI, Ty: DstType));
6263 // Floating-point multiply-add without intermediate rounding.
6264 bool HasFMA = TLI.isFMAFasterThanFMulAndFAdd(MF: *MF, DstType) &&
6265 isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_FMA, {DstType}});
6266 // No valid opcode, do not combine.
6267 if (!HasFMAD && !HasFMA)
6268 return false;
6269
6270 // FMAD (with intermediate rounding) is always safe to form; FMA requires the
6271 // contract fast-math flag.
6272 AllowFusionGlobally = HasFMAD;
6273 // If the addition is not contractable, do not combine.
6274 if (!AllowFusionGlobally && !MI.getFlag(Flag: MachineInstr::MIFlag::FmContract))
6275 return false;
6276
6277 Aggressive = TLI.enableAggressiveFMAFusion(Ty: DstType);
6278 return true;
6279}
6280
6281bool CombinerHelper::matchCombineFAddFMulToFMadOrFMA(
6282 MachineInstr &MI,
6283 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
6284 assert(MI.getOpcode() == TargetOpcode::G_FADD);
6285
6286 bool AllowFusionGlobally, HasFMAD, Aggressive;
6287 if (!canCombineFMadOrFMA(MI, AllowFusionGlobally, HasFMAD, Aggressive))
6288 return false;
6289
6290 Register Op1 = MI.getOperand(i: 1).getReg();
6291 Register Op2 = MI.getOperand(i: 2).getReg();
6292 MachineInstr *Op1Def, *Op2Def;
6293 if (!mi_match(R: Op1, MRI, P: m_MInstr(MI&: Op1Def)) ||
6294 !mi_match(R: Op2, MRI, P: m_MInstr(MI&: Op2Def)))
6295 return false;
6296 DefinitionAndSourceRegister LHS = {.MI: Op1Def, .Reg: Op1};
6297 DefinitionAndSourceRegister RHS = {.MI: Op2Def, .Reg: Op2};
6298 unsigned PreferredFusedOpcode =
6299 HasFMAD ? TargetOpcode::G_FMAD : TargetOpcode::G_FMA;
6300
6301 // If we have two choices trying to fold (fadd (fmul u, v), (fmul x, y)),
6302 // prefer to fold the multiply with fewer uses.
6303 if (Aggressive && isContractableFMul(MI&: *LHS.MI, AllowFusionGlobally) &&
6304 isContractableFMul(MI&: *RHS.MI, AllowFusionGlobally)) {
6305 if (hasMoreUses(MI0: *LHS.MI, MI1: *RHS.MI, MRI))
6306 std::swap(a&: LHS, b&: RHS);
6307 }
6308
6309 // fold (fadd (fmul x, y), z) -> (fma x, y, z)
6310 if (isContractableFMul(MI&: *LHS.MI, AllowFusionGlobally) &&
6311 (Aggressive || MRI.hasOneNonDBGUse(RegNo: LHS.Reg))) {
6312 unsigned Flags = MI.getFlags() & LHS.MI->getFlags();
6313 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6314 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {MI.getOperand(i: 0).getReg()},
6315 SrcOps: {LHS.MI->getOperand(i: 1).getReg(),
6316 LHS.MI->getOperand(i: 2).getReg(), RHS.Reg},
6317 Flags);
6318 };
6319 return true;
6320 }
6321
6322 // fold (fadd x, (fmul y, z)) -> (fma y, z, x)
6323 if (isContractableFMul(MI&: *RHS.MI, AllowFusionGlobally) &&
6324 (Aggressive || MRI.hasOneNonDBGUse(RegNo: RHS.Reg))) {
6325 unsigned Flags = MI.getFlags() & RHS.MI->getFlags();
6326 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6327 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {MI.getOperand(i: 0).getReg()},
6328 SrcOps: {RHS.MI->getOperand(i: 1).getReg(),
6329 RHS.MI->getOperand(i: 2).getReg(), LHS.Reg},
6330 Flags);
6331 };
6332 return true;
6333 }
6334
6335 return false;
6336}
6337
6338bool CombinerHelper::matchCombineFAddFpExtFMulToFMadOrFMA(
6339 MachineInstr &MI,
6340 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
6341 assert(MI.getOpcode() == TargetOpcode::G_FADD);
6342
6343 bool AllowFusionGlobally, HasFMAD, Aggressive;
6344 if (!canCombineFMadOrFMA(MI, AllowFusionGlobally, HasFMAD, Aggressive))
6345 return false;
6346
6347 const auto &TLI = *MI.getMF()->getSubtarget().getTargetLowering();
6348 Register Op1 = MI.getOperand(i: 1).getReg();
6349 Register Op2 = MI.getOperand(i: 2).getReg();
6350 MachineInstr *Op1Def, *Op2Def;
6351 if (!mi_match(R: Op1, MRI, P: m_MInstr(MI&: Op1Def)) ||
6352 !mi_match(R: Op2, MRI, P: m_MInstr(MI&: Op2Def)))
6353 return false;
6354 DefinitionAndSourceRegister LHS = {.MI: Op1Def, .Reg: Op1};
6355 DefinitionAndSourceRegister RHS = {.MI: Op2Def, .Reg: Op2};
6356 LLT DstType = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6357
6358 unsigned PreferredFusedOpcode =
6359 HasFMAD ? TargetOpcode::G_FMAD : TargetOpcode::G_FMA;
6360
6361 MachineInstr *LHSFpExtSrc;
6362 bool LHSContractable =
6363 mi_match(R: LHS.Reg, MRI, P: m_GFPExt(Src: m_MInstr(MI&: LHSFpExtSrc))) &&
6364 isContractableFMul(MI&: *LHSFpExtSrc, AllowFusionGlobally) &&
6365 TLI.isFPExtFoldable(MI, Opcode: PreferredFusedOpcode, DestTy: DstType,
6366 SrcTy: MRI.getType(Reg: LHSFpExtSrc->getOperand(i: 1).getReg()));
6367 MachineInstr *RHSFpExtSrc;
6368 bool RHSContractable =
6369 mi_match(R: RHS.Reg, MRI, P: m_GFPExt(Src: m_MInstr(MI&: RHSFpExtSrc))) &&
6370 isContractableFMul(MI&: *RHSFpExtSrc, AllowFusionGlobally) &&
6371 TLI.isFPExtFoldable(MI, Opcode: PreferredFusedOpcode, DestTy: DstType,
6372 SrcTy: MRI.getType(Reg: RHSFpExtSrc->getOperand(i: 1).getReg()));
6373
6374 // fold (fadd (fpext (fmul x, y)), z) -> (fma (fpext x), (fpext y), z)
6375 if (LHSContractable || RHSContractable) {
6376 // Ensure that the contractable fmul with the fewest uses (if both are
6377 // contractable) is the LHS operand.
6378 if (!LHSContractable ||
6379 (RHSContractable && hasMoreUses(MI0: *LHSFpExtSrc, MI1: *RHSFpExtSrc, MRI))) {
6380 std::swap(a&: LHS, b&: RHS);
6381 LHSFpExtSrc = RHSFpExtSrc;
6382 }
6383
6384 unsigned Flags = MI.getFlags() & LHSFpExtSrc->getFlags();
6385 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6386 auto FpExtX = B.buildFPExt(Res: DstType, Op: LHSFpExtSrc->getOperand(i: 1).getReg());
6387 auto FpExtY = B.buildFPExt(Res: DstType, Op: LHSFpExtSrc->getOperand(i: 2).getReg());
6388 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {MI.getOperand(i: 0).getReg()},
6389 SrcOps: {FpExtX.getReg(Idx: 0), FpExtY.getReg(Idx: 0), RHS.Reg}, Flags);
6390 };
6391 return true;
6392 }
6393
6394 return false;
6395}
6396
6397bool CombinerHelper::matchCombineFAddFMAFMulToFMadOrFMA(
6398 MachineInstr &MI,
6399 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
6400 assert(MI.getOpcode() == TargetOpcode::G_FADD);
6401
6402 bool AllowFusionGlobally, HasFMAD, Aggressive;
6403 if (!canCombineFMadOrFMA(MI, AllowFusionGlobally, HasFMAD, Aggressive, CanReassociate: true))
6404 return false;
6405
6406 Register Op1 = MI.getOperand(i: 1).getReg();
6407 Register Op2 = MI.getOperand(i: 2).getReg();
6408 MachineInstr *Op1Def, *Op2Def;
6409 if (!mi_match(R: Op1, MRI, P: m_MInstr(MI&: Op1Def)) ||
6410 !mi_match(R: Op2, MRI, P: m_MInstr(MI&: Op2Def)))
6411 return false;
6412 DefinitionAndSourceRegister LHS = {.MI: Op1Def, .Reg: Op1};
6413 DefinitionAndSourceRegister RHS = {.MI: Op2Def, .Reg: Op2};
6414 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6415
6416 unsigned PreferredFusedOpcode =
6417 HasFMAD ? TargetOpcode::G_FMAD : TargetOpcode::G_FMA;
6418
6419 MachineInstr *FMA = nullptr;
6420 Register Z;
6421 // fold (fadd (fma x, y, (fmul u, v)), z) -> (fma x, y, (fma u, v, z))
6422 if (LHS.MI->getOpcode() == PreferredFusedOpcode &&
6423 mi_match(R: LHS.MI->getOperand(i: 3).getReg(), MRI,
6424 P: m_GFMul(L: m_Reg(), R: m_Reg())) &&
6425 MRI.hasOneNonDBGUse(RegNo: LHS.MI->getOperand(i: 0).getReg()) &&
6426 MRI.hasOneNonDBGUse(RegNo: LHS.MI->getOperand(i: 3).getReg())) {
6427 FMA = LHS.MI;
6428 Z = RHS.Reg;
6429 }
6430 // fold (fadd z, (fma x, y, (fmul u, v))) -> (fma x, y, (fma u, v, z))
6431 else if (RHS.MI->getOpcode() == PreferredFusedOpcode &&
6432 mi_match(R: RHS.MI->getOperand(i: 3).getReg(), MRI,
6433 P: m_GFMul(L: m_Reg(), R: m_Reg())) &&
6434 MRI.hasOneNonDBGUse(RegNo: RHS.MI->getOperand(i: 0).getReg()) &&
6435 MRI.hasOneNonDBGUse(RegNo: RHS.MI->getOperand(i: 3).getReg())) {
6436 Z = LHS.Reg;
6437 FMA = RHS.MI;
6438 }
6439
6440 if (FMA) {
6441 MachineInstr *FMulMI;
6442 if (!mi_match(R: FMA->getOperand(i: 3).getReg(), MRI, P: m_MInstr(MI&: FMulMI)))
6443 return false;
6444 Register X = FMA->getOperand(i: 1).getReg();
6445 Register Y = FMA->getOperand(i: 2).getReg();
6446 Register U = FMulMI->getOperand(i: 1).getReg();
6447 Register V = FMulMI->getOperand(i: 2).getReg();
6448 unsigned InnerFlags = MI.getFlags() & FMulMI->getFlags();
6449 unsigned OuterFlags = MI.getFlags() & FMA->getFlags();
6450
6451 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6452 Register InnerFMA = MRI.createGenericVirtualRegister(Ty: DstTy);
6453 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {InnerFMA}, SrcOps: {U, V, Z}, Flags: InnerFlags);
6454 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {MI.getOperand(i: 0).getReg()},
6455 SrcOps: {X, Y, InnerFMA}, Flags: OuterFlags);
6456 };
6457 return true;
6458 }
6459
6460 return false;
6461}
6462
6463bool CombinerHelper::matchCombineFAddFpExtFMulToFMadOrFMAAggressive(
6464 MachineInstr &MI,
6465 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
6466 assert(MI.getOpcode() == TargetOpcode::G_FADD);
6467
6468 bool AllowFusionGlobally, HasFMAD, Aggressive;
6469 if (!canCombineFMadOrFMA(MI, AllowFusionGlobally, HasFMAD, Aggressive))
6470 return false;
6471
6472 if (!Aggressive)
6473 return false;
6474
6475 const auto &TLI = *MI.getMF()->getSubtarget().getTargetLowering();
6476 LLT DstType = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6477 Register Op1 = MI.getOperand(i: 1).getReg();
6478 Register Op2 = MI.getOperand(i: 2).getReg();
6479 MachineInstr *Op1Def, *Op2Def;
6480 if (!mi_match(R: Op1, MRI, P: m_MInstr(MI&: Op1Def)) ||
6481 !mi_match(R: Op2, MRI, P: m_MInstr(MI&: Op2Def)))
6482 return false;
6483 DefinitionAndSourceRegister LHS = {.MI: Op1Def, .Reg: Op1};
6484 DefinitionAndSourceRegister RHS = {.MI: Op2Def, .Reg: Op2};
6485
6486 unsigned PreferredFusedOpcode =
6487 HasFMAD ? TargetOpcode::G_FMAD : TargetOpcode::G_FMA;
6488
6489 // If we have two choices trying to fold (fadd (fmul u, v), (fmul x, y)),
6490 // prefer to fold the multiply with fewer uses.
6491 if (Aggressive && isContractableFMul(MI&: *LHS.MI, AllowFusionGlobally) &&
6492 isContractableFMul(MI&: *RHS.MI, AllowFusionGlobally)) {
6493 if (hasMoreUses(MI0: *LHS.MI, MI1: *RHS.MI, MRI))
6494 std::swap(a&: LHS, b&: RHS);
6495 }
6496
6497 // Builds: (fma x, y, (fma (fpext u), (fpext v), z))
6498 auto buildMatchInfo = [=, &MI](Register U, Register V, Register Z, Register X,
6499 Register Y, unsigned InnerFlags,
6500 unsigned OuterFlags, MachineIRBuilder &B) {
6501 Register FpExtU = B.buildFPExt(Res: DstType, Op: U).getReg(Idx: 0);
6502 Register FpExtV = B.buildFPExt(Res: DstType, Op: V).getReg(Idx: 0);
6503 Register InnerFMA = B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {DstType},
6504 SrcOps: {FpExtU, FpExtV, Z}, Flags: InnerFlags)
6505 .getReg(Idx: 0);
6506 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {MI.getOperand(i: 0).getReg()},
6507 SrcOps: {X, Y, InnerFMA}, Flags: OuterFlags);
6508 };
6509
6510 MachineInstr *FMulMI, *FMAMI;
6511 // fold (fadd (fma x, y, (fpext (fmul u, v))), z)
6512 // -> (fma x, y, (fma (fpext u), (fpext v), z))
6513 if (LHS.MI->getOpcode() == PreferredFusedOpcode &&
6514 mi_match(R: LHS.MI->getOperand(i: 3).getReg(), MRI,
6515 P: m_GFPExt(Src: m_MInstr(MI&: FMulMI))) &&
6516 isContractableFMul(MI&: *FMulMI, AllowFusionGlobally) &&
6517 TLI.isFPExtFoldable(MI, Opcode: PreferredFusedOpcode, DestTy: DstType,
6518 SrcTy: MRI.getType(Reg: FMulMI->getOperand(i: 0).getReg()))) {
6519 unsigned InnerFlags = MI.getFlags() & FMulMI->getFlags();
6520 unsigned OuterFlags = MI.getFlags() & LHS.MI->getFlags();
6521 MatchInfo = [=](MachineIRBuilder &B) {
6522 buildMatchInfo(FMulMI->getOperand(i: 1).getReg(),
6523 FMulMI->getOperand(i: 2).getReg(), RHS.Reg,
6524 LHS.MI->getOperand(i: 1).getReg(),
6525 LHS.MI->getOperand(i: 2).getReg(), InnerFlags, OuterFlags, B);
6526 };
6527 return true;
6528 }
6529
6530 // fold (fadd (fpext (fma x, y, (fmul u, v))), z)
6531 // -> (fma (fpext x), (fpext y), (fma (fpext u), (fpext v), z))
6532 // FIXME: This turns two single-precision and one double-precision
6533 // operation into two double-precision operations, which might not be
6534 // interesting for all targets, especially GPUs.
6535 if (mi_match(R: LHS.Reg, MRI, P: m_GFPExt(Src: m_MInstr(MI&: FMAMI))) &&
6536 FMAMI->getOpcode() == PreferredFusedOpcode) {
6537 MachineInstr *FMulMI;
6538 if (!mi_match(R: FMAMI->getOperand(i: 3).getReg(), MRI, P: m_MInstr(MI&: FMulMI)))
6539 return false;
6540 if (isContractableFMul(MI&: *FMulMI, AllowFusionGlobally) &&
6541 TLI.isFPExtFoldable(MI, Opcode: PreferredFusedOpcode, DestTy: DstType,
6542 SrcTy: MRI.getType(Reg: FMAMI->getOperand(i: 0).getReg()))) {
6543 unsigned InnerFlags = MI.getFlags() & FMulMI->getFlags();
6544 unsigned OuterFlags = MI.getFlags() & FMAMI->getFlags();
6545 MatchInfo = [=](MachineIRBuilder &B) {
6546 Register X = FMAMI->getOperand(i: 1).getReg();
6547 Register Y = FMAMI->getOperand(i: 2).getReg();
6548 X = B.buildFPExt(Res: DstType, Op: X).getReg(Idx: 0);
6549 Y = B.buildFPExt(Res: DstType, Op: Y).getReg(Idx: 0);
6550 buildMatchInfo(FMulMI->getOperand(i: 1).getReg(),
6551 FMulMI->getOperand(i: 2).getReg(), RHS.Reg, X, Y,
6552 InnerFlags, OuterFlags, B);
6553 };
6554
6555 return true;
6556 }
6557 }
6558
6559 // fold (fadd z, (fma x, y, (fpext (fmul u, v)))
6560 // -> (fma x, y, (fma (fpext u), (fpext v), z))
6561 if (RHS.MI->getOpcode() == PreferredFusedOpcode &&
6562 mi_match(R: RHS.MI->getOperand(i: 3).getReg(), MRI,
6563 P: m_GFPExt(Src: m_MInstr(MI&: FMulMI))) &&
6564 isContractableFMul(MI&: *FMulMI, AllowFusionGlobally) &&
6565 TLI.isFPExtFoldable(MI, Opcode: PreferredFusedOpcode, DestTy: DstType,
6566 SrcTy: MRI.getType(Reg: FMulMI->getOperand(i: 0).getReg()))) {
6567 unsigned InnerFlags = MI.getFlags() & FMulMI->getFlags();
6568 unsigned OuterFlags = MI.getFlags() & RHS.MI->getFlags();
6569 MatchInfo = [=](MachineIRBuilder &B) {
6570 buildMatchInfo(FMulMI->getOperand(i: 1).getReg(),
6571 FMulMI->getOperand(i: 2).getReg(), LHS.Reg,
6572 RHS.MI->getOperand(i: 1).getReg(),
6573 RHS.MI->getOperand(i: 2).getReg(), InnerFlags, OuterFlags, B);
6574 };
6575 return true;
6576 }
6577
6578 // fold (fadd z, (fpext (fma x, y, (fmul u, v)))
6579 // -> (fma (fpext x), (fpext y), (fma (fpext u), (fpext v), z))
6580 // FIXME: This turns two single-precision and one double-precision
6581 // operation into two double-precision operations, which might not be
6582 // interesting for all targets, especially GPUs.
6583 if (mi_match(R: RHS.Reg, MRI, P: m_GFPExt(Src: m_MInstr(MI&: FMAMI))) &&
6584 FMAMI->getOpcode() == PreferredFusedOpcode) {
6585 MachineInstr *FMulMI;
6586 if (!mi_match(R: FMAMI->getOperand(i: 3).getReg(), MRI, P: m_MInstr(MI&: FMulMI)))
6587 return false;
6588 if (isContractableFMul(MI&: *FMulMI, AllowFusionGlobally) &&
6589 TLI.isFPExtFoldable(MI, Opcode: PreferredFusedOpcode, DestTy: DstType,
6590 SrcTy: MRI.getType(Reg: FMAMI->getOperand(i: 0).getReg()))) {
6591 unsigned InnerFlags = MI.getFlags() & FMulMI->getFlags();
6592 unsigned OuterFlags = MI.getFlags() & FMAMI->getFlags();
6593 MatchInfo = [=](MachineIRBuilder &B) {
6594 Register X = FMAMI->getOperand(i: 1).getReg();
6595 Register Y = FMAMI->getOperand(i: 2).getReg();
6596 X = B.buildFPExt(Res: DstType, Op: X).getReg(Idx: 0);
6597 Y = B.buildFPExt(Res: DstType, Op: Y).getReg(Idx: 0);
6598 buildMatchInfo(FMulMI->getOperand(i: 1).getReg(),
6599 FMulMI->getOperand(i: 2).getReg(), LHS.Reg, X, Y,
6600 InnerFlags, OuterFlags, B);
6601 };
6602 return true;
6603 }
6604 }
6605
6606 return false;
6607}
6608
6609bool CombinerHelper::matchCombineFSubFMulToFMadOrFMA(
6610 MachineInstr &MI,
6611 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
6612 assert(MI.getOpcode() == TargetOpcode::G_FSUB);
6613
6614 bool AllowFusionGlobally, HasFMAD, Aggressive;
6615 if (!canCombineFMadOrFMA(MI, AllowFusionGlobally, HasFMAD, Aggressive))
6616 return false;
6617
6618 Register Op1 = MI.getOperand(i: 1).getReg();
6619 Register Op2 = MI.getOperand(i: 2).getReg();
6620 MachineInstr *Op1Def, *Op2Def;
6621 if (!mi_match(R: Op1, MRI, P: m_MInstr(MI&: Op1Def)) ||
6622 !mi_match(R: Op2, MRI, P: m_MInstr(MI&: Op2Def)))
6623 return false;
6624 DefinitionAndSourceRegister LHS = {.MI: Op1Def, .Reg: Op1};
6625 DefinitionAndSourceRegister RHS = {.MI: Op2Def, .Reg: Op2};
6626 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6627
6628 // If we have two choices trying to fold (fsub (fmul u, v), (fmul x, y)),
6629 // prefer to fold the multiply with fewer uses.
6630 int FirstMulHasFewerUses = true;
6631 if (isContractableFMul(MI&: *LHS.MI, AllowFusionGlobally) &&
6632 isContractableFMul(MI&: *RHS.MI, AllowFusionGlobally) &&
6633 hasMoreUses(MI0: *LHS.MI, MI1: *RHS.MI, MRI))
6634 FirstMulHasFewerUses = false;
6635
6636 unsigned PreferredFusedOpcode =
6637 HasFMAD ? TargetOpcode::G_FMAD : TargetOpcode::G_FMA;
6638
6639 // fold (fsub (fmul x, y), z) -> (fma x, y, -z)
6640 if (FirstMulHasFewerUses &&
6641 (isContractableFMul(MI&: *LHS.MI, AllowFusionGlobally) &&
6642 (Aggressive || MRI.hasOneNonDBGUse(RegNo: LHS.Reg)))) {
6643 unsigned Flags = MI.getFlags() & LHS.MI->getFlags();
6644 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6645 Register NegZ = B.buildFNeg(Dst: DstTy, Src0: RHS.Reg).getReg(Idx: 0);
6646 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {MI.getOperand(i: 0).getReg()},
6647 SrcOps: {LHS.MI->getOperand(i: 1).getReg(),
6648 LHS.MI->getOperand(i: 2).getReg(), NegZ},
6649 Flags);
6650 };
6651 return true;
6652 }
6653 // fold (fsub x, (fmul y, z)) -> (fma -y, z, x)
6654 else if ((isContractableFMul(MI&: *RHS.MI, AllowFusionGlobally) &&
6655 (Aggressive || MRI.hasOneNonDBGUse(RegNo: RHS.Reg)))) {
6656 unsigned Flags = MI.getFlags() & RHS.MI->getFlags();
6657 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6658 Register NegY =
6659 B.buildFNeg(Dst: DstTy, Src0: RHS.MI->getOperand(i: 1).getReg()).getReg(Idx: 0);
6660 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {MI.getOperand(i: 0).getReg()},
6661 SrcOps: {NegY, RHS.MI->getOperand(i: 2).getReg(), LHS.Reg}, Flags);
6662 };
6663 return true;
6664 }
6665
6666 return false;
6667}
6668
6669bool CombinerHelper::matchCombineFSubFNegFMulToFMadOrFMA(
6670 MachineInstr &MI,
6671 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
6672 assert(MI.getOpcode() == TargetOpcode::G_FSUB);
6673
6674 bool AllowFusionGlobally, HasFMAD, Aggressive;
6675 if (!canCombineFMadOrFMA(MI, AllowFusionGlobally, HasFMAD, Aggressive))
6676 return false;
6677
6678 Register LHSReg = MI.getOperand(i: 1).getReg();
6679 Register RHSReg = MI.getOperand(i: 2).getReg();
6680 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6681
6682 unsigned PreferredFusedOpcode =
6683 HasFMAD ? TargetOpcode::G_FMAD : TargetOpcode::G_FMA;
6684
6685 MachineInstr *FMulMI;
6686 // fold (fsub (fneg (fmul x, y)), z) -> (fma (fneg x), y, (fneg z))
6687 if (mi_match(R: LHSReg, MRI, P: m_GFNeg(Src: m_MInstr(MI&: FMulMI))) &&
6688 (Aggressive || (MRI.hasOneNonDBGUse(RegNo: LHSReg) &&
6689 MRI.hasOneNonDBGUse(RegNo: FMulMI->getOperand(i: 0).getReg()))) &&
6690 isContractableFMul(MI&: *FMulMI, AllowFusionGlobally)) {
6691 unsigned Flags = MI.getFlags() & FMulMI->getFlags();
6692 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6693 Register NegX =
6694 B.buildFNeg(Dst: DstTy, Src0: FMulMI->getOperand(i: 1).getReg()).getReg(Idx: 0);
6695 Register NegZ = B.buildFNeg(Dst: DstTy, Src0: RHSReg).getReg(Idx: 0);
6696 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {MI.getOperand(i: 0).getReg()},
6697 SrcOps: {NegX, FMulMI->getOperand(i: 2).getReg(), NegZ}, Flags);
6698 };
6699 return true;
6700 }
6701
6702 // fold (fsub x, (fneg (fmul, y, z))) -> (fma y, z, x)
6703 if (mi_match(R: RHSReg, MRI, P: m_GFNeg(Src: m_MInstr(MI&: FMulMI))) &&
6704 (Aggressive || (MRI.hasOneNonDBGUse(RegNo: RHSReg) &&
6705 MRI.hasOneNonDBGUse(RegNo: FMulMI->getOperand(i: 0).getReg()))) &&
6706 isContractableFMul(MI&: *FMulMI, AllowFusionGlobally)) {
6707 unsigned Flags = MI.getFlags() & FMulMI->getFlags();
6708 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6709 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {MI.getOperand(i: 0).getReg()},
6710 SrcOps: {FMulMI->getOperand(i: 1).getReg(),
6711 FMulMI->getOperand(i: 2).getReg(), LHSReg},
6712 Flags);
6713 };
6714 return true;
6715 }
6716
6717 return false;
6718}
6719
6720bool CombinerHelper::matchCombineFSubFpExtFMulToFMadOrFMA(
6721 MachineInstr &MI,
6722 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
6723 assert(MI.getOpcode() == TargetOpcode::G_FSUB);
6724
6725 bool AllowFusionGlobally, HasFMAD, Aggressive;
6726 if (!canCombineFMadOrFMA(MI, AllowFusionGlobally, HasFMAD, Aggressive))
6727 return false;
6728
6729 Register LHSReg = MI.getOperand(i: 1).getReg();
6730 Register RHSReg = MI.getOperand(i: 2).getReg();
6731 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6732
6733 unsigned PreferredFusedOpcode =
6734 HasFMAD ? TargetOpcode::G_FMAD : TargetOpcode::G_FMA;
6735
6736 MachineInstr *FMulMI;
6737 // fold (fsub (fpext (fmul x, y)), z) -> (fma (fpext x), (fpext y), (fneg z))
6738 if (mi_match(R: LHSReg, MRI, P: m_GFPExt(Src: m_MInstr(MI&: FMulMI))) &&
6739 isContractableFMul(MI&: *FMulMI, AllowFusionGlobally) &&
6740 (Aggressive || MRI.hasOneNonDBGUse(RegNo: LHSReg))) {
6741 unsigned Flags = MI.getFlags() & FMulMI->getFlags();
6742 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6743 Register FpExtX =
6744 B.buildFPExt(Res: DstTy, Op: FMulMI->getOperand(i: 1).getReg()).getReg(Idx: 0);
6745 Register FpExtY =
6746 B.buildFPExt(Res: DstTy, Op: FMulMI->getOperand(i: 2).getReg()).getReg(Idx: 0);
6747 Register NegZ = B.buildFNeg(Dst: DstTy, Src0: RHSReg).getReg(Idx: 0);
6748 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {MI.getOperand(i: 0).getReg()},
6749 SrcOps: {FpExtX, FpExtY, NegZ}, Flags);
6750 };
6751 return true;
6752 }
6753
6754 // fold (fsub x, (fpext (fmul y, z))) -> (fma (fneg (fpext y)), (fpext z), x)
6755 if (mi_match(R: RHSReg, MRI, P: m_GFPExt(Src: m_MInstr(MI&: FMulMI))) &&
6756 isContractableFMul(MI&: *FMulMI, AllowFusionGlobally) &&
6757 (Aggressive || MRI.hasOneNonDBGUse(RegNo: RHSReg))) {
6758 unsigned Flags = MI.getFlags() & FMulMI->getFlags();
6759 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6760 Register FpExtY =
6761 B.buildFPExt(Res: DstTy, Op: FMulMI->getOperand(i: 1).getReg()).getReg(Idx: 0);
6762 Register NegY = B.buildFNeg(Dst: DstTy, Src0: FpExtY).getReg(Idx: 0);
6763 Register FpExtZ =
6764 B.buildFPExt(Res: DstTy, Op: FMulMI->getOperand(i: 2).getReg()).getReg(Idx: 0);
6765 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {MI.getOperand(i: 0).getReg()},
6766 SrcOps: {NegY, FpExtZ, LHSReg}, Flags);
6767 };
6768 return true;
6769 }
6770
6771 return false;
6772}
6773
6774bool CombinerHelper::matchCombineFSubFpExtFNegFMulToFMadOrFMA(
6775 MachineInstr &MI,
6776 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
6777 assert(MI.getOpcode() == TargetOpcode::G_FSUB);
6778
6779 bool AllowFusionGlobally, HasFMAD, Aggressive;
6780 if (!canCombineFMadOrFMA(MI, AllowFusionGlobally, HasFMAD, Aggressive))
6781 return false;
6782
6783 const auto &TLI = *MI.getMF()->getSubtarget().getTargetLowering();
6784 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6785 Register LHSReg = MI.getOperand(i: 1).getReg();
6786 Register RHSReg = MI.getOperand(i: 2).getReg();
6787
6788 unsigned PreferredFusedOpcode =
6789 HasFMAD ? TargetOpcode::G_FMAD : TargetOpcode::G_FMA;
6790
6791 auto buildMatchInfo = [=](Register Dst, Register X, Register Y, Register Z,
6792 unsigned Flags, MachineIRBuilder &B) {
6793 Register FpExtX = B.buildFPExt(Res: DstTy, Op: X).getReg(Idx: 0);
6794 Register FpExtY = B.buildFPExt(Res: DstTy, Op: Y).getReg(Idx: 0);
6795 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {Dst}, SrcOps: {FpExtX, FpExtY, Z}, Flags);
6796 };
6797
6798 MachineInstr *FMulMI;
6799 // fold (fsub (fpext (fneg (fmul x, y))), z) ->
6800 // (fneg (fma (fpext x), (fpext y), z))
6801 // fold (fsub (fneg (fpext (fmul x, y))), z) ->
6802 // (fneg (fma (fpext x), (fpext y), z))
6803 if ((mi_match(R: LHSReg, MRI, P: m_GFPExt(Src: m_GFNeg(Src: m_MInstr(MI&: FMulMI)))) ||
6804 mi_match(R: LHSReg, MRI, P: m_GFNeg(Src: m_GFPExt(Src: m_MInstr(MI&: FMulMI))))) &&
6805 isContractableFMul(MI&: *FMulMI, AllowFusionGlobally) &&
6806 TLI.isFPExtFoldable(MI, Opcode: PreferredFusedOpcode, DestTy: DstTy,
6807 SrcTy: MRI.getType(Reg: FMulMI->getOperand(i: 0).getReg()))) {
6808 unsigned Flags = MI.getFlags() & FMulMI->getFlags();
6809 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6810 Register FMAReg = MRI.createGenericVirtualRegister(Ty: DstTy);
6811 buildMatchInfo(FMAReg, FMulMI->getOperand(i: 1).getReg(),
6812 FMulMI->getOperand(i: 2).getReg(), RHSReg, Flags, B);
6813 B.buildFNeg(Dst: MI.getOperand(i: 0).getReg(), Src0: FMAReg);
6814 };
6815 return true;
6816 }
6817
6818 // fold (fsub x, (fpext (fneg (fmul y, z)))) -> (fma (fpext y), (fpext z), x)
6819 // fold (fsub x, (fneg (fpext (fmul y, z)))) -> (fma (fpext y), (fpext z), x)
6820 if ((mi_match(R: RHSReg, MRI, P: m_GFPExt(Src: m_GFNeg(Src: m_MInstr(MI&: FMulMI)))) ||
6821 mi_match(R: RHSReg, MRI, P: m_GFNeg(Src: m_GFPExt(Src: m_MInstr(MI&: FMulMI))))) &&
6822 isContractableFMul(MI&: *FMulMI, AllowFusionGlobally) &&
6823 TLI.isFPExtFoldable(MI, Opcode: PreferredFusedOpcode, DestTy: DstTy,
6824 SrcTy: MRI.getType(Reg: FMulMI->getOperand(i: 0).getReg()))) {
6825 unsigned Flags = MI.getFlags() & FMulMI->getFlags();
6826 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6827 buildMatchInfo(MI.getOperand(i: 0).getReg(), FMulMI->getOperand(i: 1).getReg(),
6828 FMulMI->getOperand(i: 2).getReg(), LHSReg, Flags, B);
6829 };
6830 return true;
6831 }
6832
6833 return false;
6834}
6835
6836bool CombinerHelper::matchCombineFMinMaxNaN(MachineInstr &MI,
6837 unsigned &IdxToPropagate) const {
6838 bool PropagateNaN;
6839 switch (MI.getOpcode()) {
6840 default:
6841 return false;
6842 case TargetOpcode::G_FMINNUM:
6843 case TargetOpcode::G_FMAXNUM:
6844 PropagateNaN = false;
6845 break;
6846 case TargetOpcode::G_FMINIMUM:
6847 case TargetOpcode::G_FMAXIMUM:
6848 PropagateNaN = true;
6849 break;
6850 }
6851
6852 auto MatchNaN = [&](unsigned Idx) {
6853 Register MaybeNaNReg = MI.getOperand(i: Idx).getReg();
6854 const ConstantFP *MaybeCst = getConstantFPVRegVal(VReg: MaybeNaNReg, MRI);
6855 if (!MaybeCst || !MaybeCst->getValueAPF().isNaN())
6856 return false;
6857 IdxToPropagate = PropagateNaN ? Idx : (Idx == 1 ? 2 : 1);
6858 return true;
6859 };
6860
6861 return MatchNaN(1) || MatchNaN(2);
6862}
6863
6864// Combine multiple FDIVs with the same divisor into multiple FMULs by the
6865// reciprocal.
6866// E.g., (a / Y; b / Y;) -> (recip = 1.0 / Y; a * recip; b * recip)
6867bool CombinerHelper::matchRepeatedFPDivisor(
6868 MachineInstr &MI, SmallVector<MachineInstr *> &MatchInfo) const {
6869 assert(MI.getOpcode() == TargetOpcode::G_FDIV);
6870
6871 Register X = MI.getOperand(i: 1).getReg();
6872 Register Y = MI.getOperand(i: 2).getReg();
6873
6874 if (!MI.getFlag(Flag: MachineInstr::MIFlag::FmArcp))
6875 return false;
6876
6877 auto IsOne = [this](Register X) {
6878 auto N0CFP = isConstantOrConstantSplatVectorFP(Def: X, MRI);
6879 return N0CFP && (N0CFP->isOne() || N0CFP->isMinusOne());
6880 };
6881
6882 // Skip if current node is a reciprocal/fneg-reciprocal.
6883 if (IsOne(X))
6884 return false;
6885
6886 // Exit early if the target does not want this transform or if there can't
6887 // possibly be enough uses of the divisor to make the transform worthwhile.
6888 unsigned MinUses = getTargetLowering().combineRepeatedFPDivisors();
6889 if (!MinUses)
6890 return false;
6891
6892 // Find all FDIV users of the same divisor. For the moment we limit all
6893 // instructions to a single BB and use the first Instr in MatchInfo as the
6894 // dominating position.
6895 MatchInfo.push_back(Elt: &MI);
6896 for (auto &U : MRI.use_nodbg_instructions(Reg: Y)) {
6897 if (&U == &MI || U.getParent() != MI.getParent())
6898 continue;
6899 if (U.getOpcode() == TargetOpcode::G_FDIV &&
6900 U.getOperand(i: 2).getReg() == Y && U.getOperand(i: 1).getReg() != Y &&
6901 !IsOne(U.getOperand(i: 1).getReg())) {
6902 // This division is eligible for optimization only if global unsafe math
6903 // is enabled or if this division allows reciprocal formation.
6904 if (U.getFlag(Flag: MachineInstr::MIFlag::FmArcp)) {
6905 MatchInfo.push_back(Elt: &U);
6906 if (dominates(DefMI: U, UseMI: *MatchInfo[0]))
6907 std::swap(a&: MatchInfo[0], b&: MatchInfo.back());
6908 }
6909 }
6910 }
6911
6912 // Now that we have the actual number of divisor uses, make sure it meets
6913 // the minimum threshold specified by the target.
6914 return MatchInfo.size() >= MinUses;
6915}
6916
6917void CombinerHelper::applyRepeatedFPDivisor(
6918 SmallVector<MachineInstr *> &MatchInfo) const {
6919 // Generate the new div at the position of the first instruction, that we have
6920 // ensured will dominate all other instructions.
6921 Builder.setInsertPt(MBB&: *MatchInfo[0]->getParent(), II: MatchInfo[0]);
6922 LLT Ty = MRI.getType(Reg: MatchInfo[0]->getOperand(i: 0).getReg());
6923 auto Div = Builder.buildFDiv(Dst: Ty, Src0: Builder.buildFConstant(Res: Ty, Val: 1.0),
6924 Src1: MatchInfo[0]->getOperand(i: 2).getReg(),
6925 Flags: MatchInfo[0]->getFlags());
6926
6927 // Replace all found div's with fmul instructions.
6928 for (MachineInstr *MI : MatchInfo) {
6929 Builder.setInsertPt(MBB&: *MI->getParent(), II: MI);
6930 Builder.buildFMul(Dst: MI->getOperand(i: 0).getReg(), Src0: MI->getOperand(i: 1).getReg(),
6931 Src1: Div->getOperand(i: 0).getReg(), Flags: MI->getFlags());
6932 MI->eraseFromParent();
6933 }
6934}
6935
6936bool CombinerHelper::matchBuildVectorIdentityFold(MachineInstr &MI,
6937 Register &MatchInfo) const {
6938 // This combine folds the following patterns:
6939 //
6940 // G_BUILD_VECTOR_TRUNC (G_BITCAST(x), G_LSHR(G_BITCAST(x), k))
6941 // G_BUILD_VECTOR(G_TRUNC(G_BITCAST(x)), G_TRUNC(G_LSHR(G_BITCAST(x), k)))
6942 // into
6943 // x
6944 // if
6945 // k == sizeof(VecEltTy)/2
6946 // type(x) == type(dst)
6947 //
6948 // G_BUILD_VECTOR(G_TRUNC(G_BITCAST(x)), undef)
6949 // into
6950 // x
6951 // if
6952 // type(x) == type(dst)
6953
6954 LLT DstVecTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6955 LLT DstEltTy = DstVecTy.getElementType();
6956
6957 Register Lo, Hi;
6958
6959 if (mi_match(
6960 MI, MRI,
6961 P: m_GBuildVector(L: m_GTrunc(Src: m_GBitcast(Src: m_Reg(R&: Lo))), R: m_GImplicitDef()))) {
6962 MatchInfo = Lo;
6963 return MRI.getType(Reg: MatchInfo) == DstVecTy;
6964 }
6965
6966 std::optional<ValueAndVReg> ShiftAmount;
6967 const auto LoPattern = m_GBitcast(Src: m_Reg(R&: Lo));
6968 const auto HiPattern = m_GLShr(L: m_GBitcast(Src: m_Reg(R&: Hi)), R: m_GCst(ValReg&: ShiftAmount));
6969 if (mi_match(
6970 MI, MRI,
6971 P: m_any_of(preds: m_GBuildVectorTrunc(L: LoPattern, R: HiPattern),
6972 preds: m_GBuildVector(L: m_GTrunc(Src: LoPattern), R: m_GTrunc(Src: HiPattern))))) {
6973 if (Lo == Hi && ShiftAmount->Value == DstEltTy.getSizeInBits()) {
6974 MatchInfo = Lo;
6975 return MRI.getType(Reg: MatchInfo) == DstVecTy;
6976 }
6977 }
6978
6979 return false;
6980}
6981
6982bool CombinerHelper::matchTruncBuildVectorFold(MachineInstr &MI,
6983 Register &MatchInfo) const {
6984 // Replace (G_TRUNC (G_BITCAST (G_BUILD_VECTOR x, y)) with just x
6985 // if type(x) == type(G_TRUNC)
6986 if (!mi_match(R: MI.getOperand(i: 1).getReg(), MRI,
6987 P: m_GBitcast(Src: m_GBuildVector(L: m_Reg(R&: MatchInfo), R: m_Reg()))))
6988 return false;
6989
6990 return MRI.getType(Reg: MatchInfo) == MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6991}
6992
6993bool CombinerHelper::matchTruncLshrBuildVectorFold(MachineInstr &MI,
6994 Register &MatchInfo) const {
6995 // Replace (G_TRUNC (G_LSHR (G_BITCAST (G_BUILD_VECTOR x, y)), K)) with
6996 // y if K == size of vector element type
6997 std::optional<ValueAndVReg> ShiftAmt;
6998 if (!mi_match(R: MI.getOperand(i: 1).getReg(), MRI,
6999 P: m_GLShr(L: m_GBitcast(Src: m_GBuildVector(L: m_Reg(), R: m_Reg(R&: MatchInfo))),
7000 R: m_GCst(ValReg&: ShiftAmt))))
7001 return false;
7002
7003 LLT MatchTy = MRI.getType(Reg: MatchInfo);
7004 return ShiftAmt->Value.getZExtValue() == MatchTy.getSizeInBits() &&
7005 MatchTy == MRI.getType(Reg: MI.getOperand(i: 0).getReg());
7006}
7007
7008unsigned CombinerHelper::getFPMinMaxOpcForSelect(
7009 CmpInst::Predicate Pred, LLT DstTy,
7010 SelectPatternNaNBehaviour VsNaNRetVal) const {
7011 assert(VsNaNRetVal != SelectPatternNaNBehaviour::NOT_APPLICABLE &&
7012 "Expected a NaN behaviour?");
7013 // Choose an opcode based off of legality or the behaviour when one of the
7014 // LHS/RHS may be NaN.
7015 switch (Pred) {
7016 default:
7017 return 0;
7018 case CmpInst::FCMP_UGT:
7019 case CmpInst::FCMP_UGE:
7020 case CmpInst::FCMP_OGT:
7021 case CmpInst::FCMP_OGE:
7022 if (VsNaNRetVal == SelectPatternNaNBehaviour::RETURNS_OTHER)
7023 return TargetOpcode::G_FMAXNUM;
7024 if (VsNaNRetVal == SelectPatternNaNBehaviour::RETURNS_NAN)
7025 return TargetOpcode::G_FMAXIMUM;
7026 if (isLegal(Query: {TargetOpcode::G_FMAXNUM, {DstTy}}))
7027 return TargetOpcode::G_FMAXNUM;
7028 if (isLegal(Query: {TargetOpcode::G_FMAXIMUM, {DstTy}}))
7029 return TargetOpcode::G_FMAXIMUM;
7030 return 0;
7031 case CmpInst::FCMP_ULT:
7032 case CmpInst::FCMP_ULE:
7033 case CmpInst::FCMP_OLT:
7034 case CmpInst::FCMP_OLE:
7035 if (VsNaNRetVal == SelectPatternNaNBehaviour::RETURNS_OTHER)
7036 return TargetOpcode::G_FMINNUM;
7037 if (VsNaNRetVal == SelectPatternNaNBehaviour::RETURNS_NAN)
7038 return TargetOpcode::G_FMINIMUM;
7039 if (isLegal(Query: {TargetOpcode::G_FMINNUM, {DstTy}}))
7040 return TargetOpcode::G_FMINNUM;
7041 if (!isLegal(Query: {TargetOpcode::G_FMINIMUM, {DstTy}}))
7042 return 0;
7043 return TargetOpcode::G_FMINIMUM;
7044 }
7045}
7046
7047CombinerHelper::SelectPatternNaNBehaviour
7048CombinerHelper::computeRetValAgainstNaN(Register LHS, Register RHS,
7049 bool IsOrderedComparison) const {
7050 bool LHSSafe = VT->isKnownNeverNaN(Val: LHS);
7051 bool RHSSafe = VT->isKnownNeverNaN(Val: RHS);
7052 // Completely unsafe.
7053 if (!LHSSafe && !RHSSafe)
7054 return SelectPatternNaNBehaviour::NOT_APPLICABLE;
7055 if (LHSSafe && RHSSafe)
7056 return SelectPatternNaNBehaviour::RETURNS_ANY;
7057 // An ordered comparison will return false when given a NaN, so it
7058 // returns the RHS.
7059 if (IsOrderedComparison)
7060 return LHSSafe ? SelectPatternNaNBehaviour::RETURNS_NAN
7061 : SelectPatternNaNBehaviour::RETURNS_OTHER;
7062 // An unordered comparison will return true when given a NaN, so it
7063 // returns the LHS.
7064 return LHSSafe ? SelectPatternNaNBehaviour::RETURNS_OTHER
7065 : SelectPatternNaNBehaviour::RETURNS_NAN;
7066}
7067
7068bool CombinerHelper::matchFPSelectToMinMax(Register Dst, Register Cond,
7069 Register TrueVal, Register FalseVal,
7070 BuildFnTy &MatchInfo) const {
7071 // Match: select (fcmp cond x, y) x, y
7072 // select (fcmp cond x, y) y, x
7073 // And turn it into fminnum/fmaxnum or fmin/fmax based off of the condition.
7074 LLT DstTy = MRI.getType(Reg: Dst);
7075 // Bail out early on pointers, since we'll never want to fold to a min/max.
7076 if (DstTy.isPointer())
7077 return false;
7078 // Match a floating point compare with a less-than/greater-than predicate.
7079 // TODO: Allow multiple users of the compare if they are all selects.
7080 CmpInst::Predicate Pred;
7081 Register CmpLHS, CmpRHS;
7082 if (!mi_match(R: Cond, MRI,
7083 P: m_OneNonDBGUse(
7084 SP: m_GFCmp(P: m_Pred(P&: Pred), L: m_Reg(R&: CmpLHS), R: m_Reg(R&: CmpRHS)))) ||
7085 CmpInst::isEquality(pred: Pred))
7086 return false;
7087 SelectPatternNaNBehaviour ResWithKnownNaNInfo =
7088 computeRetValAgainstNaN(LHS: CmpLHS, RHS: CmpRHS, IsOrderedComparison: CmpInst::isOrdered(predicate: Pred));
7089 if (ResWithKnownNaNInfo == SelectPatternNaNBehaviour::NOT_APPLICABLE)
7090 return false;
7091 if (TrueVal == CmpRHS && FalseVal == CmpLHS) {
7092 std::swap(a&: CmpLHS, b&: CmpRHS);
7093 Pred = CmpInst::getSwappedPredicate(pred: Pred);
7094 if (ResWithKnownNaNInfo == SelectPatternNaNBehaviour::RETURNS_NAN)
7095 ResWithKnownNaNInfo = SelectPatternNaNBehaviour::RETURNS_OTHER;
7096 else if (ResWithKnownNaNInfo == SelectPatternNaNBehaviour::RETURNS_OTHER)
7097 ResWithKnownNaNInfo = SelectPatternNaNBehaviour::RETURNS_NAN;
7098 }
7099 if (TrueVal != CmpLHS || FalseVal != CmpRHS)
7100 return false;
7101 // Decide what type of max/min this should be based off of the predicate.
7102 unsigned Opc = getFPMinMaxOpcForSelect(Pred, DstTy, VsNaNRetVal: ResWithKnownNaNInfo);
7103 if (!Opc || !isLegal(Query: {Opc, {DstTy}}))
7104 return false;
7105 // Comparisons between signed zero and zero may have different results...
7106 // unless we have fmaximum/fminimum. In that case, we know -0 < 0.
7107 if (Opc != TargetOpcode::G_FMAXIMUM && Opc != TargetOpcode::G_FMINIMUM) {
7108 // We don't know if a comparison between two 0s will give us a consistent
7109 // result. Be conservative and only proceed if at least one side is
7110 // non-zero.
7111 auto KnownNonZeroSide = getFConstantVRegValWithLookThrough(VReg: CmpLHS, MRI);
7112 if (!KnownNonZeroSide || !KnownNonZeroSide->Value.isNonZero()) {
7113 KnownNonZeroSide = getFConstantVRegValWithLookThrough(VReg: CmpRHS, MRI);
7114 if (!KnownNonZeroSide || !KnownNonZeroSide->Value.isNonZero())
7115 return false;
7116 }
7117 }
7118 MatchInfo = [=](MachineIRBuilder &B) {
7119 B.buildInstr(Opc, DstOps: {Dst}, SrcOps: {CmpLHS, CmpRHS});
7120 };
7121 return true;
7122}
7123
7124bool CombinerHelper::matchSimplifySelectToMinMax(MachineInstr &MI,
7125 BuildFnTy &MatchInfo) const {
7126 // TODO: Handle integer cases.
7127 assert(MI.getOpcode() == TargetOpcode::G_SELECT);
7128 // Condition may be fed by a truncated compare.
7129 Register Cond = MI.getOperand(i: 1).getReg();
7130 Register MaybeTrunc;
7131 if (mi_match(R: Cond, MRI, P: m_OneNonDBGUse(SP: m_GTrunc(Src: m_Reg(R&: MaybeTrunc)))))
7132 Cond = MaybeTrunc;
7133 Register Dst = MI.getOperand(i: 0).getReg();
7134 Register TrueVal = MI.getOperand(i: 2).getReg();
7135 Register FalseVal = MI.getOperand(i: 3).getReg();
7136 return matchFPSelectToMinMax(Dst, Cond, TrueVal, FalseVal, MatchInfo);
7137}
7138
7139bool CombinerHelper::matchRedundantBinOpInEquality(MachineInstr &MI,
7140 BuildFnTy &MatchInfo) const {
7141 assert(MI.getOpcode() == TargetOpcode::G_ICMP);
7142 // (X + Y) == X --> Y == 0
7143 // (X + Y) != X --> Y != 0
7144 // (X - Y) == X --> Y == 0
7145 // (X - Y) != X --> Y != 0
7146 // (X ^ Y) == X --> Y == 0
7147 // (X ^ Y) != X --> Y != 0
7148 Register Dst = MI.getOperand(i: 0).getReg();
7149 CmpInst::Predicate Pred;
7150 Register X, Y, OpLHS, OpRHS;
7151 bool MatchedSub = mi_match(
7152 R: Dst, MRI,
7153 P: m_c_GICmp(P: m_Pred(P&: Pred), L: m_Reg(R&: X), R: m_GSub(L: m_Reg(R&: OpLHS), R: m_Reg(R&: Y))));
7154 if (MatchedSub && X != OpLHS)
7155 return false;
7156 if (!MatchedSub) {
7157 if (!mi_match(R: Dst, MRI,
7158 P: m_c_GICmp(P: m_Pred(P&: Pred), L: m_Reg(R&: X),
7159 R: m_any_of(preds: m_GAdd(L: m_Reg(R&: OpLHS), R: m_Reg(R&: OpRHS)),
7160 preds: m_GXor(L: m_Reg(R&: OpLHS), R: m_Reg(R&: OpRHS))))))
7161 return false;
7162 Y = X == OpLHS ? OpRHS : X == OpRHS ? OpLHS : Register();
7163 }
7164 MatchInfo = [=](MachineIRBuilder &B) {
7165 auto Zero = B.buildConstant(Res: MRI.getType(Reg: Y), Val: 0);
7166 B.buildICmp(Pred, Res: Dst, Op0: Y, Op1: Zero);
7167 };
7168 return CmpInst::isEquality(pred: Pred) && Y.isValid();
7169}
7170
7171/// Return the minimum useless shift amount that results in complete loss of the
7172/// source value. Return std::nullopt when it cannot determine a value.
7173static std::optional<unsigned>
7174getMinUselessShift(KnownBits ValueKB, unsigned Opcode,
7175 std::optional<int64_t> &Result) {
7176 assert((Opcode == TargetOpcode::G_SHL || Opcode == TargetOpcode::G_LSHR ||
7177 Opcode == TargetOpcode::G_ASHR) &&
7178 "Expect G_SHL, G_LSHR or G_ASHR.");
7179 auto SignificantBits = 0;
7180 switch (Opcode) {
7181 case TargetOpcode::G_SHL:
7182 SignificantBits = ValueKB.countMinTrailingZeros();
7183 Result = 0;
7184 break;
7185 case TargetOpcode::G_LSHR:
7186 Result = 0;
7187 SignificantBits = ValueKB.countMinLeadingZeros();
7188 break;
7189 case TargetOpcode::G_ASHR:
7190 if (ValueKB.isNonNegative()) {
7191 SignificantBits = ValueKB.countMinLeadingZeros();
7192 Result = 0;
7193 } else if (ValueKB.isNegative()) {
7194 SignificantBits = ValueKB.countMinLeadingOnes();
7195 Result = -1;
7196 } else {
7197 // Cannot determine shift result.
7198 Result = std::nullopt;
7199 }
7200 break;
7201 default:
7202 break;
7203 }
7204 return ValueKB.getBitWidth() - SignificantBits;
7205}
7206
7207bool CombinerHelper::matchShiftsTooBig(
7208 MachineInstr &MI, std::optional<int64_t> &MatchInfo) const {
7209 Register ShiftVal = MI.getOperand(i: 1).getReg();
7210 Register ShiftReg = MI.getOperand(i: 2).getReg();
7211 LLT ResTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
7212 auto IsShiftTooBig = [&](const Constant *C) {
7213 auto *CI = dyn_cast<ConstantInt>(Val: C);
7214 if (!CI)
7215 return false;
7216 if (CI->uge(Num: ResTy.getScalarSizeInBits())) {
7217 MatchInfo = std::nullopt;
7218 return true;
7219 }
7220 auto OptMaxUsefulShift = getMinUselessShift(ValueKB: VT->getKnownBits(R: ShiftVal),
7221 Opcode: MI.getOpcode(), Result&: MatchInfo);
7222 return OptMaxUsefulShift && CI->uge(Num: *OptMaxUsefulShift);
7223 };
7224 return matchUnaryPredicate(MRI, Reg: ShiftReg, Match: IsShiftTooBig);
7225}
7226
7227bool CombinerHelper::matchCommuteConstantToRHS(MachineInstr &MI) const {
7228 unsigned LHSOpndIdx = 1;
7229 unsigned RHSOpndIdx = 2;
7230 switch (MI.getOpcode()) {
7231 case TargetOpcode::G_UADDO:
7232 case TargetOpcode::G_SADDO:
7233 case TargetOpcode::G_UMULO:
7234 case TargetOpcode::G_SMULO:
7235 LHSOpndIdx = 2;
7236 RHSOpndIdx = 3;
7237 break;
7238 default:
7239 break;
7240 }
7241 Register LHS = MI.getOperand(i: LHSOpndIdx).getReg();
7242 Register RHS = MI.getOperand(i: RHSOpndIdx).getReg();
7243 MachineInstr *LHSDef, *RHSDef;
7244 if (!mi_match(R: LHS, MRI, P: m_MInstr(MI&: LHSDef)) ||
7245 !mi_match(R: RHS, MRI, P: m_MInstr(MI&: RHSDef)))
7246 return false;
7247
7248 if (!getIConstantVRegVal(VReg: LHS, MRI)) {
7249 // Skip commuting if LHS is not a constant. But, LHS may be a
7250 // G_CONSTANT_FOLD_BARRIER. If so we commute as long as we don't already
7251 // have a constant on the RHS.
7252 if (LHSDef->getOpcode() != TargetOpcode::G_CONSTANT_FOLD_BARRIER)
7253 return false;
7254 }
7255 // Commute as long as RHS is not a constant or G_CONSTANT_FOLD_BARRIER.
7256 return RHSDef->getOpcode() != TargetOpcode::G_CONSTANT_FOLD_BARRIER &&
7257 !getIConstantVRegVal(VReg: RHS, MRI);
7258}
7259
7260bool CombinerHelper::matchCommuteFPConstantToRHS(MachineInstr &MI) const {
7261 Register LHS = MI.getOperand(i: 1).getReg();
7262 Register RHS = MI.getOperand(i: 2).getReg();
7263 std::optional<FPValueAndVReg> ValAndVReg;
7264 if (!mi_match(R: LHS, MRI, P: m_GFCstOrSplat(FPValReg&: ValAndVReg)))
7265 return false;
7266 return !mi_match(R: RHS, MRI, P: m_GFCstOrSplat(FPValReg&: ValAndVReg));
7267}
7268
7269void CombinerHelper::applyCommuteBinOpOperands(MachineInstr &MI) const {
7270 Observer.changingInstr(MI);
7271 unsigned LHSOpndIdx = 1;
7272 unsigned RHSOpndIdx = 2;
7273 switch (MI.getOpcode()) {
7274 case TargetOpcode::G_UADDO:
7275 case TargetOpcode::G_SADDO:
7276 case TargetOpcode::G_UMULO:
7277 case TargetOpcode::G_SMULO:
7278 LHSOpndIdx = 2;
7279 RHSOpndIdx = 3;
7280 break;
7281 default:
7282 break;
7283 }
7284 Register LHSReg = MI.getOperand(i: LHSOpndIdx).getReg();
7285 Register RHSReg = MI.getOperand(i: RHSOpndIdx).getReg();
7286 MI.getOperand(i: LHSOpndIdx).setReg(RHSReg);
7287 MI.getOperand(i: RHSOpndIdx).setReg(LHSReg);
7288 Observer.changedInstr(MI);
7289}
7290
7291bool CombinerHelper::isOneOrOneSplat(Register Src, bool AllowUndefs) const {
7292 LLT SrcTy = MRI.getType(Reg: Src);
7293 if (SrcTy.isFixedVector())
7294 return isConstantSplatVector(Src, SplatValue: 1, AllowUndefs);
7295 if (SrcTy.isScalar()) {
7296 if (AllowUndefs && getOpcodeDef<GImplicitDef>(Reg: Src, MRI) != nullptr)
7297 return true;
7298 auto IConstant = getIConstantVRegValWithLookThrough(VReg: Src, MRI);
7299 return IConstant && IConstant->Value == 1;
7300 }
7301 return false; // scalable vector
7302}
7303
7304bool CombinerHelper::isZeroOrZeroSplat(Register Src, bool AllowUndefs) const {
7305 LLT SrcTy = MRI.getType(Reg: Src);
7306 if (SrcTy.isFixedVector())
7307 return isConstantSplatVector(Src, SplatValue: 0, AllowUndefs);
7308 if (SrcTy.isScalar()) {
7309 if (AllowUndefs && getOpcodeDef<GImplicitDef>(Reg: Src, MRI) != nullptr)
7310 return true;
7311 auto IConstant = getIConstantVRegValWithLookThrough(VReg: Src, MRI);
7312 return IConstant && IConstant->Value == 0;
7313 }
7314 return false; // scalable vector
7315}
7316
7317// Ignores COPYs during conformance checks.
7318// FIXME scalable vectors.
7319bool CombinerHelper::isConstantSplatVector(Register Src, int64_t SplatValue,
7320 bool AllowUndefs) const {
7321 GBuildVector *BuildVector = getOpcodeDef<GBuildVector>(Reg: Src, MRI);
7322 if (!BuildVector)
7323 return false;
7324 unsigned NumSources = BuildVector->getNumSources();
7325
7326 for (unsigned I = 0; I < NumSources; ++I) {
7327 GImplicitDef *ImplicitDef =
7328 getOpcodeDef<GImplicitDef>(Reg: BuildVector->getSourceReg(I), MRI);
7329 if (ImplicitDef && AllowUndefs)
7330 continue;
7331 if (ImplicitDef && !AllowUndefs)
7332 return false;
7333 std::optional<ValueAndVReg> IConstant =
7334 getIConstantVRegValWithLookThrough(VReg: BuildVector->getSourceReg(I), MRI);
7335 if (IConstant && IConstant->Value == SplatValue)
7336 continue;
7337 return false;
7338 }
7339 return true;
7340}
7341
7342// Ignores COPYs during lookups.
7343// FIXME scalable vectors
7344std::optional<APInt>
7345CombinerHelper::getConstantOrConstantSplatVector(Register Src) const {
7346 auto IConstant = getIConstantVRegValWithLookThrough(VReg: Src, MRI);
7347 if (IConstant)
7348 return IConstant->Value;
7349
7350 GBuildVector *BuildVector = getOpcodeDef<GBuildVector>(Reg: Src, MRI);
7351 if (!BuildVector)
7352 return std::nullopt;
7353 unsigned NumSources = BuildVector->getNumSources();
7354
7355 std::optional<APInt> Value = std::nullopt;
7356 for (unsigned I = 0; I < NumSources; ++I) {
7357 std::optional<ValueAndVReg> IConstant =
7358 getIConstantVRegValWithLookThrough(VReg: BuildVector->getSourceReg(I), MRI);
7359 if (!IConstant)
7360 return std::nullopt;
7361 if (!Value)
7362 Value = IConstant->Value;
7363 else if (*Value != IConstant->Value)
7364 return std::nullopt;
7365 }
7366 return Value;
7367}
7368
7369// FIXME G_SPLAT_VECTOR
7370bool CombinerHelper::isConstantOrConstantVectorI(Register Src) const {
7371 auto IConstant = getIConstantVRegValWithLookThrough(VReg: Src, MRI);
7372 if (IConstant)
7373 return true;
7374
7375 GBuildVector *BuildVector = getOpcodeDef<GBuildVector>(Reg: Src, MRI);
7376 if (!BuildVector)
7377 return false;
7378
7379 unsigned NumSources = BuildVector->getNumSources();
7380 for (unsigned I = 0; I < NumSources; ++I) {
7381 std::optional<ValueAndVReg> IConstant =
7382 getIConstantVRegValWithLookThrough(VReg: BuildVector->getSourceReg(I), MRI);
7383 if (!IConstant)
7384 return false;
7385 }
7386 return true;
7387}
7388
7389// TODO: use knownbits to determine zeros
7390bool CombinerHelper::tryFoldSelectOfConstants(GSelect *Select,
7391 BuildFnTy &MatchInfo) const {
7392 uint32_t Flags = Select->getFlags();
7393 Register Dest = Select->getReg(Idx: 0);
7394 Register Cond = Select->getCondReg();
7395 Register True = Select->getTrueReg();
7396 Register False = Select->getFalseReg();
7397 LLT CondTy = MRI.getType(Reg: Select->getCondReg());
7398 LLT TrueTy = MRI.getType(Reg: Select->getTrueReg());
7399
7400 // We only do this combine for scalar boolean conditions.
7401 if (CondTy != LLT::scalar(SizeInBits: 1))
7402 return false;
7403
7404 if (TrueTy.isPointer())
7405 return false;
7406
7407 // Both are scalars.
7408 std::optional<ValueAndVReg> TrueOpt =
7409 getIConstantVRegValWithLookThrough(VReg: True, MRI);
7410 std::optional<ValueAndVReg> FalseOpt =
7411 getIConstantVRegValWithLookThrough(VReg: False, MRI);
7412
7413 if (!TrueOpt || !FalseOpt)
7414 return false;
7415
7416 APInt TrueValue = TrueOpt->Value;
7417 APInt FalseValue = FalseOpt->Value;
7418
7419 // select Cond, 1, 0 --> zext (Cond)
7420 if (TrueValue.isOne() && FalseValue.isZero()) {
7421 MatchInfo = [=](MachineIRBuilder &B) {
7422 B.setInstrAndDebugLoc(*Select);
7423 B.buildZExtOrTrunc(Res: Dest, Op: Cond);
7424 };
7425 return true;
7426 }
7427
7428 // select Cond, -1, 0 --> sext (Cond)
7429 if (TrueValue.isAllOnes() && FalseValue.isZero()) {
7430 MatchInfo = [=](MachineIRBuilder &B) {
7431 B.setInstrAndDebugLoc(*Select);
7432 B.buildSExtOrTrunc(Res: Dest, Op: Cond);
7433 };
7434 return true;
7435 }
7436
7437 // select Cond, 0, 1 --> zext (!Cond)
7438 if (TrueValue.isZero() && FalseValue.isOne()) {
7439 MatchInfo = [=](MachineIRBuilder &B) {
7440 B.setInstrAndDebugLoc(*Select);
7441 Register Inner = MRI.createGenericVirtualRegister(Ty: CondTy);
7442 B.buildNot(Dst: Inner, Src0: Cond);
7443 B.buildZExtOrTrunc(Res: Dest, Op: Inner);
7444 };
7445 return true;
7446 }
7447
7448 // select Cond, 0, -1 --> sext (!Cond)
7449 if (TrueValue.isZero() && FalseValue.isAllOnes()) {
7450 MatchInfo = [=](MachineIRBuilder &B) {
7451 B.setInstrAndDebugLoc(*Select);
7452 Register Inner = MRI.createGenericVirtualRegister(Ty: CondTy);
7453 B.buildNot(Dst: Inner, Src0: Cond);
7454 B.buildSExtOrTrunc(Res: Dest, Op: Inner);
7455 };
7456 return true;
7457 }
7458
7459 // select Cond, C1, C1-1 --> add (zext Cond), C1-1
7460 if (TrueValue - 1 == FalseValue) {
7461 MatchInfo = [=](MachineIRBuilder &B) {
7462 B.setInstrAndDebugLoc(*Select);
7463 Register Inner = MRI.createGenericVirtualRegister(Ty: TrueTy);
7464 B.buildZExtOrTrunc(Res: Inner, Op: Cond);
7465 B.buildAdd(Dst: Dest, Src0: Inner, Src1: False);
7466 };
7467 return true;
7468 }
7469
7470 // select Cond, C1, C1+1 --> add (sext Cond), C1+1
7471 if (TrueValue + 1 == FalseValue) {
7472 MatchInfo = [=](MachineIRBuilder &B) {
7473 B.setInstrAndDebugLoc(*Select);
7474 Register Inner = MRI.createGenericVirtualRegister(Ty: TrueTy);
7475 B.buildSExtOrTrunc(Res: Inner, Op: Cond);
7476 B.buildAdd(Dst: Dest, Src0: Inner, Src1: False);
7477 };
7478 return true;
7479 }
7480
7481 // select Cond, Pow2, 0 --> (zext Cond) << log2(Pow2)
7482 if (TrueValue.isPowerOf2() && FalseValue.isZero()) {
7483 MatchInfo = [=](MachineIRBuilder &B) {
7484 B.setInstrAndDebugLoc(*Select);
7485 Register Inner = MRI.createGenericVirtualRegister(Ty: TrueTy);
7486 B.buildZExtOrTrunc(Res: Inner, Op: Cond);
7487 // The shift amount must be scalar.
7488 LLT ShiftTy = TrueTy.isVector() ? TrueTy.getElementType() : TrueTy;
7489 auto ShAmtC = B.buildConstant(Res: ShiftTy, Val: TrueValue.exactLogBase2());
7490 B.buildShl(Dst: Dest, Src0: Inner, Src1: ShAmtC, Flags);
7491 };
7492 return true;
7493 }
7494
7495 // select Cond, 0, Pow2 --> (zext (!Cond)) << log2(Pow2)
7496 if (FalseValue.isPowerOf2() && TrueValue.isZero()) {
7497 MatchInfo = [=](MachineIRBuilder &B) {
7498 B.setInstrAndDebugLoc(*Select);
7499 Register Not = MRI.createGenericVirtualRegister(Ty: CondTy);
7500 B.buildNot(Dst: Not, Src0: Cond);
7501 Register Inner = MRI.createGenericVirtualRegister(Ty: TrueTy);
7502 B.buildZExtOrTrunc(Res: Inner, Op: Not);
7503 // The shift amount must be scalar.
7504 LLT ShiftTy = TrueTy.isVector() ? TrueTy.getElementType() : TrueTy;
7505 auto ShAmtC = B.buildConstant(Res: ShiftTy, Val: FalseValue.exactLogBase2());
7506 B.buildShl(Dst: Dest, Src0: Inner, Src1: ShAmtC, Flags);
7507 };
7508 return true;
7509 }
7510
7511 // select Cond, -1, C --> or (sext Cond), C
7512 if (TrueValue.isAllOnes()) {
7513 MatchInfo = [=](MachineIRBuilder &B) {
7514 B.setInstrAndDebugLoc(*Select);
7515 Register Inner = MRI.createGenericVirtualRegister(Ty: TrueTy);
7516 B.buildSExtOrTrunc(Res: Inner, Op: Cond);
7517 B.buildOr(Dst: Dest, Src0: Inner, Src1: False, Flags);
7518 };
7519 return true;
7520 }
7521
7522 // select Cond, C, -1 --> or (sext (not Cond)), C
7523 if (FalseValue.isAllOnes()) {
7524 MatchInfo = [=](MachineIRBuilder &B) {
7525 B.setInstrAndDebugLoc(*Select);
7526 Register Not = MRI.createGenericVirtualRegister(Ty: CondTy);
7527 B.buildNot(Dst: Not, Src0: Cond);
7528 Register Inner = MRI.createGenericVirtualRegister(Ty: TrueTy);
7529 B.buildSExtOrTrunc(Res: Inner, Op: Not);
7530 B.buildOr(Dst: Dest, Src0: Inner, Src1: True, Flags);
7531 };
7532 return true;
7533 }
7534
7535 return false;
7536}
7537
7538// TODO: use knownbits to determine zeros
7539bool CombinerHelper::tryFoldBoolSelectToLogic(GSelect *Select,
7540 BuildFnTy &MatchInfo) const {
7541 uint32_t Flags = Select->getFlags();
7542 Register DstReg = Select->getReg(Idx: 0);
7543 Register Cond = Select->getCondReg();
7544 Register True = Select->getTrueReg();
7545 Register False = Select->getFalseReg();
7546 LLT CondTy = MRI.getType(Reg: Select->getCondReg());
7547 LLT TrueTy = MRI.getType(Reg: Select->getTrueReg());
7548
7549 // Boolean or fixed vector of booleans.
7550 if (CondTy.isScalableVector() ||
7551 (CondTy.isFixedVector() &&
7552 CondTy.getElementType().getScalarSizeInBits() != 1) ||
7553 CondTy.getScalarSizeInBits() != 1)
7554 return false;
7555
7556 if (CondTy != TrueTy)
7557 return false;
7558
7559 // select Cond, Cond, F --> or Cond, F
7560 // select Cond, 1, F --> or Cond, F
7561 if ((Cond == True) || isOneOrOneSplat(Src: True, /* AllowUndefs */ true)) {
7562 MatchInfo = [=](MachineIRBuilder &B) {
7563 B.setInstrAndDebugLoc(*Select);
7564 Register Ext = MRI.createGenericVirtualRegister(Ty: TrueTy);
7565 B.buildZExtOrTrunc(Res: Ext, Op: Cond);
7566 auto FreezeFalse = B.buildFreeze(Dst: TrueTy, Src: False);
7567 B.buildOr(Dst: DstReg, Src0: Ext, Src1: FreezeFalse, Flags);
7568 };
7569 return true;
7570 }
7571
7572 // select Cond, T, Cond --> and Cond, T
7573 // select Cond, T, 0 --> and Cond, T
7574 if ((Cond == False) || isZeroOrZeroSplat(Src: False, /* AllowUndefs */ true)) {
7575 MatchInfo = [=](MachineIRBuilder &B) {
7576 B.setInstrAndDebugLoc(*Select);
7577 Register Ext = MRI.createGenericVirtualRegister(Ty: TrueTy);
7578 B.buildZExtOrTrunc(Res: Ext, Op: Cond);
7579 auto FreezeTrue = B.buildFreeze(Dst: TrueTy, Src: True);
7580 B.buildAnd(Dst: DstReg, Src0: Ext, Src1: FreezeTrue);
7581 };
7582 return true;
7583 }
7584
7585 // select Cond, T, 1 --> or (not Cond), T
7586 if (isOneOrOneSplat(Src: False, /* AllowUndefs */ true)) {
7587 MatchInfo = [=](MachineIRBuilder &B) {
7588 B.setInstrAndDebugLoc(*Select);
7589 // First the not.
7590 Register Inner = MRI.createGenericVirtualRegister(Ty: CondTy);
7591 B.buildNot(Dst: Inner, Src0: Cond);
7592 // Then an ext to match the destination register.
7593 Register Ext = MRI.createGenericVirtualRegister(Ty: TrueTy);
7594 B.buildZExtOrTrunc(Res: Ext, Op: Inner);
7595 auto FreezeTrue = B.buildFreeze(Dst: TrueTy, Src: True);
7596 B.buildOr(Dst: DstReg, Src0: Ext, Src1: FreezeTrue, Flags);
7597 };
7598 return true;
7599 }
7600
7601 // select Cond, 0, F --> and (not Cond), F
7602 if (isZeroOrZeroSplat(Src: True, /* AllowUndefs */ true)) {
7603 MatchInfo = [=](MachineIRBuilder &B) {
7604 B.setInstrAndDebugLoc(*Select);
7605 // First the not.
7606 Register Inner = MRI.createGenericVirtualRegister(Ty: CondTy);
7607 B.buildNot(Dst: Inner, Src0: Cond);
7608 // Then an ext to match the destination register.
7609 Register Ext = MRI.createGenericVirtualRegister(Ty: TrueTy);
7610 B.buildZExtOrTrunc(Res: Ext, Op: Inner);
7611 auto FreezeFalse = B.buildFreeze(Dst: TrueTy, Src: False);
7612 B.buildAnd(Dst: DstReg, Src0: Ext, Src1: FreezeFalse);
7613 };
7614 return true;
7615 }
7616
7617 return false;
7618}
7619
7620bool CombinerHelper::matchSelectIMinMax(const MachineOperand &MO,
7621 BuildFnTy &MatchInfo) const {
7622 Register DstReg = MO.getReg();
7623 Register CondReg, True, False;
7624 if (!mi_match(R: DstReg, MRI,
7625 P: m_GISelect(Src0: m_Reg(R&: CondReg), Src1: m_Reg(R&: True), Src2: m_Reg(R&: False))))
7626 return false;
7627
7628 CmpInst::Predicate Pred;
7629 Register CmpLHS, CmpRHS;
7630 if (!mi_match(R: CondReg, MRI,
7631 P: m_GICmp(P: m_Pred(P&: Pred), L: m_Reg(R&: CmpLHS), R: m_Reg(R&: CmpRHS))))
7632 return false;
7633
7634 LLT DstTy = MRI.getType(Reg: DstReg);
7635 if (DstTy.isPointerOrPointerVector())
7636 return false;
7637
7638 // We want to fold the icmp and replace the select.
7639 if (!MRI.hasOneNonDBGUse(RegNo: CondReg))
7640 return false;
7641
7642 // We need a larger or smaller predicate for
7643 // canonicalization.
7644 if (CmpInst::isEquality(pred: Pred))
7645 return false;
7646
7647 // We can swap CmpLHS and CmpRHS for higher hitrate.
7648 if (True == CmpRHS && False == CmpLHS) {
7649 std::swap(a&: CmpLHS, b&: CmpRHS);
7650 Pred = CmpInst::getSwappedPredicate(pred: Pred);
7651 }
7652
7653 // (icmp X, Y) ? X : Y -> integer minmax.
7654 // see matchSelectPattern in ValueTracking.
7655 // Legality between G_SELECT and integer minmax can differ.
7656 if (True != CmpLHS || False != CmpRHS)
7657 return false;
7658
7659 switch (Pred) {
7660 case ICmpInst::ICMP_UGT:
7661 case ICmpInst::ICMP_UGE: {
7662 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_UMAX, DstTy}))
7663 return false;
7664 MatchInfo = [=](MachineIRBuilder &B) { B.buildUMax(Dst: DstReg, Src0: True, Src1: False); };
7665 return true;
7666 }
7667 case ICmpInst::ICMP_SGT:
7668 case ICmpInst::ICMP_SGE: {
7669 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_SMAX, DstTy}))
7670 return false;
7671 MatchInfo = [=](MachineIRBuilder &B) { B.buildSMax(Dst: DstReg, Src0: True, Src1: False); };
7672 return true;
7673 }
7674 case ICmpInst::ICMP_ULT:
7675 case ICmpInst::ICMP_ULE: {
7676 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_UMIN, DstTy}))
7677 return false;
7678 MatchInfo = [=](MachineIRBuilder &B) { B.buildUMin(Dst: DstReg, Src0: True, Src1: False); };
7679 return true;
7680 }
7681 case ICmpInst::ICMP_SLT:
7682 case ICmpInst::ICMP_SLE: {
7683 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_SMIN, DstTy}))
7684 return false;
7685 MatchInfo = [=](MachineIRBuilder &B) { B.buildSMin(Dst: DstReg, Src0: True, Src1: False); };
7686 return true;
7687 }
7688 default:
7689 return false;
7690 }
7691}
7692
7693// (neg (min/max x, (neg x))) --> (max/min x, (neg x))
7694bool CombinerHelper::matchSimplifyNegMinMax(MachineInstr &MI,
7695 BuildFnTy &MatchInfo) const {
7696 assert(MI.getOpcode() == TargetOpcode::G_SUB);
7697 Register DestReg = MI.getOperand(i: 0).getReg();
7698 LLT DestTy = MRI.getType(Reg: DestReg);
7699
7700 Register X;
7701 Register Sub0;
7702 auto NegPattern = m_all_of(preds: m_Neg(Src: m_DeferredReg(R&: X)), preds: m_Reg(R&: Sub0));
7703 if (mi_match(R: DestReg, MRI,
7704 P: m_Neg(Src: m_OneUse(SP: m_any_of(preds: m_GSMin(L: m_Reg(R&: X), R: NegPattern),
7705 preds: m_GSMax(L: m_Reg(R&: X), R: NegPattern),
7706 preds: m_GUMin(L: m_Reg(R&: X), R: NegPattern),
7707 preds: m_GUMax(L: m_Reg(R&: X), R: NegPattern)))))) {
7708 MachineInstr *MinMaxMI;
7709 if (!mi_match(R: MI.getOperand(i: 2).getReg(), MRI, P: m_MInstr(MI&: MinMaxMI)))
7710 return false;
7711 unsigned NewOpc = getInverseGMinMaxOpcode(MinMaxOpc: MinMaxMI->getOpcode());
7712 if (isLegal(Query: {NewOpc, {DestTy}})) {
7713 MatchInfo = [=](MachineIRBuilder &B) {
7714 B.buildInstr(Opc: NewOpc, DstOps: {DestReg}, SrcOps: {X, Sub0});
7715 };
7716 return true;
7717 }
7718 }
7719
7720 return false;
7721}
7722
7723bool CombinerHelper::matchSelect(MachineInstr &MI, BuildFnTy &MatchInfo) const {
7724 GSelect *Select = cast<GSelect>(Val: &MI);
7725
7726 if (tryFoldSelectOfConstants(Select, MatchInfo))
7727 return true;
7728
7729 if (tryFoldBoolSelectToLogic(Select, MatchInfo))
7730 return true;
7731
7732 return false;
7733}
7734
7735/// Fold (icmp Pred1 V1, C1) && (icmp Pred2 V2, C2)
7736/// or (icmp Pred1 V1, C1) || (icmp Pred2 V2, C2)
7737/// into a single comparison using range-based reasoning.
7738/// see InstCombinerImpl::foldAndOrOfICmpsUsingRanges.
7739bool CombinerHelper::tryFoldAndOrOrICmpsUsingRanges(
7740 GLogicalBinOp *Logic, BuildFnTy &MatchInfo) const {
7741 assert(Logic->getOpcode() != TargetOpcode::G_XOR && "unexpected xor");
7742 bool IsAnd = Logic->getOpcode() == TargetOpcode::G_AND;
7743 Register DstReg = Logic->getReg(Idx: 0);
7744 Register LHS = Logic->getLHSReg();
7745 Register RHS = Logic->getRHSReg();
7746 unsigned Flags = Logic->getFlags();
7747
7748 // We need an G_ICMP on the LHS register.
7749 GICmp *Cmp1 = getOpcodeDef<GICmp>(Reg: LHS, MRI);
7750 if (!Cmp1)
7751 return false;
7752
7753 // We need an G_ICMP on the RHS register.
7754 GICmp *Cmp2 = getOpcodeDef<GICmp>(Reg: RHS, MRI);
7755 if (!Cmp2)
7756 return false;
7757
7758 // We want to fold the icmps.
7759 if (!MRI.hasOneNonDBGUse(RegNo: Cmp1->getReg(Idx: 0)) ||
7760 !MRI.hasOneNonDBGUse(RegNo: Cmp2->getReg(Idx: 0)))
7761 return false;
7762
7763 APInt C1;
7764 APInt C2;
7765 std::optional<ValueAndVReg> MaybeC1 =
7766 getIConstantVRegValWithLookThrough(VReg: Cmp1->getRHSReg(), MRI);
7767 if (!MaybeC1)
7768 return false;
7769 C1 = MaybeC1->Value;
7770
7771 std::optional<ValueAndVReg> MaybeC2 =
7772 getIConstantVRegValWithLookThrough(VReg: Cmp2->getRHSReg(), MRI);
7773 if (!MaybeC2)
7774 return false;
7775 C2 = MaybeC2->Value;
7776
7777 Register R1 = Cmp1->getLHSReg();
7778 Register R2 = Cmp2->getLHSReg();
7779 CmpInst::Predicate Pred1 = Cmp1->getCond();
7780 CmpInst::Predicate Pred2 = Cmp2->getCond();
7781 LLT CmpTy = MRI.getType(Reg: Cmp1->getReg(Idx: 0));
7782 LLT CmpOperandTy = MRI.getType(Reg: R1);
7783
7784 if (CmpOperandTy.isPointer())
7785 return false;
7786
7787 // We build ands, adds, and constants of type CmpOperandTy.
7788 // They must be legal to build.
7789 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_AND, CmpOperandTy}) ||
7790 !isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_ADD, CmpOperandTy}) ||
7791 !isConstantLegalOrBeforeLegalizer(Ty: CmpOperandTy))
7792 return false;
7793
7794 // Look through add of a constant offset on R1, R2, or both operands. This
7795 // allows us to interpret the R + C' < C'' range idiom into a proper range.
7796 std::optional<APInt> Offset1;
7797 std::optional<APInt> Offset2;
7798 if (R1 != R2) {
7799 if (GAdd *Add = getOpcodeDef<GAdd>(Reg: R1, MRI)) {
7800 std::optional<ValueAndVReg> MaybeOffset1 =
7801 getIConstantVRegValWithLookThrough(VReg: Add->getRHSReg(), MRI);
7802 if (MaybeOffset1) {
7803 R1 = Add->getLHSReg();
7804 Offset1 = MaybeOffset1->Value;
7805 }
7806 }
7807 if (GAdd *Add = getOpcodeDef<GAdd>(Reg: R2, MRI)) {
7808 std::optional<ValueAndVReg> MaybeOffset2 =
7809 getIConstantVRegValWithLookThrough(VReg: Add->getRHSReg(), MRI);
7810 if (MaybeOffset2) {
7811 R2 = Add->getLHSReg();
7812 Offset2 = MaybeOffset2->Value;
7813 }
7814 }
7815 }
7816
7817 if (R1 != R2)
7818 return false;
7819
7820 // We calculate the icmp ranges including maybe offsets.
7821 ConstantRange CR1 = ConstantRange::makeExactICmpRegion(
7822 Pred: IsAnd ? ICmpInst::getInversePredicate(pred: Pred1) : Pred1, Other: C1);
7823 if (Offset1)
7824 CR1 = CR1.subtract(CI: *Offset1);
7825
7826 ConstantRange CR2 = ConstantRange::makeExactICmpRegion(
7827 Pred: IsAnd ? ICmpInst::getInversePredicate(pred: Pred2) : Pred2, Other: C2);
7828 if (Offset2)
7829 CR2 = CR2.subtract(CI: *Offset2);
7830
7831 bool CreateMask = false;
7832 APInt LowerDiff;
7833 std::optional<ConstantRange> CR = CR1.exactUnionWith(CR: CR2);
7834 if (!CR) {
7835 // We need non-wrapping ranges.
7836 if (CR1.isWrappedSet() || CR2.isWrappedSet())
7837 return false;
7838
7839 // Check whether we have equal-size ranges that only differ by one bit.
7840 // In that case we can apply a mask to map one range onto the other.
7841 LowerDiff = CR1.getLower() ^ CR2.getLower();
7842 APInt UpperDiff = (CR1.getUpper() - 1) ^ (CR2.getUpper() - 1);
7843 APInt CR1Size = CR1.getUpper() - CR1.getLower();
7844 if (!LowerDiff.isPowerOf2() || LowerDiff != UpperDiff ||
7845 CR1Size != CR2.getUpper() - CR2.getLower())
7846 return false;
7847
7848 CR = CR1.getLower().ult(RHS: CR2.getLower()) ? CR1 : CR2;
7849 CreateMask = true;
7850 }
7851
7852 if (IsAnd)
7853 CR = CR->inverse();
7854
7855 CmpInst::Predicate NewPred;
7856 APInt NewC, Offset;
7857 CR->getEquivalentICmp(Pred&: NewPred, RHS&: NewC, Offset);
7858
7859 // We take the result type of one of the original icmps, CmpTy, for
7860 // the to be build icmp. The operand type, CmpOperandTy, is used for
7861 // the other instructions and constants to be build. The types of
7862 // the parameters and output are the same for add and and. CmpTy
7863 // and the type of DstReg might differ. That is why we zext or trunc
7864 // the icmp into the destination register.
7865
7866 MatchInfo = [=](MachineIRBuilder &B) {
7867 if (CreateMask && Offset != 0) {
7868 auto TildeLowerDiff = B.buildConstant(Res: CmpOperandTy, Val: ~LowerDiff);
7869 auto And = B.buildAnd(Dst: CmpOperandTy, Src0: R1, Src1: TildeLowerDiff); // the mask.
7870 auto OffsetC = B.buildConstant(Res: CmpOperandTy, Val: Offset);
7871 auto Add = B.buildAdd(Dst: CmpOperandTy, Src0: And, Src1: OffsetC, Flags);
7872 auto NewCon = B.buildConstant(Res: CmpOperandTy, Val: NewC);
7873 auto ICmp = B.buildICmp(Pred: NewPred, Res: CmpTy, Op0: Add, Op1: NewCon);
7874 B.buildZExtOrTrunc(Res: DstReg, Op: ICmp);
7875 } else if (CreateMask && Offset == 0) {
7876 auto TildeLowerDiff = B.buildConstant(Res: CmpOperandTy, Val: ~LowerDiff);
7877 auto And = B.buildAnd(Dst: CmpOperandTy, Src0: R1, Src1: TildeLowerDiff); // the mask.
7878 auto NewCon = B.buildConstant(Res: CmpOperandTy, Val: NewC);
7879 auto ICmp = B.buildICmp(Pred: NewPred, Res: CmpTy, Op0: And, Op1: NewCon);
7880 B.buildZExtOrTrunc(Res: DstReg, Op: ICmp);
7881 } else if (!CreateMask && Offset != 0) {
7882 auto OffsetC = B.buildConstant(Res: CmpOperandTy, Val: Offset);
7883 auto Add = B.buildAdd(Dst: CmpOperandTy, Src0: R1, Src1: OffsetC, Flags);
7884 auto NewCon = B.buildConstant(Res: CmpOperandTy, Val: NewC);
7885 auto ICmp = B.buildICmp(Pred: NewPred, Res: CmpTy, Op0: Add, Op1: NewCon);
7886 B.buildZExtOrTrunc(Res: DstReg, Op: ICmp);
7887 } else if (!CreateMask && Offset == 0) {
7888 auto NewCon = B.buildConstant(Res: CmpOperandTy, Val: NewC);
7889 auto ICmp = B.buildICmp(Pred: NewPred, Res: CmpTy, Op0: R1, Op1: NewCon);
7890 B.buildZExtOrTrunc(Res: DstReg, Op: ICmp);
7891 } else {
7892 llvm_unreachable("unexpected configuration of CreateMask and Offset");
7893 }
7894 };
7895 return true;
7896}
7897
7898bool CombinerHelper::tryFoldLogicOfFCmps(GLogicalBinOp *Logic,
7899 BuildFnTy &MatchInfo) const {
7900 assert(Logic->getOpcode() != TargetOpcode::G_XOR && "unexpecte xor");
7901 Register DestReg = Logic->getReg(Idx: 0);
7902 Register LHS = Logic->getLHSReg();
7903 Register RHS = Logic->getRHSReg();
7904 bool IsAnd = Logic->getOpcode() == TargetOpcode::G_AND;
7905
7906 // We need a compare on the LHS register.
7907 GFCmp *Cmp1 = getOpcodeDef<GFCmp>(Reg: LHS, MRI);
7908 if (!Cmp1)
7909 return false;
7910
7911 // We need a compare on the RHS register.
7912 GFCmp *Cmp2 = getOpcodeDef<GFCmp>(Reg: RHS, MRI);
7913 if (!Cmp2)
7914 return false;
7915
7916 LLT CmpTy = MRI.getType(Reg: Cmp1->getReg(Idx: 0));
7917 LLT CmpOperandTy = MRI.getType(Reg: Cmp1->getLHSReg());
7918
7919 // We build one fcmp, want to fold the fcmps, replace the logic op,
7920 // and the fcmps must have the same shape.
7921 if (!isLegalOrBeforeLegalizer(
7922 Query: {TargetOpcode::G_FCMP, {CmpTy, CmpOperandTy}}) ||
7923 !MRI.hasOneNonDBGUse(RegNo: Logic->getReg(Idx: 0)) ||
7924 !MRI.hasOneNonDBGUse(RegNo: Cmp1->getReg(Idx: 0)) ||
7925 !MRI.hasOneNonDBGUse(RegNo: Cmp2->getReg(Idx: 0)) ||
7926 MRI.getType(Reg: Cmp1->getLHSReg()) != MRI.getType(Reg: Cmp2->getLHSReg()))
7927 return false;
7928
7929 CmpInst::Predicate PredL = Cmp1->getCond();
7930 CmpInst::Predicate PredR = Cmp2->getCond();
7931 Register LHS0 = Cmp1->getLHSReg();
7932 Register LHS1 = Cmp1->getRHSReg();
7933 Register RHS0 = Cmp2->getLHSReg();
7934 Register RHS1 = Cmp2->getRHSReg();
7935
7936 if (LHS0 == RHS1 && LHS1 == RHS0) {
7937 // Swap RHS operands to match LHS.
7938 PredR = CmpInst::getSwappedPredicate(pred: PredR);
7939 std::swap(a&: RHS0, b&: RHS1);
7940 }
7941
7942 if (LHS0 == RHS0 && LHS1 == RHS1) {
7943 // We determine the new predicate.
7944 unsigned CmpCodeL = getFCmpCode(CC: PredL);
7945 unsigned CmpCodeR = getFCmpCode(CC: PredR);
7946 unsigned NewPred = IsAnd ? CmpCodeL & CmpCodeR : CmpCodeL | CmpCodeR;
7947 unsigned Flags = Cmp1->getFlags() | Cmp2->getFlags();
7948 MatchInfo = [=](MachineIRBuilder &B) {
7949 // The fcmp predicates fill the lower part of the enum.
7950 FCmpInst::Predicate Pred = static_cast<FCmpInst::Predicate>(NewPred);
7951 if (Pred == FCmpInst::FCMP_FALSE &&
7952 isConstantLegalOrBeforeLegalizer(Ty: CmpTy)) {
7953 auto False = B.buildConstant(Res: CmpTy, Val: 0);
7954 B.buildZExtOrTrunc(Res: DestReg, Op: False);
7955 } else if (Pred == FCmpInst::FCMP_TRUE &&
7956 isConstantLegalOrBeforeLegalizer(Ty: CmpTy)) {
7957 auto True =
7958 B.buildConstant(Res: CmpTy, Val: getICmpTrueVal(TLI: getTargetLowering(),
7959 IsVector: CmpTy.isVector() /*isVector*/,
7960 IsFP: true /*isFP*/));
7961 B.buildZExtOrTrunc(Res: DestReg, Op: True);
7962 } else { // We take the predicate without predicate optimizations.
7963 auto Cmp = B.buildFCmp(Pred, Res: CmpTy, Op0: LHS0, Op1: LHS1, Flags);
7964 B.buildZExtOrTrunc(Res: DestReg, Op: Cmp);
7965 }
7966 };
7967 return true;
7968 }
7969
7970 return false;
7971}
7972
7973bool CombinerHelper::matchAnd(MachineInstr &MI, BuildFnTy &MatchInfo) const {
7974 GAnd *And = cast<GAnd>(Val: &MI);
7975
7976 if (tryFoldAndOrOrICmpsUsingRanges(Logic: And, MatchInfo))
7977 return true;
7978
7979 if (tryFoldLogicOfFCmps(Logic: And, MatchInfo))
7980 return true;
7981
7982 return false;
7983}
7984
7985bool CombinerHelper::matchOr(MachineInstr &MI, BuildFnTy &MatchInfo) const {
7986 GOr *Or = cast<GOr>(Val: &MI);
7987
7988 if (tryFoldAndOrOrICmpsUsingRanges(Logic: Or, MatchInfo))
7989 return true;
7990
7991 if (tryFoldLogicOfFCmps(Logic: Or, MatchInfo))
7992 return true;
7993
7994 return false;
7995}
7996
7997bool CombinerHelper::matchAddOverflow(MachineInstr &MI,
7998 BuildFnTy &MatchInfo) const {
7999 GAddCarryOut *Add = cast<GAddCarryOut>(Val: &MI);
8000
8001 // Addo has no flags
8002 Register Dst = Add->getReg(Idx: 0);
8003 Register Carry = Add->getReg(Idx: 1);
8004 Register LHS = Add->getLHSReg();
8005 Register RHS = Add->getRHSReg();
8006 bool IsSigned = Add->isSigned();
8007 LLT DstTy = MRI.getType(Reg: Dst);
8008 LLT CarryTy = MRI.getType(Reg: Carry);
8009
8010 // Fold addo, if the carry is dead -> add, undef.
8011 if (MRI.use_nodbg_empty(RegNo: Carry) &&
8012 isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_ADD, {DstTy}})) {
8013 MatchInfo = [=](MachineIRBuilder &B) {
8014 B.buildAdd(Dst, Src0: LHS, Src1: RHS);
8015 B.buildUndef(Res: Carry);
8016 };
8017 return true;
8018 }
8019
8020 // Canonicalize constant to RHS.
8021 if (isConstantOrConstantVectorI(Src: LHS) && !isConstantOrConstantVectorI(Src: RHS)) {
8022 if (IsSigned) {
8023 MatchInfo = [=](MachineIRBuilder &B) {
8024 B.buildSAddo(Res: Dst, CarryOut: Carry, Op0: RHS, Op1: LHS);
8025 };
8026 return true;
8027 }
8028 // !IsSigned
8029 MatchInfo = [=](MachineIRBuilder &B) {
8030 B.buildUAddo(Res: Dst, CarryOut: Carry, Op0: RHS, Op1: LHS);
8031 };
8032 return true;
8033 }
8034
8035 std::optional<APInt> MaybeLHS = getConstantOrConstantSplatVector(Src: LHS);
8036 std::optional<APInt> MaybeRHS = getConstantOrConstantSplatVector(Src: RHS);
8037
8038 // Fold addo(c1, c2) -> c3, carry.
8039 if (MaybeLHS && MaybeRHS && isConstantLegalOrBeforeLegalizer(Ty: DstTy) &&
8040 isConstantLegalOrBeforeLegalizer(Ty: CarryTy)) {
8041 bool Overflow;
8042 APInt Result = IsSigned ? MaybeLHS->sadd_ov(RHS: *MaybeRHS, Overflow)
8043 : MaybeLHS->uadd_ov(RHS: *MaybeRHS, Overflow);
8044 MatchInfo = [=](MachineIRBuilder &B) {
8045 B.buildConstant(Res: Dst, Val: Result);
8046 B.buildConstant(Res: Carry, Val: Overflow);
8047 };
8048 return true;
8049 }
8050
8051 // Fold (addo x, 0) -> x, no carry
8052 if (MaybeRHS && *MaybeRHS == 0 && isConstantLegalOrBeforeLegalizer(Ty: CarryTy)) {
8053 MatchInfo = [=](MachineIRBuilder &B) {
8054 B.buildCopy(Res: Dst, Op: LHS);
8055 B.buildConstant(Res: Carry, Val: 0);
8056 };
8057 return true;
8058 }
8059
8060 // Given 2 constant operands whose sum does not overflow:
8061 // uaddo (X +nuw C0), C1 -> uaddo X, C0 + C1
8062 // saddo (X +nsw C0), C1 -> saddo X, C0 + C1
8063 GAdd *AddLHS = getOpcodeDef<GAdd>(Reg: LHS, MRI);
8064 if (MaybeRHS && AddLHS && MRI.hasOneNonDBGUse(RegNo: Add->getReg(Idx: 0)) &&
8065 ((IsSigned && AddLHS->getFlag(Flag: MachineInstr::MIFlag::NoSWrap)) ||
8066 (!IsSigned && AddLHS->getFlag(Flag: MachineInstr::MIFlag::NoUWrap)))) {
8067 std::optional<APInt> MaybeAddRHS =
8068 getConstantOrConstantSplatVector(Src: AddLHS->getRHSReg());
8069 if (MaybeAddRHS) {
8070 bool Overflow;
8071 APInt NewC = IsSigned ? MaybeAddRHS->sadd_ov(RHS: *MaybeRHS, Overflow)
8072 : MaybeAddRHS->uadd_ov(RHS: *MaybeRHS, Overflow);
8073 if (!Overflow && isConstantLegalOrBeforeLegalizer(Ty: DstTy)) {
8074 if (IsSigned) {
8075 MatchInfo = [=](MachineIRBuilder &B) {
8076 auto ConstRHS = B.buildConstant(Res: DstTy, Val: NewC);
8077 B.buildSAddo(Res: Dst, CarryOut: Carry, Op0: AddLHS->getLHSReg(), Op1: ConstRHS);
8078 };
8079 return true;
8080 }
8081 // !IsSigned
8082 MatchInfo = [=](MachineIRBuilder &B) {
8083 auto ConstRHS = B.buildConstant(Res: DstTy, Val: NewC);
8084 B.buildUAddo(Res: Dst, CarryOut: Carry, Op0: AddLHS->getLHSReg(), Op1: ConstRHS);
8085 };
8086 return true;
8087 }
8088 }
8089 };
8090
8091 // We try to combine addo to non-overflowing add.
8092 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_ADD, {DstTy}}) ||
8093 !isConstantLegalOrBeforeLegalizer(Ty: CarryTy))
8094 return false;
8095
8096 // We try to combine uaddo to non-overflowing add.
8097 if (!IsSigned) {
8098 ConstantRange CRLHS =
8099 ConstantRange::fromKnownBits(Known: VT->getKnownBits(R: LHS), /*IsSigned=*/false);
8100 ConstantRange CRRHS =
8101 ConstantRange::fromKnownBits(Known: VT->getKnownBits(R: RHS), /*IsSigned=*/false);
8102
8103 switch (CRLHS.unsignedAddMayOverflow(Other: CRRHS)) {
8104 case ConstantRange::OverflowResult::MayOverflow:
8105 return false;
8106 case ConstantRange::OverflowResult::NeverOverflows: {
8107 MatchInfo = [=](MachineIRBuilder &B) {
8108 B.buildAdd(Dst, Src0: LHS, Src1: RHS, Flags: MachineInstr::MIFlag::NoUWrap);
8109 B.buildConstant(Res: Carry, Val: 0);
8110 };
8111 return true;
8112 }
8113 case ConstantRange::OverflowResult::AlwaysOverflowsLow:
8114 case ConstantRange::OverflowResult::AlwaysOverflowsHigh: {
8115 MatchInfo = [=](MachineIRBuilder &B) {
8116 B.buildAdd(Dst, Src0: LHS, Src1: RHS);
8117 B.buildConstant(Res: Carry, Val: 1);
8118 };
8119 return true;
8120 }
8121 }
8122 return false;
8123 }
8124
8125 // We try to combine saddo to non-overflowing add.
8126
8127 // If LHS and RHS each have at least two sign bits, then there is no signed
8128 // overflow.
8129 if (VT->computeNumSignBits(R: RHS) > 1 && VT->computeNumSignBits(R: LHS) > 1) {
8130 MatchInfo = [=](MachineIRBuilder &B) {
8131 B.buildAdd(Dst, Src0: LHS, Src1: RHS, Flags: MachineInstr::MIFlag::NoSWrap);
8132 B.buildConstant(Res: Carry, Val: 0);
8133 };
8134 return true;
8135 }
8136
8137 ConstantRange CRLHS =
8138 ConstantRange::fromKnownBits(Known: VT->getKnownBits(R: LHS), /*IsSigned=*/true);
8139 ConstantRange CRRHS =
8140 ConstantRange::fromKnownBits(Known: VT->getKnownBits(R: RHS), /*IsSigned=*/true);
8141
8142 switch (CRLHS.signedAddMayOverflow(Other: CRRHS)) {
8143 case ConstantRange::OverflowResult::MayOverflow:
8144 return false;
8145 case ConstantRange::OverflowResult::NeverOverflows: {
8146 MatchInfo = [=](MachineIRBuilder &B) {
8147 B.buildAdd(Dst, Src0: LHS, Src1: RHS, Flags: MachineInstr::MIFlag::NoSWrap);
8148 B.buildConstant(Res: Carry, Val: 0);
8149 };
8150 return true;
8151 }
8152 case ConstantRange::OverflowResult::AlwaysOverflowsLow:
8153 case ConstantRange::OverflowResult::AlwaysOverflowsHigh: {
8154 MatchInfo = [=](MachineIRBuilder &B) {
8155 B.buildAdd(Dst, Src0: LHS, Src1: RHS);
8156 B.buildConstant(Res: Carry, Val: 1);
8157 };
8158 return true;
8159 }
8160 }
8161
8162 return false;
8163}
8164
8165void CombinerHelper::applyBuildFnMO(const MachineOperand &MO,
8166 BuildFnTy &MatchInfo) const {
8167 MachineInstr *Root = getDefIgnoringCopies(Reg: MO.getReg(), MRI);
8168 MatchInfo(Builder);
8169 Root->eraseFromParent();
8170}
8171
8172bool CombinerHelper::matchFPowIExpansion(MachineInstr &MI,
8173 int64_t Exponent) const {
8174 bool OptForSize = MI.getMF()->getFunction().hasOptSize();
8175 return getTargetLowering().isBeneficialToExpandPowI(Exponent, OptForSize);
8176}
8177
8178void CombinerHelper::applyExpandFPowI(MachineInstr &MI,
8179 int64_t Exponent) const {
8180 auto [Dst, Base] = MI.getFirst2Regs();
8181 LLT Ty = MRI.getType(Reg: Dst);
8182 int64_t ExpVal = Exponent;
8183
8184 if (ExpVal == 0) {
8185 Builder.buildFConstant(Res: Dst, Val: 1.0);
8186 MI.removeFromParent();
8187 return;
8188 }
8189
8190 if (ExpVal < 0)
8191 ExpVal = -ExpVal;
8192
8193 // We use the simple binary decomposition method from SelectionDAG ExpandPowI
8194 // to generate the multiply sequence. There are more optimal ways to do this
8195 // (for example, powi(x,15) generates one more multiply than it should), but
8196 // this has the benefit of being both really simple and much better than a
8197 // libcall.
8198 std::optional<SrcOp> Res;
8199 SrcOp CurSquare = Base;
8200 while (ExpVal > 0) {
8201 if (ExpVal & 1) {
8202 if (!Res)
8203 Res = CurSquare;
8204 else
8205 Res = Builder.buildFMul(Dst: Ty, Src0: *Res, Src1: CurSquare);
8206 }
8207
8208 CurSquare = Builder.buildFMul(Dst: Ty, Src0: CurSquare, Src1: CurSquare);
8209 ExpVal >>= 1;
8210 }
8211
8212 // If the original exponent was negative, invert the result, producing
8213 // 1/(x*x*x).
8214 if (Exponent < 0)
8215 Res = Builder.buildFDiv(Dst: Ty, Src0: Builder.buildFConstant(Res: Ty, Val: 1.0), Src1: *Res,
8216 Flags: MI.getFlags());
8217
8218 Builder.buildCopy(Res: Dst, Op: *Res);
8219 MI.eraseFromParent();
8220}
8221
8222bool CombinerHelper::matchFoldAPlusC1MinusC2(const MachineInstr &MI,
8223 BuildFnTy &MatchInfo) const {
8224 // fold (A+C1)-C2 -> A+(C1-C2)
8225 const GSub *Sub = cast<GSub>(Val: &MI);
8226 Register A, C1Reg;
8227 if (!mi_match(R: Sub->getLHSReg(), MRI, P: m_GAdd(L: m_Reg(R&: A), R: m_Reg(R&: C1Reg))))
8228 return false;
8229
8230 if (!MRI.hasOneNonDBGUse(RegNo: Sub->getLHSReg()))
8231 return false;
8232
8233 APInt C2 = getIConstantFromReg(VReg: Sub->getRHSReg(), MRI);
8234 APInt C1 = getIConstantFromReg(VReg: C1Reg, MRI);
8235
8236 Register Dst = Sub->getReg(Idx: 0);
8237 LLT DstTy = MRI.getType(Reg: Dst);
8238
8239 MatchInfo = [=](MachineIRBuilder &B) {
8240 auto Const = B.buildConstant(Res: DstTy, Val: C1 - C2);
8241 B.buildAdd(Dst, Src0: A, Src1: Const);
8242 };
8243
8244 return true;
8245}
8246
8247bool CombinerHelper::matchFoldC2MinusAPlusC1(const MachineInstr &MI,
8248 BuildFnTy &MatchInfo) const {
8249 // fold C2-(A+C1) -> (C2-C1)-A
8250 const GSub *Sub = cast<GSub>(Val: &MI);
8251 Register A, C1Reg;
8252 if (!mi_match(R: Sub->getRHSReg(), MRI, P: m_GAdd(L: m_Reg(R&: A), R: m_Reg(R&: C1Reg))))
8253 return false;
8254
8255 if (!MRI.hasOneNonDBGUse(RegNo: Sub->getRHSReg()))
8256 return false;
8257
8258 APInt C2 = getIConstantFromReg(VReg: Sub->getLHSReg(), MRI);
8259 APInt C1 = getIConstantFromReg(VReg: C1Reg, MRI);
8260
8261 Register Dst = Sub->getReg(Idx: 0);
8262 LLT DstTy = MRI.getType(Reg: Dst);
8263
8264 MatchInfo = [=](MachineIRBuilder &B) {
8265 auto Const = B.buildConstant(Res: DstTy, Val: C2 - C1);
8266 B.buildSub(Dst, Src0: Const, Src1: A);
8267 };
8268
8269 return true;
8270}
8271
8272bool CombinerHelper::matchFoldAMinusC1MinusC2(const MachineInstr &MI,
8273 BuildFnTy &MatchInfo) const {
8274 // fold (A-C1)-C2 -> A-(C1+C2)
8275 const GSub *Sub1 = cast<GSub>(Val: &MI);
8276 Register A, C1Reg;
8277 if (!mi_match(R: Sub1->getLHSReg(), MRI, P: m_GSub(L: m_Reg(R&: A), R: m_Reg(R&: C1Reg))))
8278 return false;
8279
8280 if (!MRI.hasOneNonDBGUse(RegNo: Sub1->getLHSReg()))
8281 return false;
8282
8283 APInt C2 = getIConstantFromReg(VReg: Sub1->getRHSReg(), MRI);
8284 APInt C1 = getIConstantFromReg(VReg: C1Reg, MRI);
8285
8286 Register Dst = Sub1->getReg(Idx: 0);
8287 LLT DstTy = MRI.getType(Reg: Dst);
8288
8289 MatchInfo = [=](MachineIRBuilder &B) {
8290 auto Const = B.buildConstant(Res: DstTy, Val: C1 + C2);
8291 B.buildSub(Dst, Src0: A, Src1: Const);
8292 };
8293
8294 return true;
8295}
8296
8297bool CombinerHelper::matchFoldC1Minus2MinusC2(const MachineInstr &MI,
8298 BuildFnTy &MatchInfo) const {
8299 // fold (C1-A)-C2 -> (C1-C2)-A
8300 const GSub *Sub1 = cast<GSub>(Val: &MI);
8301 Register C1Reg, A;
8302 if (!mi_match(R: Sub1->getLHSReg(), MRI, P: m_GSub(L: m_Reg(R&: C1Reg), R: m_Reg(R&: A))))
8303 return false;
8304
8305 if (!MRI.hasOneNonDBGUse(RegNo: Sub1->getLHSReg()))
8306 return false;
8307
8308 APInt C2 = getIConstantFromReg(VReg: Sub1->getRHSReg(), MRI);
8309 APInt C1 = getIConstantFromReg(VReg: C1Reg, MRI);
8310
8311 Register Dst = Sub1->getReg(Idx: 0);
8312 LLT DstTy = MRI.getType(Reg: Dst);
8313
8314 MatchInfo = [=](MachineIRBuilder &B) {
8315 auto Const = B.buildConstant(Res: DstTy, Val: C1 - C2);
8316 B.buildSub(Dst, Src0: Const, Src1: A);
8317 };
8318
8319 return true;
8320}
8321
8322bool CombinerHelper::matchFoldAMinusC1PlusC2(const MachineInstr &MI,
8323 BuildFnTy &MatchInfo) const {
8324 // fold ((A-C1)+C2) -> (A+(C2-C1))
8325 const GAdd *Add = cast<GAdd>(Val: &MI);
8326 Register A, C1Reg;
8327 if (!mi_match(R: Add->getLHSReg(), MRI, P: m_GSub(L: m_Reg(R&: A), R: m_Reg(R&: C1Reg))))
8328 return false;
8329
8330 if (!MRI.hasOneNonDBGUse(RegNo: Add->getLHSReg()))
8331 return false;
8332
8333 APInt C2 = getIConstantFromReg(VReg: Add->getRHSReg(), MRI);
8334 APInt C1 = getIConstantFromReg(VReg: C1Reg, MRI);
8335
8336 Register Dst = Add->getReg(Idx: 0);
8337 LLT DstTy = MRI.getType(Reg: Dst);
8338
8339 MatchInfo = [=](MachineIRBuilder &B) {
8340 auto Const = B.buildConstant(Res: DstTy, Val: C2 - C1);
8341 B.buildAdd(Dst, Src0: A, Src1: Const);
8342 };
8343
8344 return true;
8345}
8346
8347bool CombinerHelper::matchUnmergeValuesAnyExtBuildVector(
8348 const MachineInstr &MI, BuildFnTy &MatchInfo) const {
8349 const GUnmerge *Unmerge = cast<GUnmerge>(Val: &MI);
8350
8351 if (!MRI.hasOneNonDBGUse(RegNo: Unmerge->getSourceReg()))
8352 return false;
8353
8354 LLT DstTy = MRI.getType(Reg: Unmerge->getReg(Idx: 0));
8355
8356 // $bv:_(<8 x s8>) = G_BUILD_VECTOR ....
8357 // $any:_(<8 x s16>) = G_ANYEXT $bv
8358 // $uv:_(<4 x s16>), $uv1:_(<4 x s16>) = G_UNMERGE_VALUES $any
8359 //
8360 // ->
8361 //
8362 // $any:_(s16) = G_ANYEXT $bv[0]
8363 // $any1:_(s16) = G_ANYEXT $bv[1]
8364 // $any2:_(s16) = G_ANYEXT $bv[2]
8365 // $any3:_(s16) = G_ANYEXT $bv[3]
8366 // $any4:_(s16) = G_ANYEXT $bv[4]
8367 // $any5:_(s16) = G_ANYEXT $bv[5]
8368 // $any6:_(s16) = G_ANYEXT $bv[6]
8369 // $any7:_(s16) = G_ANYEXT $bv[7]
8370 // $uv:_(<4 x s16>) = G_BUILD_VECTOR $any, $any1, $any2, $any3
8371 // $uv1:_(<4 x s16>) = G_BUILD_VECTOR $any4, $any5, $any6, $any7
8372
8373 // We want to unmerge into vectors.
8374 if (!DstTy.isFixedVector())
8375 return false;
8376
8377 Register AnySrcReg;
8378 if (!mi_match(R: Unmerge->getSourceReg(), MRI, P: m_GAnyExt(Src: m_Reg(R&: AnySrcReg))))
8379 return false;
8380
8381 GBuildVector *BV;
8382 if (mi_match(R: AnySrcReg, MRI, P: m_GBuildVector(Inst&: BV))) {
8383 // G_UNMERGE_VALUES G_ANYEXT G_BUILD_VECTOR
8384
8385 if (!MRI.hasOneNonDBGUse(RegNo: BV->getReg(Idx: 0)))
8386 return false;
8387
8388 // FIXME: check element types?
8389 if (BV->getNumSources() % Unmerge->getNumDefs() != 0)
8390 return false;
8391
8392 LLT BigBvTy = MRI.getType(Reg: BV->getReg(Idx: 0));
8393 LLT SmallBvTy = DstTy;
8394 LLT SmallBvElemenTy = SmallBvTy.getElementType();
8395
8396 if (!isLegalOrBeforeLegalizer(
8397 Query: {TargetOpcode::G_BUILD_VECTOR, {SmallBvTy, SmallBvElemenTy}}))
8398 return false;
8399
8400 // We check the legality of scalar anyext.
8401 if (!isLegalOrBeforeLegalizer(
8402 Query: {TargetOpcode::G_ANYEXT,
8403 {SmallBvElemenTy, BigBvTy.getElementType()}}))
8404 return false;
8405
8406 MatchInfo = [=](MachineIRBuilder &B) {
8407 // Build into each G_UNMERGE_VALUES def
8408 // a small build vector with anyext from the source build vector.
8409 for (unsigned I = 0; I < Unmerge->getNumDefs(); ++I) {
8410 SmallVector<Register> Ops;
8411 for (unsigned J = 0; J < SmallBvTy.getNumElements(); ++J) {
8412 Register SourceArray =
8413 BV->getSourceReg(I: I * SmallBvTy.getNumElements() + J);
8414 auto AnyExt = B.buildAnyExt(Res: SmallBvElemenTy, Op: SourceArray);
8415 Ops.push_back(Elt: AnyExt.getReg(Idx: 0));
8416 }
8417 B.buildBuildVector(Res: Unmerge->getOperand(i: I).getReg(), Ops);
8418 };
8419 };
8420 return true;
8421 };
8422
8423 return false;
8424}
8425
8426bool CombinerHelper::matchShuffleUndefRHS(MachineInstr &MI,
8427 BuildFnTy &MatchInfo) const {
8428
8429 bool Changed = false;
8430 auto &Shuffle = cast<GShuffleVector>(Val&: MI);
8431 ArrayRef<int> OrigMask = Shuffle.getMask();
8432 SmallVector<int, 16> NewMask;
8433 const LLT SrcTy = MRI.getType(Reg: Shuffle.getSrc1Reg());
8434 const unsigned NumSrcElems = SrcTy.isVector() ? SrcTy.getNumElements() : 1;
8435 const unsigned NumDstElts = OrigMask.size();
8436 for (unsigned i = 0; i != NumDstElts; ++i) {
8437 int Idx = OrigMask[i];
8438 if (Idx >= (int)NumSrcElems) {
8439 Idx = -1;
8440 Changed = true;
8441 }
8442 NewMask.push_back(Elt: Idx);
8443 }
8444
8445 if (!Changed)
8446 return false;
8447
8448 MatchInfo = [&, NewMask = std::move(NewMask)](MachineIRBuilder &B) {
8449 B.buildShuffleVector(Res: MI.getOperand(i: 0), Src1: MI.getOperand(i: 1), Src2: MI.getOperand(i: 2),
8450 Mask: std::move(NewMask));
8451 };
8452
8453 return true;
8454}
8455
8456static void commuteMask(MutableArrayRef<int> Mask, const unsigned NumElems) {
8457 const unsigned MaskSize = Mask.size();
8458 for (unsigned I = 0; I < MaskSize; ++I) {
8459 int Idx = Mask[I];
8460 if (Idx < 0)
8461 continue;
8462
8463 if (Idx < (int)NumElems)
8464 Mask[I] = Idx + NumElems;
8465 else
8466 Mask[I] = Idx - NumElems;
8467 }
8468}
8469
8470bool CombinerHelper::matchShuffleDisjointMask(MachineInstr &MI,
8471 BuildFnTy &MatchInfo) const {
8472
8473 auto &Shuffle = cast<GShuffleVector>(Val&: MI);
8474 // If any of the two inputs is already undef, don't check the mask again to
8475 // prevent infinite loop
8476 if (getOpcodeDef(Opcode: TargetOpcode::G_IMPLICIT_DEF, Reg: Shuffle.getSrc1Reg(), MRI))
8477 return false;
8478
8479 if (getOpcodeDef(Opcode: TargetOpcode::G_IMPLICIT_DEF, Reg: Shuffle.getSrc2Reg(), MRI))
8480 return false;
8481
8482 const LLT DstTy = MRI.getType(Reg: Shuffle.getReg(Idx: 0));
8483 const LLT Src1Ty = MRI.getType(Reg: Shuffle.getSrc1Reg());
8484 if (!isLegalOrBeforeLegalizer(
8485 Query: {TargetOpcode::G_SHUFFLE_VECTOR, {DstTy, Src1Ty}}))
8486 return false;
8487
8488 ArrayRef<int> Mask = Shuffle.getMask();
8489 const unsigned NumSrcElems = Src1Ty.getNumElements();
8490
8491 bool TouchesSrc1 = false;
8492 bool TouchesSrc2 = false;
8493 const unsigned NumElems = Mask.size();
8494 for (unsigned Idx = 0; Idx < NumElems; ++Idx) {
8495 if (Mask[Idx] < 0)
8496 continue;
8497
8498 if (Mask[Idx] < (int)NumSrcElems)
8499 TouchesSrc1 = true;
8500 else
8501 TouchesSrc2 = true;
8502 }
8503
8504 if (TouchesSrc1 == TouchesSrc2)
8505 return false;
8506
8507 Register NewSrc1 = Shuffle.getSrc1Reg();
8508 SmallVector<int, 16> NewMask(Mask);
8509 if (TouchesSrc2) {
8510 NewSrc1 = Shuffle.getSrc2Reg();
8511 commuteMask(Mask: NewMask, NumElems: NumSrcElems);
8512 }
8513
8514 MatchInfo = [=, &Shuffle](MachineIRBuilder &B) {
8515 auto Undef = B.buildUndef(Res: Src1Ty);
8516 B.buildShuffleVector(Res: Shuffle.getReg(Idx: 0), Src1: NewSrc1, Src2: Undef, Mask: NewMask);
8517 };
8518
8519 return true;
8520}
8521
8522bool CombinerHelper::matchSuboCarryOut(const MachineInstr &MI,
8523 BuildFnTy &MatchInfo) const {
8524 const GSubCarryOut *Subo = cast<GSubCarryOut>(Val: &MI);
8525
8526 Register Dst = Subo->getReg(Idx: 0);
8527 Register LHS = Subo->getLHSReg();
8528 Register RHS = Subo->getRHSReg();
8529 Register Carry = Subo->getCarryOutReg();
8530 LLT DstTy = MRI.getType(Reg: Dst);
8531 LLT CarryTy = MRI.getType(Reg: Carry);
8532
8533 // Check legality before known bits.
8534 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_SUB, {DstTy}}) ||
8535 !isConstantLegalOrBeforeLegalizer(Ty: CarryTy))
8536 return false;
8537
8538 ConstantRange KBLHS =
8539 ConstantRange::fromKnownBits(Known: VT->getKnownBits(R: LHS),
8540 /* IsSigned=*/Subo->isSigned());
8541 ConstantRange KBRHS =
8542 ConstantRange::fromKnownBits(Known: VT->getKnownBits(R: RHS),
8543 /* IsSigned=*/Subo->isSigned());
8544
8545 if (Subo->isSigned()) {
8546 // G_SSUBO
8547 switch (KBLHS.signedSubMayOverflow(Other: KBRHS)) {
8548 case ConstantRange::OverflowResult::MayOverflow:
8549 return false;
8550 case ConstantRange::OverflowResult::NeverOverflows: {
8551 MatchInfo = [=](MachineIRBuilder &B) {
8552 B.buildSub(Dst, Src0: LHS, Src1: RHS, Flags: MachineInstr::MIFlag::NoSWrap);
8553 B.buildConstant(Res: Carry, Val: 0);
8554 };
8555 return true;
8556 }
8557 case ConstantRange::OverflowResult::AlwaysOverflowsLow:
8558 case ConstantRange::OverflowResult::AlwaysOverflowsHigh: {
8559 MatchInfo = [=](MachineIRBuilder &B) {
8560 B.buildSub(Dst, Src0: LHS, Src1: RHS);
8561 B.buildConstant(Res: Carry, Val: getICmpTrueVal(TLI: getTargetLowering(),
8562 /*isVector=*/IsVector: CarryTy.isVector(),
8563 /*isFP=*/IsFP: false));
8564 };
8565 return true;
8566 }
8567 }
8568 return false;
8569 }
8570
8571 // G_USUBO
8572 switch (KBLHS.unsignedSubMayOverflow(Other: KBRHS)) {
8573 case ConstantRange::OverflowResult::MayOverflow:
8574 return false;
8575 case ConstantRange::OverflowResult::NeverOverflows: {
8576 MatchInfo = [=](MachineIRBuilder &B) {
8577 B.buildSub(Dst, Src0: LHS, Src1: RHS, Flags: MachineInstr::MIFlag::NoUWrap);
8578 B.buildConstant(Res: Carry, Val: 0);
8579 };
8580 return true;
8581 }
8582 case ConstantRange::OverflowResult::AlwaysOverflowsLow:
8583 case ConstantRange::OverflowResult::AlwaysOverflowsHigh: {
8584 MatchInfo = [=](MachineIRBuilder &B) {
8585 B.buildSub(Dst, Src0: LHS, Src1: RHS);
8586 B.buildConstant(Res: Carry, Val: getICmpTrueVal(TLI: getTargetLowering(),
8587 /*isVector=*/IsVector: CarryTy.isVector(),
8588 /*isFP=*/IsFP: false));
8589 };
8590 return true;
8591 }
8592 }
8593
8594 return false;
8595}
8596
8597// Fold (ctlz (xor x, (sra x, bitwidth-1))) -> (add (ctls x), 1).
8598// Fold (ctlz (or (shl (xor x, (sra x, bitwidth-1)), 1), 1) -> (ctls x)
8599bool CombinerHelper::matchCtls(MachineInstr &CtlzMI,
8600 BuildFnTy &MatchInfo) const {
8601 assert((CtlzMI.getOpcode() == TargetOpcode::G_CTLZ ||
8602 CtlzMI.getOpcode() == TargetOpcode::G_CTLZ_ZERO_POISON) &&
8603 "Expected G_CTLZ variant");
8604
8605 const Register Dst = CtlzMI.getOperand(i: 0).getReg();
8606 Register Src = CtlzMI.getOperand(i: 1).getReg();
8607
8608 LLT Ty = MRI.getType(Reg: Dst);
8609 LLT SrcTy = MRI.getType(Reg: Src);
8610
8611 if (!(Ty.isValid() && Ty.isScalar()))
8612 return false;
8613
8614 if (!LI)
8615 return false;
8616
8617 SmallVector<LLT, 2> QueryTypes = {Ty, SrcTy};
8618 LegalityQuery Query(TargetOpcode::G_CTLS, QueryTypes);
8619
8620 switch (LI->getAction(Query).Action) {
8621 default:
8622 return false;
8623 case LegalizeActions::Legal:
8624 case LegalizeActions::Custom:
8625 case LegalizeActions::WidenScalar:
8626 break;
8627 }
8628
8629 // Src = or(shl(V, 1), 1) -> Src=V; NeedAdd = False
8630 Register V;
8631 bool NeedAdd = true;
8632 if (mi_match(R: Src, MRI,
8633 P: m_OneUse(SP: m_GOr(L: m_OneUse(SP: m_GShl(L: m_Reg(R&: V), R: m_SpecificICst(RequestedValue: 1))),
8634 R: m_SpecificICst(RequestedValue: 1))))) {
8635 NeedAdd = false;
8636 Src = V;
8637 }
8638
8639 unsigned BitWidth = Ty.getScalarSizeInBits();
8640
8641 Register X;
8642 if (!mi_match(R: Src, MRI,
8643 P: m_OneUse(SP: m_GXor(L: m_Reg(R&: X), R: m_OneUse(SP: m_GAShr(
8644 L: m_DeferredReg(R&: X),
8645 R: m_SpecificICst(RequestedValue: BitWidth - 1)))))))
8646 return false;
8647
8648 MatchInfo = [=](MachineIRBuilder &B) {
8649 if (!NeedAdd) {
8650 B.buildCTLS(Dst, Src0: X);
8651 return;
8652 }
8653
8654 auto Ctls = B.buildCTLS(Dst: Ty, Src0: X);
8655 auto One = B.buildConstant(Res: Ty, Val: 1);
8656
8657 B.buildAdd(Dst, Src0: Ctls, Src1: One);
8658 };
8659
8660 return true;
8661}
8662
8663// Fold shr ( add ( ext X, ext Y ), 1 ) -> avgfloor ( x, y )
8664// Fold shr ( add ( ext X, ext Y, 1 ), 1 ) -> avgceil ( x, y )
8665bool CombinerHelper::matchAVG(MachineInstr &MI, MachineRegisterInfo &MRI,
8666 Register X, Register Y,
8667 unsigned TargetOpc) const {
8668 assert((MI.getOpcode() == TargetOpcode::G_LSHR ||
8669 MI.getOpcode() == TargetOpcode::G_ASHR) &&
8670 "Expected G_LSHR/G_ASHR");
8671
8672 LLT XTy = MRI.getType(Reg: X);
8673 return XTy == MRI.getType(Reg: Y) && isLegal(Query: {TargetOpc, {XTy}});
8674}
8675
8676static unsigned getCountZeroPoisonOpcode(const MachineInstr &MI) {
8677 assert((MI.getOpcode() == TargetOpcode::G_CTLZ ||
8678 MI.getOpcode() == TargetOpcode::G_CTTZ) &&
8679 "Expected count-zero opcode");
8680 switch (MI.getOpcode()) {
8681 case TargetOpcode::G_CTLZ:
8682 return TargetOpcode::G_CTLZ_ZERO_POISON;
8683 case TargetOpcode::G_CTTZ:
8684 return TargetOpcode::G_CTTZ_ZERO_POISON;
8685 default:
8686 llvm_unreachable("Unexpected count-zero opcode");
8687 }
8688}
8689
8690bool CombinerHelper::matchCountZeroToZeroPoison(MachineInstr &MI) const {
8691 if (!VT)
8692 return false;
8693
8694 unsigned ZPOpc = getCountZeroPoisonOpcode(MI);
8695 Register Src = MI.getOperand(i: 1).getReg();
8696 if (!VT->isKnownNeverZero(R: Src))
8697 return false;
8698
8699 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
8700 LLT SrcTy = MRI.getType(Reg: Src);
8701 return isLegalOrBeforeLegalizer(Query: {ZPOpc, {DstTy, SrcTy}});
8702}
8703
8704void CombinerHelper::applyCountZeroToZeroPoison(MachineInstr &MI) const {
8705 replaceOpcodeWith(FromMI&: MI, ToOpcode: getCountZeroPoisonOpcode(MI));
8706}
8707