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 if (!cast<GStore>(Val&: MI).isUnordered())
2749 return false;
2750 return getOpcodeDef(Opcode: TargetOpcode::G_IMPLICIT_DEF, Reg: MI.getOperand(i: 0).getReg(),
2751 MRI);
2752}
2753
2754bool CombinerHelper::matchInsertExtractVecEltOutOfBounds(
2755 MachineInstr &MI) const {
2756 assert((MI.getOpcode() == TargetOpcode::G_INSERT_VECTOR_ELT ||
2757 MI.getOpcode() == TargetOpcode::G_EXTRACT_VECTOR_ELT) &&
2758 "Expected an insert/extract element op");
2759 LLT VecTy = MRI.getType(Reg: MI.getOperand(i: 1).getReg());
2760 if (VecTy.isScalableVector())
2761 return false;
2762
2763 unsigned IdxIdx =
2764 MI.getOpcode() == TargetOpcode::G_EXTRACT_VECTOR_ELT ? 2 : 3;
2765 auto Idx = getIConstantVRegVal(VReg: MI.getOperand(i: IdxIdx).getReg(), MRI);
2766 if (!Idx)
2767 return false;
2768 return Idx->getZExtValue() >= VecTy.getNumElements();
2769}
2770
2771bool CombinerHelper::matchConstantSelectCmp(MachineInstr &MI,
2772 unsigned &OpIdx) const {
2773 GSelect &SelMI = cast<GSelect>(Val&: MI);
2774 auto Cst = isConstantOrConstantSplatVector(Def: SelMI.getCondReg(), MRI);
2775 if (!Cst)
2776 return false;
2777 OpIdx = Cst->isZero() ? 3 : 2;
2778 return true;
2779}
2780
2781void CombinerHelper::eraseInst(MachineInstr &MI) const { MI.eraseFromParent(); }
2782
2783bool CombinerHelper::matchEqualDefs(const MachineOperand &MOP1,
2784 const MachineOperand &MOP2) const {
2785 if (!MOP1.isReg() || !MOP2.isReg())
2786 return false;
2787 auto InstAndDef1 = getDefSrcRegIgnoringCopies(Reg: MOP1.getReg(), MRI);
2788 if (!InstAndDef1)
2789 return false;
2790 auto InstAndDef2 = getDefSrcRegIgnoringCopies(Reg: MOP2.getReg(), MRI);
2791 if (!InstAndDef2)
2792 return false;
2793 MachineInstr *I1 = InstAndDef1->MI;
2794 MachineInstr *I2 = InstAndDef2->MI;
2795
2796 // Handle a case like this:
2797 //
2798 // %0:_(s64), %1:_(s64) = G_UNMERGE_VALUES %2:_(<2 x s64>)
2799 //
2800 // Even though %0 and %1 are produced by the same instruction they are not
2801 // the same values.
2802 if (I1 == I2)
2803 return MOP1.getReg() == MOP2.getReg();
2804
2805 // If we have an instruction which loads or stores, we can't guarantee that
2806 // it is identical.
2807 //
2808 // For example, we may have
2809 //
2810 // %x1 = G_LOAD %addr (load N from @somewhere)
2811 // ...
2812 // call @foo
2813 // ...
2814 // %x2 = G_LOAD %addr (load N from @somewhere)
2815 // ...
2816 // %or = G_OR %x1, %x2
2817 //
2818 // It's possible that @foo will modify whatever lives at the address we're
2819 // loading from. To be safe, let's just assume that all loads and stores
2820 // are different (unless we have something which is guaranteed to not
2821 // change.)
2822 if (I1->mayLoadOrStore() && !I1->isDereferenceableInvariantLoad())
2823 return false;
2824
2825 // If both instructions are loads or stores, they are equal only if both
2826 // are dereferenceable invariant loads with the same number of bits.
2827 if (I1->mayLoadOrStore() && I2->mayLoadOrStore()) {
2828 GLoadStore *LS1 = dyn_cast<GLoadStore>(Val: I1);
2829 GLoadStore *LS2 = dyn_cast<GLoadStore>(Val: I2);
2830 if (!LS1 || !LS2)
2831 return false;
2832
2833 if (!I2->isDereferenceableInvariantLoad() ||
2834 (LS1->getMemSizeInBits() != LS2->getMemSizeInBits()))
2835 return false;
2836 }
2837
2838 // Check for physical registers on the instructions first to avoid cases
2839 // like this:
2840 //
2841 // %a = COPY $physreg
2842 // ...
2843 // SOMETHING implicit-def $physreg
2844 // ...
2845 // %b = COPY $physreg
2846 //
2847 // These copies are not equivalent.
2848 if (any_of(Range: I1->uses(), P: [](const MachineOperand &MO) {
2849 return MO.isReg() && MO.getReg().isPhysical();
2850 })) {
2851 // Check if we have a case like this:
2852 //
2853 // %a = COPY $physreg
2854 // %b = COPY %a
2855 //
2856 // In this case, I1 and I2 will both be equal to %a = COPY $physreg.
2857 // From that, we know that they must have the same value, since they must
2858 // have come from the same COPY.
2859 return I1->isIdenticalTo(Other: *I2);
2860 }
2861
2862 // We don't have any physical registers, so we don't necessarily need the
2863 // same vreg defs.
2864 //
2865 // On the off-chance that there's some target instruction feeding into the
2866 // instruction, let's use produceSameValue instead of isIdenticalTo.
2867 if (Builder.getTII().produceSameValue(MI0: *I1, MI1: *I2, MRI: &MRI)) {
2868 // Handle instructions with multiple defs that produce same values. Values
2869 // are same for operands with same index.
2870 // %0:_(s8), %1:_(s8), %2:_(s8), %3:_(s8) = G_UNMERGE_VALUES %4:_(<4 x s8>)
2871 // %5:_(s8), %6:_(s8), %7:_(s8), %8:_(s8) = G_UNMERGE_VALUES %4:_(<4 x s8>)
2872 // I1 and I2 are different instructions but produce same values,
2873 // %1 and %6 are same, %1 and %7 are not the same value.
2874 return I1->findRegisterDefOperandIdx(Reg: InstAndDef1->Reg, /*TRI=*/nullptr) ==
2875 I2->findRegisterDefOperandIdx(Reg: InstAndDef2->Reg, /*TRI=*/nullptr);
2876 }
2877 return false;
2878}
2879
2880bool CombinerHelper::matchConstantFPOp(const MachineOperand &MOP,
2881 double C) const {
2882 if (!MOP.isReg())
2883 return false;
2884 std::optional<FPValueAndVReg> MaybeCst;
2885 if (!mi_match(R: MOP.getReg(), MRI, P: m_GFCstOrSplat(FPValReg&: MaybeCst)))
2886 return false;
2887
2888 return MaybeCst->Value.isExactlyValue(V: C);
2889}
2890
2891void CombinerHelper::replaceSingleDefInstWithOperand(MachineInstr &MI,
2892 unsigned OpIdx) const {
2893 assert(MI.getNumExplicitDefs() == 1 && "Expected one explicit def?");
2894 Register OldReg = MI.getOperand(i: 0).getReg();
2895 Register Replacement = MI.getOperand(i: OpIdx).getReg();
2896 assert(canReplaceReg(OldReg, Replacement, MRI) && "Cannot replace register?");
2897 replaceRegWith(MRI, FromReg: OldReg, ToReg: Replacement);
2898 MI.eraseFromParent();
2899}
2900
2901void CombinerHelper::replaceSingleDefInstWithReg(MachineInstr &MI,
2902 Register Replacement) const {
2903 assert(MI.getNumExplicitDefs() == 1 && "Expected one explicit def?");
2904 Register OldReg = MI.getOperand(i: 0).getReg();
2905 assert(canReplaceReg(OldReg, Replacement, MRI) && "Cannot replace register?");
2906 replaceRegWith(MRI, FromReg: OldReg, ToReg: Replacement);
2907 MI.eraseFromParent();
2908}
2909
2910bool CombinerHelper::matchConstantLargerBitWidth(MachineInstr &MI,
2911 unsigned ConstIdx) const {
2912 Register ConstReg = MI.getOperand(i: ConstIdx).getReg();
2913 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
2914
2915 // Get the shift amount
2916 auto VRegAndVal = getIConstantVRegValWithLookThrough(VReg: ConstReg, MRI);
2917 if (!VRegAndVal)
2918 return false;
2919
2920 // Return true of shift amount >= Bitwidth
2921 return (VRegAndVal->Value.uge(RHS: DstTy.getSizeInBits()));
2922}
2923
2924void CombinerHelper::applyFunnelShiftConstantModulo(MachineInstr &MI) const {
2925 assert((MI.getOpcode() == TargetOpcode::G_FSHL ||
2926 MI.getOpcode() == TargetOpcode::G_FSHR) &&
2927 "This is not a funnel shift operation");
2928
2929 Register ConstReg = MI.getOperand(i: 3).getReg();
2930 LLT ConstTy = MRI.getType(Reg: ConstReg);
2931 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
2932
2933 auto VRegAndVal = getIConstantVRegValWithLookThrough(VReg: ConstReg, MRI);
2934 assert((VRegAndVal) && "Value is not a constant");
2935
2936 // Calculate the new Shift Amount = Old Shift Amount % BitWidth
2937 APInt NewConst = VRegAndVal->Value.urem(
2938 RHS: APInt(ConstTy.getSizeInBits(), DstTy.getScalarSizeInBits()));
2939
2940 auto NewConstInstr = Builder.buildConstant(Res: ConstTy, Val: NewConst.getZExtValue());
2941 Builder.buildInstr(
2942 Opc: MI.getOpcode(), DstOps: {MI.getOperand(i: 0)},
2943 SrcOps: {MI.getOperand(i: 1), MI.getOperand(i: 2), NewConstInstr.getReg(Idx: 0)});
2944
2945 MI.eraseFromParent();
2946}
2947
2948bool CombinerHelper::matchSelectSameVal(MachineInstr &MI) const {
2949 assert(MI.getOpcode() == TargetOpcode::G_SELECT);
2950 // Match (cond ? x : x)
2951 return matchEqualDefs(MOP1: MI.getOperand(i: 2), MOP2: MI.getOperand(i: 3)) &&
2952 canReplaceReg(DstReg: MI.getOperand(i: 0).getReg(), SrcReg: MI.getOperand(i: 2).getReg(),
2953 MRI);
2954}
2955
2956bool CombinerHelper::matchOperandIsKnownToBeAPowerOfTwo(
2957 const MachineOperand &MO, bool OrNegative) const {
2958 return isKnownToBeAPowerOfTwo(Val: MO.getReg(), MRI, ValueTracking: VT, OrNegative);
2959}
2960
2961void CombinerHelper::replaceInstWithFConstant(MachineInstr &MI,
2962 double C) const {
2963 assert(MI.getNumDefs() == 1 && "Expected only one def?");
2964 Builder.buildFConstant(Res: MI.getOperand(i: 0), Val: C);
2965 MI.eraseFromParent();
2966}
2967
2968void CombinerHelper::replaceInstWithConstant(MachineInstr &MI,
2969 int64_t C) const {
2970 assert(MI.getNumDefs() == 1 && "Expected only one def?");
2971 Builder.buildConstant(Res: MI.getOperand(i: 0), Val: C);
2972 MI.eraseFromParent();
2973}
2974
2975void CombinerHelper::replaceInstWithConstant(MachineInstr &MI, APInt C) const {
2976 assert(MI.getNumDefs() == 1 && "Expected only one def?");
2977 Builder.buildConstant(Res: MI.getOperand(i: 0), Val: C);
2978 MI.eraseFromParent();
2979}
2980
2981void CombinerHelper::replaceInstWithFConstant(MachineInstr &MI,
2982 ConstantFP *CFP) const {
2983 assert(MI.getNumDefs() == 1 && "Expected only one def?");
2984 Builder.buildFConstant(Res: MI.getOperand(i: 0), Val: CFP->getValueAPF());
2985 MI.eraseFromParent();
2986}
2987
2988void CombinerHelper::replaceInstWithUndef(MachineInstr &MI) const {
2989 assert(MI.getNumDefs() == 1 && "Expected only one def?");
2990 Builder.buildUndef(Res: MI.getOperand(i: 0));
2991 MI.eraseFromParent();
2992}
2993
2994bool CombinerHelper::matchCombineInsertVecElts(
2995 MachineInstr &MI, SmallVectorImpl<Register> &MatchInfo) const {
2996 assert(MI.getOpcode() == TargetOpcode::G_INSERT_VECTOR_ELT &&
2997 "Invalid opcode");
2998 Register DstReg = MI.getOperand(i: 0).getReg();
2999 LLT DstTy = MRI.getType(Reg: DstReg);
3000 assert(DstTy.isVector() && "Invalid G_INSERT_VECTOR_ELT?");
3001
3002 if (DstTy.isScalableVector())
3003 return false;
3004
3005 unsigned NumElts = DstTy.getNumElements();
3006 // If this MI is part of a sequence of insert_vec_elts, then
3007 // don't do the combine in the middle of the sequence.
3008 if (MRI.hasOneUse(RegNo: DstReg) && MRI.use_instr_begin(RegNo: DstReg)->getOpcode() ==
3009 TargetOpcode::G_INSERT_VECTOR_ELT)
3010 return false;
3011 MachineInstr *CurrInst = &MI;
3012 MachineInstr *TmpInst;
3013 int64_t IntImm;
3014 Register TmpReg;
3015 MatchInfo.resize(N: NumElts);
3016 while (mi_match(
3017 MI&: *CurrInst, MRI,
3018 P: m_GInsertVecElt(Src0: m_MInstr(MI&: TmpInst), Src1: m_Reg(R&: TmpReg), Src2: m_ICst(Cst&: IntImm)))) {
3019 if (IntImm >= NumElts || IntImm < 0)
3020 return false;
3021 if (!MatchInfo[IntImm])
3022 MatchInfo[IntImm] = TmpReg;
3023 CurrInst = TmpInst;
3024 }
3025 // Variable index.
3026 if (CurrInst->getOpcode() == TargetOpcode::G_INSERT_VECTOR_ELT)
3027 return false;
3028 if (TmpInst->getOpcode() == TargetOpcode::G_BUILD_VECTOR) {
3029 for (unsigned I = 1; I < TmpInst->getNumOperands(); ++I) {
3030 if (!MatchInfo[I - 1].isValid())
3031 MatchInfo[I - 1] = TmpInst->getOperand(i: I).getReg();
3032 }
3033 return true;
3034 }
3035 // If we didn't end in a G_IMPLICIT_DEF and the source is not fully
3036 // overwritten, bail out.
3037 return TmpInst->getOpcode() == TargetOpcode::G_IMPLICIT_DEF ||
3038 all_of(Range&: MatchInfo, P: [](Register Reg) { return !!Reg; });
3039}
3040
3041void CombinerHelper::applyCombineInsertVecElts(
3042 MachineInstr &MI, SmallVectorImpl<Register> &MatchInfo) const {
3043 Register UndefReg;
3044 auto GetUndef = [&]() {
3045 if (UndefReg)
3046 return UndefReg;
3047 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
3048 UndefReg = Builder.buildUndef(Res: DstTy.getScalarType()).getReg(Idx: 0);
3049 return UndefReg;
3050 };
3051 for (Register &Reg : MatchInfo) {
3052 if (!Reg)
3053 Reg = GetUndef();
3054 }
3055 Builder.buildBuildVector(Res: MI.getOperand(i: 0).getReg(), Ops: MatchInfo);
3056 MI.eraseFromParent();
3057}
3058
3059bool CombinerHelper::matchBinopWithNegInner(Register MInner, Register Other,
3060 unsigned RootOpc, Register Dst,
3061 LLT Ty,
3062 BuildFnTy &MatchInfo) const {
3063 /// Helper function for matchBinopWithNeg: tries to match one commuted form
3064 /// of `a bitwiseop (~b +/- c)` -> `a bitwiseop ~(b -/+ c)`.
3065 MachineInstr *InnerDef;
3066 if (!mi_match(R: MInner, MRI, P: m_MInstr(MI&: InnerDef)))
3067 return false;
3068
3069 unsigned InnerOpc = InnerDef->getOpcode();
3070 if (InnerOpc != TargetOpcode::G_ADD && InnerOpc != TargetOpcode::G_SUB)
3071 return false;
3072
3073 if (!MRI.hasOneNonDBGUse(RegNo: MInner))
3074 return false;
3075
3076 Register InnerLHS = InnerDef->getOperand(i: 1).getReg();
3077 Register InnerRHS = InnerDef->getOperand(i: 2).getReg();
3078 Register NotSrc;
3079 Register B, C;
3080
3081 // Check if either operand is ~b
3082 auto TryMatch = [&](Register MaybeNot, Register Other) {
3083 if (mi_match(R: MaybeNot, MRI, P: m_Not(Src: m_Reg(R&: NotSrc)))) {
3084 if (!MRI.hasOneNonDBGUse(RegNo: MaybeNot))
3085 return false;
3086 B = NotSrc;
3087 C = Other;
3088 return true;
3089 }
3090 return false;
3091 };
3092
3093 // For SUB, the not must be the LHS. For ADD, it can be either operand.
3094 if (!TryMatch(InnerLHS, InnerRHS) &&
3095 !(InnerOpc == TargetOpcode::G_ADD && TryMatch(InnerRHS, InnerLHS)))
3096 return false;
3097
3098 // Flip add/sub
3099 unsigned FlippedOpc = (InnerOpc == TargetOpcode::G_ADD) ? TargetOpcode::G_SUB
3100 : TargetOpcode::G_ADD;
3101
3102 Register A = Other;
3103 MatchInfo = [=](MachineIRBuilder &Builder) {
3104 auto NewInner = Builder.buildInstr(Opc: FlippedOpc, DstOps: {Ty}, SrcOps: {B, C});
3105 auto NewNot = Builder.buildNot(Dst: Ty, Src0: NewInner);
3106 Builder.buildInstr(Opc: RootOpc, DstOps: {Dst}, SrcOps: {A, NewNot});
3107 };
3108 return true;
3109}
3110
3111bool CombinerHelper::matchBinopWithNeg(MachineInstr &MI,
3112 BuildFnTy &MatchInfo) const {
3113 // Fold `a bitwiseop (~b +/- c)` -> `a bitwiseop ~(b -/+ c)`
3114 // Root MI is one of G_AND, G_OR, G_XOR.
3115 // We also look for commuted forms of operations. Pattern shouldn't apply
3116 // if there are multiple reasons of inner operations.
3117
3118 unsigned RootOpc = MI.getOpcode();
3119 Register Dst = MI.getOperand(i: 0).getReg();
3120 LLT Ty = MRI.getType(Reg: Dst);
3121
3122 Register LHS = MI.getOperand(i: 1).getReg();
3123 Register RHS = MI.getOperand(i: 2).getReg();
3124 // Check the commuted and uncommuted forms of the operation.
3125 return matchBinopWithNegInner(MInner: LHS, Other: RHS, RootOpc, Dst, Ty, MatchInfo) ||
3126 matchBinopWithNegInner(MInner: RHS, Other: LHS, RootOpc, Dst, Ty, MatchInfo);
3127}
3128
3129bool CombinerHelper::matchHoistLogicOpWithSameOpcodeHands(
3130 MachineInstr &MI, InstructionStepsMatchInfo &MatchInfo) const {
3131 // Matches: logic (hand x, ...), (hand y, ...) -> hand (logic x, y), ...
3132 //
3133 // Creates the new hand + logic instruction (but does not insert them.)
3134 //
3135 // On success, MatchInfo is populated with the new instructions. These are
3136 // inserted in applyHoistLogicOpWithSameOpcodeHands.
3137 unsigned LogicOpcode = MI.getOpcode();
3138 assert(LogicOpcode == TargetOpcode::G_AND ||
3139 LogicOpcode == TargetOpcode::G_OR ||
3140 LogicOpcode == TargetOpcode::G_XOR);
3141 MachineIRBuilder MIB(MI);
3142 Register Dst = MI.getOperand(i: 0).getReg();
3143 Register LHSReg = MI.getOperand(i: 1).getReg();
3144 Register RHSReg = MI.getOperand(i: 2).getReg();
3145
3146 // Don't recompute anything.
3147 if (!MRI.hasOneNonDBGUse(RegNo: LHSReg) || !MRI.hasOneNonDBGUse(RegNo: RHSReg))
3148 return false;
3149
3150 // Make sure we have (hand x, ...), (hand y, ...)
3151 MachineInstr *LeftHandInst = getDefIgnoringCopies(Reg: LHSReg, MRI);
3152 MachineInstr *RightHandInst = getDefIgnoringCopies(Reg: RHSReg, MRI);
3153 if (!LeftHandInst || !RightHandInst)
3154 return false;
3155 unsigned HandOpcode = LeftHandInst->getOpcode();
3156 if (HandOpcode != RightHandInst->getOpcode())
3157 return false;
3158 if (LeftHandInst->getNumOperands() < 2 ||
3159 !LeftHandInst->getOperand(i: 1).isReg() ||
3160 RightHandInst->getNumOperands() < 2 ||
3161 !RightHandInst->getOperand(i: 1).isReg())
3162 return false;
3163
3164 // Make sure the types match up, and if we're doing this post-legalization,
3165 // we end up with legal types.
3166 Register X = LeftHandInst->getOperand(i: 1).getReg();
3167 Register Y = RightHandInst->getOperand(i: 1).getReg();
3168 LLT XTy = MRI.getType(Reg: X);
3169 LLT YTy = MRI.getType(Reg: Y);
3170 if (!XTy.isValid() || XTy != YTy)
3171 return false;
3172
3173 // Optional extra source register.
3174 Register ExtraHandOpSrcReg;
3175 switch (HandOpcode) {
3176 default:
3177 return false;
3178 case TargetOpcode::G_ANYEXT:
3179 case TargetOpcode::G_SEXT:
3180 case TargetOpcode::G_ZEXT: {
3181 // Match: logic (ext X), (ext Y) --> ext (logic X, Y)
3182 break;
3183 }
3184 case TargetOpcode::G_TRUNC: {
3185 // Match: logic (trunc X), (trunc Y) -> trunc (logic X, Y)
3186 const MachineFunction *MF = MI.getMF();
3187 LLVMContext &Ctx = MF->getFunction().getContext();
3188
3189 LLT DstTy = MRI.getType(Reg: Dst);
3190 const TargetLowering &TLI = getTargetLowering();
3191
3192 // Be extra careful sinking truncate. If it's free, there's no benefit in
3193 // widening a binop.
3194 if (TLI.isZExtFree(FromTy: DstTy, ToTy: XTy, Ctx) && TLI.isTruncateFree(FromTy: XTy, ToTy: DstTy, Ctx))
3195 return false;
3196 break;
3197 }
3198 case TargetOpcode::G_AND:
3199 case TargetOpcode::G_ASHR:
3200 case TargetOpcode::G_LSHR:
3201 case TargetOpcode::G_SHL: {
3202 // Match: logic (binop x, z), (binop y, z) -> binop (logic x, y), z
3203 MachineOperand &ZOp = LeftHandInst->getOperand(i: 2);
3204 if (!matchEqualDefs(MOP1: ZOp, MOP2: RightHandInst->getOperand(i: 2)))
3205 return false;
3206 ExtraHandOpSrcReg = ZOp.getReg();
3207 break;
3208 }
3209 }
3210
3211 if (!isLegalOrBeforeLegalizer(Query: {LogicOpcode, {XTy, YTy}}))
3212 return false;
3213
3214 // Record the steps to build the new instructions.
3215 //
3216 // Steps to build (logic x, y)
3217 auto NewLogicDst = MRI.createGenericVirtualRegister(Ty: XTy);
3218 OperandBuildSteps LogicBuildSteps = {
3219 [=](MachineInstrBuilder &MIB) { MIB.addDef(RegNo: NewLogicDst); },
3220 [=](MachineInstrBuilder &MIB) { MIB.addReg(RegNo: X); },
3221 [=](MachineInstrBuilder &MIB) { MIB.addReg(RegNo: Y); }};
3222 InstructionBuildSteps LogicSteps(LogicOpcode, LogicBuildSteps);
3223
3224 // Steps to build hand (logic x, y), ...z
3225 OperandBuildSteps HandBuildSteps = {
3226 [=](MachineInstrBuilder &MIB) { MIB.addDef(RegNo: Dst); },
3227 [=](MachineInstrBuilder &MIB) { MIB.addReg(RegNo: NewLogicDst); }};
3228 if (ExtraHandOpSrcReg.isValid())
3229 HandBuildSteps.push_back(
3230 Elt: [=](MachineInstrBuilder &MIB) { MIB.addReg(RegNo: ExtraHandOpSrcReg); });
3231 InstructionBuildSteps HandSteps(HandOpcode, HandBuildSteps);
3232
3233 MatchInfo = InstructionStepsMatchInfo({LogicSteps, HandSteps});
3234 return true;
3235}
3236
3237void CombinerHelper::applyBuildInstructionSteps(
3238 MachineInstr &MI, InstructionStepsMatchInfo &MatchInfo) const {
3239 assert(MatchInfo.InstrsToBuild.size() &&
3240 "Expected at least one instr to build?");
3241 for (auto &InstrToBuild : MatchInfo.InstrsToBuild) {
3242 assert(InstrToBuild.Opcode && "Expected a valid opcode?");
3243 assert(InstrToBuild.OperandFns.size() && "Expected at least one operand?");
3244 MachineInstrBuilder Instr = Builder.buildInstr(Opcode: InstrToBuild.Opcode);
3245 for (auto &OperandFn : InstrToBuild.OperandFns)
3246 OperandFn(Instr);
3247 }
3248 MI.eraseFromParent();
3249}
3250
3251bool CombinerHelper::matchAshrShlToSextInreg(
3252 MachineInstr &MI, std::tuple<Register, int64_t> &MatchInfo) const {
3253 assert(MI.getOpcode() == TargetOpcode::G_ASHR);
3254 int64_t ShlCst, AshrCst;
3255 Register Src;
3256 if (!mi_match(R: MI.getOperand(i: 0).getReg(), MRI,
3257 P: m_GAShr(L: m_GShl(L: m_Reg(R&: Src), R: m_ICstOrSplat(Cst&: ShlCst)),
3258 R: m_ICstOrSplat(Cst&: AshrCst))))
3259 return false;
3260 if (ShlCst != AshrCst)
3261 return false;
3262 if (!isLegalOrBeforeLegalizer(
3263 Query: {TargetOpcode::G_SEXT_INREG,
3264 {MRI.getType(Reg: Src)},
3265 {},
3266 {MRI.getType(Reg: Src).getScalarSizeInBits() - ShlCst}}))
3267 return false;
3268 MatchInfo = std::make_tuple(args&: Src, args&: ShlCst);
3269 return true;
3270}
3271
3272void CombinerHelper::applyAshShlToSextInreg(
3273 MachineInstr &MI, std::tuple<Register, int64_t> &MatchInfo) const {
3274 assert(MI.getOpcode() == TargetOpcode::G_ASHR);
3275 Register Src;
3276 int64_t ShiftAmt;
3277 std::tie(args&: Src, args&: ShiftAmt) = MatchInfo;
3278 unsigned Size = MRI.getType(Reg: Src).getScalarSizeInBits();
3279 Builder.buildSExtInReg(Res: MI.getOperand(i: 0).getReg(), Op: Src, ImmOp: Size - ShiftAmt);
3280 MI.eraseFromParent();
3281}
3282
3283/// and(and(x, C1), C2) -> C1&C2 ? and(x, C1&C2) : 0
3284bool CombinerHelper::matchOverlappingAnd(
3285 MachineInstr &MI,
3286 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
3287 assert(MI.getOpcode() == TargetOpcode::G_AND);
3288
3289 Register Dst = MI.getOperand(i: 0).getReg();
3290 LLT Ty = MRI.getType(Reg: Dst);
3291
3292 Register R;
3293 int64_t C1;
3294 int64_t C2;
3295 if (!mi_match(
3296 R: Dst, MRI,
3297 P: m_GAnd(L: m_GAnd(L: m_Reg(R), R: m_ICst(Cst&: C1)), R: m_ICst(Cst&: C2))))
3298 return false;
3299
3300 MatchInfo = [=](MachineIRBuilder &B) {
3301 if (C1 & C2) {
3302 B.buildAnd(Dst, Src0: R, Src1: B.buildConstant(Res: Ty, Val: C1 & C2));
3303 return;
3304 }
3305 auto Zero = B.buildConstant(Res: Ty, Val: 0);
3306 replaceRegWith(MRI, FromReg: Dst, ToReg: Zero->getOperand(i: 0).getReg());
3307 };
3308 return true;
3309}
3310
3311bool CombinerHelper::matchRedundantAnd(MachineInstr &MI,
3312 Register &Replacement) const {
3313 // Given
3314 //
3315 // %y:_(sN) = G_SOMETHING
3316 // %x:_(sN) = G_SOMETHING
3317 // %res:_(sN) = G_AND %x, %y
3318 //
3319 // Eliminate the G_AND when it is known that x & y == x or x & y == y.
3320 //
3321 // Patterns like this can appear as a result of legalization. E.g.
3322 //
3323 // %cmp:_(s32) = G_ICMP intpred(pred), %x(s32), %y
3324 // %one:_(s32) = G_CONSTANT i32 1
3325 // %and:_(s32) = G_AND %cmp, %one
3326 //
3327 // In this case, G_ICMP only produces a single bit, so x & 1 == x.
3328 assert(MI.getOpcode() == TargetOpcode::G_AND);
3329 if (!VT)
3330 return false;
3331
3332 Register AndDst = MI.getOperand(i: 0).getReg();
3333 Register LHS = MI.getOperand(i: 1).getReg();
3334 Register RHS = MI.getOperand(i: 2).getReg();
3335
3336 // Check the RHS (maybe a constant) first, and if we have no KnownBits there,
3337 // we can't do anything. If we do, then it depends on whether we have
3338 // KnownBits on the LHS.
3339 KnownBits RHSBits = VT->getKnownBits(R: RHS);
3340 if (RHSBits.isUnknown())
3341 return false;
3342
3343 KnownBits LHSBits = VT->getKnownBits(R: LHS);
3344
3345 // Check that x & Mask == x.
3346 // x & 1 == x, always
3347 // x & 0 == x, only if x is also 0
3348 // Meaning Mask has no effect if every bit is either one in Mask or zero in x.
3349 //
3350 // Check if we can replace AndDst with the LHS of the G_AND
3351 if (canReplaceReg(DstReg: AndDst, SrcReg: LHS, MRI) &&
3352 (LHSBits.Zero | RHSBits.One).isAllOnes()) {
3353 Replacement = LHS;
3354 return true;
3355 }
3356
3357 // Check if we can replace AndDst with the RHS of the G_AND
3358 if (canReplaceReg(DstReg: AndDst, SrcReg: RHS, MRI) &&
3359 (LHSBits.One | RHSBits.Zero).isAllOnes()) {
3360 Replacement = RHS;
3361 return true;
3362 }
3363
3364 return false;
3365}
3366
3367bool CombinerHelper::matchRedundantOr(MachineInstr &MI,
3368 Register &Replacement) const {
3369 // Given
3370 //
3371 // %y:_(sN) = G_SOMETHING
3372 // %x:_(sN) = G_SOMETHING
3373 // %res:_(sN) = G_OR %x, %y
3374 //
3375 // Eliminate the G_OR when it is known that x | y == x or x | y == y.
3376 assert(MI.getOpcode() == TargetOpcode::G_OR);
3377 if (!VT)
3378 return false;
3379
3380 Register OrDst = MI.getOperand(i: 0).getReg();
3381 Register LHS = MI.getOperand(i: 1).getReg();
3382 Register RHS = MI.getOperand(i: 2).getReg();
3383
3384 KnownBits LHSBits = VT->getKnownBits(R: LHS);
3385 KnownBits RHSBits = VT->getKnownBits(R: RHS);
3386
3387 // Check that x | Mask == x.
3388 // x | 0 == x, always
3389 // x | 1 == x, only if x is also 1
3390 // Meaning Mask has no effect if every bit is either zero in Mask or one in x.
3391 //
3392 // Check if we can replace OrDst with the LHS of the G_OR
3393 if (canReplaceReg(DstReg: OrDst, SrcReg: LHS, MRI) &&
3394 (LHSBits.One | RHSBits.Zero).isAllOnes()) {
3395 Replacement = LHS;
3396 return true;
3397 }
3398
3399 // Check if we can replace OrDst with the RHS of the G_OR
3400 if (canReplaceReg(DstReg: OrDst, SrcReg: RHS, MRI) &&
3401 (LHSBits.Zero | RHSBits.One).isAllOnes()) {
3402 Replacement = RHS;
3403 return true;
3404 }
3405
3406 return false;
3407}
3408
3409bool CombinerHelper::matchRedundantSExtInReg(MachineInstr &MI) const {
3410 // If the input is already sign extended, just drop the extension.
3411 Register Src = MI.getOperand(i: 1).getReg();
3412 unsigned ExtBits = MI.getOperand(i: 2).getImm();
3413 unsigned TypeSize = MRI.getType(Reg: Src).getScalarSizeInBits();
3414 return VT->computeNumSignBits(R: Src) >= (TypeSize - ExtBits + 1);
3415}
3416
3417static bool isConstValidTrue(const TargetLowering &TLI, unsigned ScalarSizeBits,
3418 int64_t Cst, bool IsVector, bool IsFP) {
3419 // For i1, Cst will always be -1 regardless of boolean contents.
3420 return (ScalarSizeBits == 1 && Cst == -1) ||
3421 isConstTrueVal(TLI, Val: Cst, IsVector, IsFP);
3422}
3423
3424// This pattern aims to match the following shape to avoid extra mov
3425// instructions
3426// G_BUILD_VECTOR(
3427// G_UNMERGE_VALUES(src, 0)
3428// G_UNMERGE_VALUES(src, 1)
3429// G_IMPLICIT_DEF
3430// G_IMPLICIT_DEF
3431// )
3432// ->
3433// G_CONCAT_VECTORS(
3434// src,
3435// undef
3436// )
3437bool CombinerHelper::matchCombineBuildUnmerge(MachineInstr &MI,
3438 MachineRegisterInfo &MRI,
3439 Register &UnmergeSrc) const {
3440 auto &BV = cast<GBuildVector>(Val&: MI);
3441
3442 unsigned BuildUseCount = BV.getNumSources();
3443 if (BuildUseCount % 2 != 0)
3444 return false;
3445
3446 unsigned NumUnmerge = BuildUseCount / 2;
3447
3448 auto *Unmerge = getOpcodeDef<GUnmerge>(Reg: BV.getSourceReg(I: 0), MRI);
3449
3450 // Check the first operand is an unmerge and has the correct number of
3451 // operands
3452 if (!Unmerge || Unmerge->getNumDefs() != NumUnmerge)
3453 return false;
3454
3455 UnmergeSrc = Unmerge->getSourceReg();
3456
3457 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
3458 LLT UnmergeSrcTy = MRI.getType(Reg: UnmergeSrc);
3459
3460 if (!UnmergeSrcTy.isVector())
3461 return false;
3462
3463 // Ensure we only generate legal instructions post-legalizer
3464 if (!IsPreLegalize &&
3465 !isLegal(Query: {TargetOpcode::G_CONCAT_VECTORS, {DstTy, UnmergeSrcTy}}))
3466 return false;
3467
3468 // Check that all of the operands before the midpoint come from the same
3469 // unmerge and are in the same order as they are used in the build_vector
3470 for (unsigned I = 0; I < NumUnmerge; ++I) {
3471 auto MaybeUnmergeReg = BV.getSourceReg(I);
3472 auto *LoopUnmerge = getOpcodeDef<GUnmerge>(Reg: MaybeUnmergeReg, MRI);
3473
3474 if (!LoopUnmerge || LoopUnmerge != Unmerge)
3475 return false;
3476
3477 if (LoopUnmerge->getOperand(i: I).getReg() != MaybeUnmergeReg)
3478 return false;
3479 }
3480
3481 // Check that all of the unmerged values are used
3482 if (Unmerge->getNumDefs() != NumUnmerge)
3483 return false;
3484
3485 // Check that all of the operands after the mid point are undefs.
3486 for (unsigned I = NumUnmerge; I < BuildUseCount; ++I) {
3487 auto *Undef = getDefIgnoringCopies(Reg: BV.getSourceReg(I), MRI);
3488
3489 if (Undef->getOpcode() != TargetOpcode::G_IMPLICIT_DEF)
3490 return false;
3491 }
3492
3493 return true;
3494}
3495
3496void CombinerHelper::applyCombineBuildUnmerge(MachineInstr &MI,
3497 MachineRegisterInfo &MRI,
3498 MachineIRBuilder &B,
3499 Register &UnmergeSrc) const {
3500 assert(UnmergeSrc && "Expected there to be one matching G_UNMERGE_VALUES");
3501 B.setInstrAndDebugLoc(MI);
3502
3503 Register UndefVec = B.buildUndef(Res: MRI.getType(Reg: UnmergeSrc)).getReg(Idx: 0);
3504 B.buildConcatVectors(Res: MI.getOperand(i: 0), Ops: {UnmergeSrc, UndefVec});
3505
3506 MI.eraseFromParent();
3507}
3508
3509// This combine tries to reduce the number of scalarised G_TRUNC instructions by
3510// using vector truncates instead
3511//
3512// EXAMPLE:
3513// %a(i32), %b(i32) = G_UNMERGE_VALUES %src(<2 x i32>)
3514// %T_a(i16) = G_TRUNC %a(i32)
3515// %T_b(i16) = G_TRUNC %b(i32)
3516// %Undef(i16) = G_IMPLICIT_DEF(i16)
3517// %dst(v4i16) = G_BUILD_VECTORS %T_a(i16), %T_b(i16), %Undef(i16), %Undef(i16)
3518//
3519// ===>
3520// %Undef(<2 x i32>) = G_IMPLICIT_DEF(<2 x i32>)
3521// %Mid(<4 x s32>) = G_CONCAT_VECTORS %src(<2 x i32>), %Undef(<2 x i32>)
3522// %dst(<4 x s16>) = G_TRUNC %Mid(<4 x s32>)
3523//
3524// Only matches sources made up of G_TRUNCs followed by G_IMPLICIT_DEFs
3525bool CombinerHelper::matchUseVectorTruncate(MachineInstr &MI,
3526 Register &MatchInfo) const {
3527 auto BuildMI = cast<GBuildVector>(Val: &MI);
3528 unsigned NumOperands = BuildMI->getNumSources();
3529 LLT DstTy = MRI.getType(Reg: BuildMI->getReg(Idx: 0));
3530
3531 // Check the G_BUILD_VECTOR sources
3532 unsigned I;
3533 GUnmerge *UnmergeMI = nullptr;
3534
3535 // Check all source TRUNCs come from the same UNMERGE instruction
3536 // and that the element order matches (BUILD_VECTOR position I
3537 // corresponds to UNMERGE result I)
3538 for (I = 0; I < NumOperands; ++I) {
3539 // Check if the G_TRUNC instructions all come from the same MI
3540 Register TruncSrcReg;
3541 if (!mi_match(R: BuildMI->getSourceReg(I), MRI, P: m_GTrunc(Src: m_Reg(R&: TruncSrcReg))))
3542 break;
3543
3544 if (!UnmergeMI) {
3545 if (!mi_match(R: TruncSrcReg, MRI, P: m_GUnmerge(Inst&: UnmergeMI)))
3546 return false;
3547 } else {
3548 MachineInstr *UnmergeSrcMI;
3549 if (!mi_match(R: TruncSrcReg, MRI, P: m_MInstr(MI&: UnmergeSrcMI)) ||
3550 UnmergeMI != UnmergeSrcMI)
3551 return false;
3552 }
3553 // Element order must match: position I must use UNMERGE result I.
3554 if (UnmergeMI->getOperand(i: I).getReg() != TruncSrcReg)
3555 return false;
3556 }
3557 if (I < 2)
3558 return false;
3559
3560 // Check the remaining source elements are only G_IMPLICIT_DEF
3561 for (; I < NumOperands; ++I) {
3562 if (!mi_match(R: BuildMI->getSourceReg(I), MRI, P: m_GImplicitDef()))
3563 return false;
3564 }
3565
3566 // Check the size of unmerge source
3567 MatchInfo = UnmergeMI->getSourceReg();
3568 LLT UnmergeSrcTy = MRI.getType(Reg: MatchInfo);
3569 if (!UnmergeSrcTy.isVector())
3570 return false;
3571
3572 if (!DstTy.getElementCount().isKnownMultipleOf(RHS: UnmergeSrcTy.getNumElements()))
3573 return false;
3574
3575 // Check the unmerge source and destination element types match
3576 LLT UnmergeSrcEltTy = UnmergeSrcTy.getElementType();
3577 Register UnmergeDstReg = UnmergeMI->getOperand(i: 0).getReg();
3578 LLT UnmergeDstEltTy = MRI.getType(Reg: UnmergeDstReg);
3579 if (UnmergeSrcEltTy != UnmergeDstEltTy)
3580 return false;
3581
3582 // Only generate legal instructions post-legalizer
3583 if (!IsPreLegalize) {
3584 LLT MidTy = DstTy.changeElementType(NewEltTy: UnmergeSrcTy.getScalarType());
3585
3586 if (DstTy.getElementCount() != UnmergeSrcTy.getElementCount() &&
3587 !isLegal(Query: {TargetOpcode::G_CONCAT_VECTORS, {MidTy, UnmergeSrcTy}}))
3588 return false;
3589
3590 if (!isLegal(Query: {TargetOpcode::G_TRUNC, {DstTy, MidTy}}))
3591 return false;
3592 }
3593
3594 return true;
3595}
3596
3597void CombinerHelper::applyUseVectorTruncate(MachineInstr &MI,
3598 Register &MatchInfo) const {
3599 Register MidReg;
3600 auto BuildMI = cast<GBuildVector>(Val: &MI);
3601 Register DstReg = BuildMI->getReg(Idx: 0);
3602 LLT DstTy = MRI.getType(Reg: DstReg);
3603 LLT UnmergeSrcTy = MRI.getType(Reg: MatchInfo);
3604 unsigned DstTyNumElt = DstTy.getNumElements();
3605 unsigned UnmergeSrcTyNumElt = UnmergeSrcTy.getNumElements();
3606
3607 // No need to pad vector if only G_TRUNC is needed
3608 if (DstTyNumElt / UnmergeSrcTyNumElt == 1) {
3609 MidReg = MatchInfo;
3610 } else {
3611 Register UndefReg = Builder.buildUndef(Res: UnmergeSrcTy).getReg(Idx: 0);
3612 SmallVector<Register> ConcatRegs = {MatchInfo};
3613 for (unsigned I = 1; I < DstTyNumElt / UnmergeSrcTyNumElt; ++I)
3614 ConcatRegs.push_back(Elt: UndefReg);
3615
3616 auto MidTy = DstTy.changeElementType(NewEltTy: UnmergeSrcTy.getScalarType());
3617 MidReg = Builder.buildConcatVectors(Res: MidTy, Ops: ConcatRegs).getReg(Idx: 0);
3618 }
3619
3620 Builder.buildTrunc(Res: DstReg, Op: MidReg);
3621 MI.eraseFromParent();
3622}
3623
3624bool CombinerHelper::matchNotCmp(
3625 MachineInstr &MI, SmallVectorImpl<Register> &RegsToNegate) const {
3626 assert(MI.getOpcode() == TargetOpcode::G_XOR);
3627 LLT Ty = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
3628 const auto &TLI = *Builder.getMF().getSubtarget().getTargetLowering();
3629 Register XorSrc;
3630 Register CstReg;
3631 // We match xor(src, true) here.
3632 if (!mi_match(R: MI.getOperand(i: 0).getReg(), MRI,
3633 P: m_GXor(L: m_Reg(R&: XorSrc), R: m_Reg(R&: CstReg))))
3634 return false;
3635
3636 if (!MRI.hasOneNonDBGUse(RegNo: XorSrc))
3637 return false;
3638
3639 // Check that XorSrc is the root of a tree of comparisons combined with ANDs
3640 // and ORs. The suffix of RegsToNegate starting from index I is used a work
3641 // list of tree nodes to visit.
3642 RegsToNegate.push_back(Elt: XorSrc);
3643 // Remember whether the comparisons are all integer or all floating point.
3644 bool IsInt = false;
3645 bool IsFP = false;
3646 for (unsigned I = 0; I < RegsToNegate.size(); ++I) {
3647 Register Reg = RegsToNegate[I];
3648 if (!MRI.hasOneNonDBGUse(RegNo: Reg))
3649 return false;
3650 MachineInstr *Def;
3651 if (!mi_match(R: Reg, MRI, P: m_MInstr(MI&: Def)))
3652 return false;
3653 switch (Def->getOpcode()) {
3654 default:
3655 // Don't match if the tree contains anything other than ANDs, ORs and
3656 // comparisons.
3657 return false;
3658 case TargetOpcode::G_ICMP:
3659 if (IsFP)
3660 return false;
3661 IsInt = true;
3662 // When we apply the combine we will invert the predicate.
3663 break;
3664 case TargetOpcode::G_FCMP:
3665 if (IsInt)
3666 return false;
3667 IsFP = true;
3668 // When we apply the combine we will invert the predicate.
3669 break;
3670 case TargetOpcode::G_AND:
3671 case TargetOpcode::G_OR:
3672 // Implement De Morgan's laws:
3673 // ~(x & y) -> ~x | ~y
3674 // ~(x | y) -> ~x & ~y
3675 // When we apply the combine we will change the opcode and recursively
3676 // negate the operands.
3677 RegsToNegate.push_back(Elt: Def->getOperand(i: 1).getReg());
3678 RegsToNegate.push_back(Elt: Def->getOperand(i: 2).getReg());
3679 break;
3680 }
3681 }
3682
3683 // Now we know whether the comparisons are integer or floating point, check
3684 // the constant in the xor.
3685 int64_t Cst;
3686 if (Ty.isVector()) {
3687 int64_t SplatCst;
3688 if (!mi_match(R: CstReg, MRI, P: m_ICstOrSplat(Cst&: SplatCst)))
3689 return false;
3690 if (!isConstValidTrue(TLI, ScalarSizeBits: Ty.getScalarSizeInBits(), Cst: SplatCst, IsVector: true, IsFP))
3691 return false;
3692 } else {
3693 if (!mi_match(R: CstReg, MRI, P: m_ICst(Cst)))
3694 return false;
3695 if (!isConstValidTrue(TLI, ScalarSizeBits: Ty.getSizeInBits(), Cst, IsVector: false, IsFP))
3696 return false;
3697 }
3698
3699 return true;
3700}
3701
3702void CombinerHelper::applyNotCmp(
3703 MachineInstr &MI, SmallVectorImpl<Register> &RegsToNegate) const {
3704 for (Register Reg : RegsToNegate) {
3705 MachineInstr *Def = MRI.getVRegDef(Reg);
3706 Observer.changingInstr(MI&: *Def);
3707 // For each comparison, invert the opcode. For each AND and OR, change the
3708 // opcode.
3709 switch (Def->getOpcode()) {
3710 default:
3711 llvm_unreachable("Unexpected opcode");
3712 case TargetOpcode::G_ICMP:
3713 case TargetOpcode::G_FCMP: {
3714 MachineOperand &PredOp = Def->getOperand(i: 1);
3715 CmpInst::Predicate NewP = CmpInst::getInversePredicate(
3716 pred: (CmpInst::Predicate)PredOp.getPredicate());
3717 PredOp.setPredicate(NewP);
3718 break;
3719 }
3720 case TargetOpcode::G_AND:
3721 Def->setDesc(Builder.getTII().get(Opcode: TargetOpcode::G_OR));
3722 break;
3723 case TargetOpcode::G_OR:
3724 Def->setDesc(Builder.getTII().get(Opcode: TargetOpcode::G_AND));
3725 break;
3726 }
3727 Observer.changedInstr(MI&: *Def);
3728 }
3729
3730 replaceRegWith(MRI, FromReg: MI.getOperand(i: 0).getReg(), ToReg: MI.getOperand(i: 1).getReg());
3731 MI.eraseFromParent();
3732}
3733
3734bool CombinerHelper::matchXorOfAndWithSameReg(
3735 MachineInstr &MI, std::pair<Register, Register> &MatchInfo) const {
3736 // Match (xor (and x, y), y) (or any of its commuted cases)
3737 assert(MI.getOpcode() == TargetOpcode::G_XOR);
3738 Register &X = MatchInfo.first;
3739 Register &Y = MatchInfo.second;
3740 Register AndReg = MI.getOperand(i: 1).getReg();
3741 Register SharedReg = MI.getOperand(i: 2).getReg();
3742
3743 // Find a G_AND on either side of the G_XOR.
3744 // Look for one of
3745 //
3746 // (xor (and x, y), SharedReg)
3747 // (xor SharedReg, (and x, y))
3748 if (!mi_match(R: AndReg, MRI, P: m_GAnd(L: m_Reg(R&: X), R: m_Reg(R&: Y)))) {
3749 std::swap(a&: AndReg, b&: SharedReg);
3750 if (!mi_match(R: AndReg, MRI, P: m_GAnd(L: m_Reg(R&: X), R: m_Reg(R&: Y))))
3751 return false;
3752 }
3753
3754 // Only do this if we'll eliminate the G_AND.
3755 if (!MRI.hasOneNonDBGUse(RegNo: AndReg))
3756 return false;
3757
3758 // We can combine if SharedReg is the same as either the LHS or RHS of the
3759 // G_AND.
3760 if (Y != SharedReg)
3761 std::swap(a&: X, b&: Y);
3762 return Y == SharedReg;
3763}
3764
3765void CombinerHelper::applyXorOfAndWithSameReg(
3766 MachineInstr &MI, std::pair<Register, Register> &MatchInfo) const {
3767 // Fold (xor (and x, y), y) -> (and (not x), y)
3768 Register X, Y;
3769 std::tie(args&: X, args&: Y) = MatchInfo;
3770 auto Not = Builder.buildNot(Dst: MRI.getType(Reg: X), Src0: X);
3771 Observer.changingInstr(MI);
3772 MI.setDesc(Builder.getTII().get(Opcode: TargetOpcode::G_AND));
3773 MI.getOperand(i: 1).setReg(Not->getOperand(i: 0).getReg());
3774 MI.getOperand(i: 2).setReg(Y);
3775 Observer.changedInstr(MI);
3776}
3777
3778bool CombinerHelper::matchPtrAddZero(MachineInstr &MI) const {
3779 auto &PtrAdd = cast<GPtrAdd>(Val&: MI);
3780 Register DstReg = PtrAdd.getReg(Idx: 0);
3781 LLT Ty = MRI.getType(Reg: DstReg);
3782 const DataLayout &DL = Builder.getMF().getDataLayout();
3783
3784 if (DL.isNonIntegralAddressSpace(AddrSpace: Ty.getScalarType().getAddressSpace()))
3785 return false;
3786
3787 if (Ty.isPointer()) {
3788 auto ConstVal = getIConstantVRegVal(VReg: PtrAdd.getBaseReg(), MRI);
3789 return ConstVal && *ConstVal == 0;
3790 }
3791
3792 assert(Ty.isVector() && "Expecting a vector type");
3793 const MachineInstr *VecMI;
3794 if (!mi_match(R: PtrAdd.getBaseReg(), MRI, P: m_MInstr(MI&: VecMI)))
3795 return false;
3796 return isBuildVectorAllZeros(MI: *VecMI, MRI);
3797}
3798
3799/// The second source operand is known to be a power of 2.
3800void CombinerHelper::applySimplifyURemByPow2(MachineInstr &MI) const {
3801 Register DstReg = MI.getOperand(i: 0).getReg();
3802 Register Src0 = MI.getOperand(i: 1).getReg();
3803 Register Pow2Src1 = MI.getOperand(i: 2).getReg();
3804 LLT Ty = MRI.getType(Reg: DstReg);
3805
3806 // Fold (urem x, pow2) -> (and x, pow2-1)
3807 auto NegOne = Builder.buildConstant(Res: Ty, Val: -1);
3808 auto Add = Builder.buildAdd(Dst: Ty, Src0: Pow2Src1, Src1: NegOne);
3809 Builder.buildAnd(Dst: DstReg, Src0, Src1: Add);
3810 MI.eraseFromParent();
3811}
3812
3813bool CombinerHelper::matchFoldBinOpIntoSelect(MachineInstr &MI,
3814 unsigned &SelectOpNo) const {
3815 Register LHS = MI.getOperand(i: 1).getReg();
3816 Register RHS = MI.getOperand(i: 2).getReg();
3817
3818 Register OtherOperandReg = RHS;
3819 SelectOpNo = 1;
3820 Register SelectTrue, SelectFalse;
3821
3822 // Don't do this unless the old select is going away. We want to eliminate the
3823 // binary operator, not replace a binop with a select.
3824 if (!mi_match(R: LHS, MRI,
3825 P: m_GISelect(Src0: m_Reg(), Src1: m_Reg(R&: SelectTrue), Src2: m_Reg(R&: SelectFalse))) ||
3826 !MRI.hasOneNonDBGUse(RegNo: LHS)) {
3827 OtherOperandReg = LHS;
3828 SelectOpNo = 2;
3829 if (!mi_match(R: RHS, MRI,
3830 P: m_GISelect(Src0: m_Reg(), Src1: m_Reg(R&: SelectTrue), Src2: m_Reg(R&: SelectFalse))) ||
3831 !MRI.hasOneNonDBGUse(RegNo: RHS))
3832 return false;
3833 }
3834
3835 MachineInstr *SelectLHS, *SelectRHS;
3836 if (!mi_match(R: SelectTrue, MRI, P: m_MInstr(MI&: SelectLHS)) ||
3837 !mi_match(R: SelectFalse, MRI, P: m_MInstr(MI&: SelectRHS)))
3838 return false;
3839
3840 if (!isConstantOrConstantVector(MI: *SelectLHS, MRI,
3841 /*AllowFP*/ true,
3842 /*AllowOpaqueConstants*/ false))
3843 return false;
3844 if (!isConstantOrConstantVector(MI: *SelectRHS, MRI,
3845 /*AllowFP*/ true,
3846 /*AllowOpaqueConstants*/ false))
3847 return false;
3848
3849 unsigned BinOpcode = MI.getOpcode();
3850
3851 // We know that one of the operands is a select of constants. Now verify that
3852 // the other binary operator operand is either a constant, or we can handle a
3853 // variable.
3854 bool CanFoldNonConst =
3855 (BinOpcode == TargetOpcode::G_AND || BinOpcode == TargetOpcode::G_OR) &&
3856 (isNullOrNullSplat(MI: *SelectLHS, MRI) ||
3857 isAllOnesOrAllOnesSplat(MI: *SelectLHS, MRI)) &&
3858 (isNullOrNullSplat(MI: *SelectRHS, MRI) ||
3859 isAllOnesOrAllOnesSplat(MI: *SelectRHS, MRI));
3860 if (CanFoldNonConst)
3861 return true;
3862
3863 MachineInstr *OtherOperandDef;
3864 if (!mi_match(R: OtherOperandReg, MRI, P: m_MInstr(MI&: OtherOperandDef)))
3865 return false;
3866 return isConstantOrConstantVector(MI: *OtherOperandDef, MRI,
3867 /*AllowFP*/ true,
3868 /*AllowOpaqueConstants*/ false);
3869}
3870
3871/// \p SelectOperand is the operand in binary operator \p MI that is the select
3872/// to fold.
3873void CombinerHelper::applyFoldBinOpIntoSelect(
3874 MachineInstr &MI, const unsigned &SelectOperand) const {
3875 Register Dst = MI.getOperand(i: 0).getReg();
3876 Register LHS = MI.getOperand(i: 1).getReg();
3877 Register RHS = MI.getOperand(i: 2).getReg();
3878 GSelect *Select =
3879 cast<GSelect>(Val: MRI.getVRegDef(Reg: MI.getOperand(i: SelectOperand).getReg()));
3880
3881 Register SelectCond = Select->getCondReg();
3882 Register SelectTrue = Select->getTrueReg();
3883 Register SelectFalse = Select->getFalseReg();
3884
3885 LLT Ty = MRI.getType(Reg: Dst);
3886 unsigned BinOpcode = MI.getOpcode();
3887
3888 Register FoldTrue, FoldFalse;
3889
3890 // We have a select-of-constants followed by a binary operator with a
3891 // constant. Eliminate the binop by pulling the constant math into the select.
3892 // Example: add (select Cond, CT, CF), CBO --> select Cond, CT + CBO, CF + CBO
3893 if (SelectOperand == 1) {
3894 // TODO: SelectionDAG verifies this actually constant folds before
3895 // committing to the combine.
3896
3897 FoldTrue = Builder.buildInstr(Opc: BinOpcode, DstOps: {Ty}, SrcOps: {SelectTrue, RHS}).getReg(Idx: 0);
3898 FoldFalse =
3899 Builder.buildInstr(Opc: BinOpcode, DstOps: {Ty}, SrcOps: {SelectFalse, RHS}).getReg(Idx: 0);
3900 } else {
3901 FoldTrue = Builder.buildInstr(Opc: BinOpcode, DstOps: {Ty}, SrcOps: {LHS, SelectTrue}).getReg(Idx: 0);
3902 FoldFalse =
3903 Builder.buildInstr(Opc: BinOpcode, DstOps: {Ty}, SrcOps: {LHS, SelectFalse}).getReg(Idx: 0);
3904 }
3905
3906 Builder.buildSelect(Res: Dst, Tst: SelectCond, Op0: FoldTrue, Op1: FoldFalse, Flags: MI.getFlags());
3907 MI.eraseFromParent();
3908}
3909
3910std::optional<SmallVector<Register, 8>>
3911CombinerHelper::findCandidatesForLoadOrCombine(const MachineInstr *Root) const {
3912 assert(Root->getOpcode() == TargetOpcode::G_OR && "Expected G_OR only!");
3913 // We want to detect if Root is part of a tree which represents a bunch
3914 // of loads being merged into a larger load. We'll try to recognize patterns
3915 // like, for example:
3916 //
3917 // Reg Reg
3918 // \ /
3919 // OR_1 Reg
3920 // \ /
3921 // OR_2
3922 // \ Reg
3923 // .. /
3924 // Root
3925 //
3926 // Reg Reg Reg Reg
3927 // \ / \ /
3928 // OR_1 OR_2
3929 // \ /
3930 // \ /
3931 // ...
3932 // Root
3933 //
3934 // Each "Reg" may have been produced by a load + some arithmetic. This
3935 // function will save each of them.
3936 SmallVector<Register, 8> RegsToVisit;
3937 SmallVector<const MachineInstr *, 7> Ors = {Root};
3938
3939 // In the "worst" case, we're dealing with a load for each byte. So, there
3940 // are at most #bytes - 1 ORs.
3941 const unsigned MaxIter =
3942 MRI.getType(Reg: Root->getOperand(i: 0).getReg()).getSizeInBytes() - 1;
3943 for (unsigned Iter = 0; Iter < MaxIter; ++Iter) {
3944 if (Ors.empty())
3945 break;
3946 const MachineInstr *Curr = Ors.pop_back_val();
3947 Register OrLHS = Curr->getOperand(i: 1).getReg();
3948 Register OrRHS = Curr->getOperand(i: 2).getReg();
3949
3950 // In the combine, we want to elimate the entire tree.
3951 if (!MRI.hasOneNonDBGUse(RegNo: OrLHS) || !MRI.hasOneNonDBGUse(RegNo: OrRHS))
3952 return std::nullopt;
3953
3954 // If it's a G_OR, save it and continue to walk. If it's not, then it's
3955 // something that may be a load + arithmetic.
3956 if (const MachineInstr *Or = getOpcodeDef(Opcode: TargetOpcode::G_OR, Reg: OrLHS, MRI))
3957 Ors.push_back(Elt: Or);
3958 else
3959 RegsToVisit.push_back(Elt: OrLHS);
3960 if (const MachineInstr *Or = getOpcodeDef(Opcode: TargetOpcode::G_OR, Reg: OrRHS, MRI))
3961 Ors.push_back(Elt: Or);
3962 else
3963 RegsToVisit.push_back(Elt: OrRHS);
3964 }
3965
3966 // We're going to try and merge each register into a wider power-of-2 type,
3967 // so we ought to have an even number of registers.
3968 if (RegsToVisit.empty() || RegsToVisit.size() % 2 != 0)
3969 return std::nullopt;
3970 return RegsToVisit;
3971}
3972
3973/// Helper function for findLoadOffsetsForLoadOrCombine.
3974///
3975/// Check if \p Reg is the result of loading a \p MemSizeInBits wide value,
3976/// and then moving that value into a specific byte offset.
3977///
3978/// e.g. x[i] << 24
3979///
3980/// \returns The load instruction and the byte offset it is moved into.
3981static std::optional<std::pair<GZExtLoad *, int64_t>>
3982matchLoadAndBytePosition(Register Reg, unsigned MemSizeInBits,
3983 const MachineRegisterInfo &MRI) {
3984 assert(MRI.hasOneNonDBGUse(Reg) &&
3985 "Expected Reg to only have one non-debug use?");
3986 Register MaybeLoad;
3987 int64_t Shift;
3988 if (!mi_match(R: Reg, MRI,
3989 P: m_OneNonDBGUse(SP: m_GShl(L: m_Reg(R&: MaybeLoad), R: m_ICst(Cst&: Shift))))) {
3990 Shift = 0;
3991 MaybeLoad = Reg;
3992 }
3993
3994 if (Shift % MemSizeInBits != 0)
3995 return std::nullopt;
3996
3997 // TODO: Handle other types of loads.
3998 auto *Load = getOpcodeDef<GZExtLoad>(Reg: MaybeLoad, MRI);
3999 if (!Load)
4000 return std::nullopt;
4001
4002 if (!Load->isUnordered() || Load->getMemSizeInBits() != MemSizeInBits)
4003 return std::nullopt;
4004
4005 return std::make_pair(x&: Load, y: Shift / MemSizeInBits);
4006}
4007
4008std::optional<std::tuple<GZExtLoad *, int64_t, GZExtLoad *>>
4009CombinerHelper::findLoadOffsetsForLoadOrCombine(
4010 SmallDenseMap<int64_t, int64_t, 8> &MemOffset2Idx,
4011 const SmallVector<Register, 8> &RegsToVisit,
4012 const unsigned MemSizeInBits) const {
4013
4014 // Each load found for the pattern. There should be one for each RegsToVisit.
4015 SmallSetVector<const MachineInstr *, 8> Loads;
4016
4017 // The lowest index used in any load. (The lowest "i" for each x[i].)
4018 int64_t LowestIdx = INT64_MAX;
4019
4020 // The load which uses the lowest index.
4021 GZExtLoad *LowestIdxLoad = nullptr;
4022
4023 // Keeps track of the load indices we see. We shouldn't see any indices twice.
4024 SmallSet<int64_t, 8> SeenIdx;
4025
4026 // Ensure each load is in the same MBB.
4027 // TODO: Support multiple MachineBasicBlocks.
4028 MachineBasicBlock *MBB = nullptr;
4029 const MachineMemOperand *MMO = nullptr;
4030
4031 // Earliest instruction-order load in the pattern.
4032 GZExtLoad *EarliestLoad = nullptr;
4033
4034 // Latest instruction-order load in the pattern.
4035 GZExtLoad *LatestLoad = nullptr;
4036
4037 // Base pointer which every load should share.
4038 Register BasePtr;
4039
4040 // We want to find a load for each register. Each load should have some
4041 // appropriate bit twiddling arithmetic. During this loop, we will also keep
4042 // track of the load which uses the lowest index. Later, we will check if we
4043 // can use its pointer in the final, combined load.
4044 for (auto Reg : RegsToVisit) {
4045 // Find the load, and find the position that it will end up in (e.g. a
4046 // shifted) value.
4047 auto LoadAndPos = matchLoadAndBytePosition(Reg, MemSizeInBits, MRI);
4048 if (!LoadAndPos)
4049 return std::nullopt;
4050 GZExtLoad *Load;
4051 int64_t DstPos;
4052 std::tie(args&: Load, args&: DstPos) = *LoadAndPos;
4053
4054 // TODO: Handle multiple MachineBasicBlocks. Currently not handled because
4055 // it is difficult to check for stores/calls/etc between loads.
4056 MachineBasicBlock *LoadMBB = Load->getParent();
4057 if (!MBB)
4058 MBB = LoadMBB;
4059 if (LoadMBB != MBB)
4060 return std::nullopt;
4061
4062 // Make sure that the MachineMemOperands of every seen load are compatible.
4063 auto &LoadMMO = Load->getMMO();
4064 if (!MMO)
4065 MMO = &LoadMMO;
4066 if (MMO->getAddrSpace() != LoadMMO.getAddrSpace())
4067 return std::nullopt;
4068
4069 // Find out what the base pointer and index for the load is.
4070 Register LoadPtr;
4071 int64_t Idx;
4072 if (!mi_match(R: Load->getOperand(i: 1).getReg(), MRI,
4073 P: m_GPtrAdd(L: m_Reg(R&: LoadPtr), R: m_ICst(Cst&: Idx)))) {
4074 LoadPtr = Load->getOperand(i: 1).getReg();
4075 Idx = 0;
4076 }
4077
4078 // Don't combine things like a[i], a[i] -> a bigger load.
4079 if (!SeenIdx.insert(V: Idx).second)
4080 return std::nullopt;
4081
4082 // Every load must share the same base pointer; don't combine things like:
4083 //
4084 // a[i], b[i + 1] -> a bigger load.
4085 if (!BasePtr.isValid())
4086 BasePtr = LoadPtr;
4087 if (BasePtr != LoadPtr)
4088 return std::nullopt;
4089
4090 if (Idx < LowestIdx) {
4091 LowestIdx = Idx;
4092 LowestIdxLoad = Load;
4093 }
4094
4095 // Keep track of the byte offset that this load ends up at. If we have seen
4096 // the byte offset, then stop here. We do not want to combine:
4097 //
4098 // a[i] << 16, a[i + k] << 16 -> a bigger load.
4099 if (!MemOffset2Idx.try_emplace(Key: DstPos, Args&: Idx).second)
4100 return std::nullopt;
4101 Loads.insert(X: Load);
4102
4103 // Keep track of the position of the earliest/latest loads in the pattern.
4104 // We will check that there are no load fold barriers between them later
4105 // on.
4106 //
4107 // FIXME: Is there a better way to check for load fold barriers?
4108 if (!EarliestLoad || dominates(DefMI: *Load, UseMI: *EarliestLoad))
4109 EarliestLoad = Load;
4110 if (!LatestLoad || dominates(DefMI: *LatestLoad, UseMI: *Load))
4111 LatestLoad = Load;
4112 }
4113
4114 // We found a load for each register. Let's check if each load satisfies the
4115 // pattern.
4116 assert(Loads.size() == RegsToVisit.size() &&
4117 "Expected to find a load for each register?");
4118 assert(EarliestLoad != LatestLoad && EarliestLoad &&
4119 LatestLoad && "Expected at least two loads?");
4120
4121 // Check if there are any stores, calls, etc. between any of the loads. If
4122 // there are, then we can't safely perform the combine.
4123 //
4124 // MaxIter is chosen based off the (worst case) number of iterations it
4125 // typically takes to succeed in the LLVM test suite plus some padding.
4126 //
4127 // FIXME: Is there a better way to check for load fold barriers?
4128 const unsigned MaxIter = 20;
4129 unsigned Iter = 0;
4130 for (const auto &MI : instructionsWithoutDebug(It: EarliestLoad->getIterator(),
4131 End: LatestLoad->getIterator())) {
4132 if (Loads.count(key: &MI))
4133 continue;
4134 if (MI.isLoadFoldBarrier())
4135 return std::nullopt;
4136 if (Iter++ == MaxIter)
4137 return std::nullopt;
4138 }
4139
4140 return std::make_tuple(args&: LowestIdxLoad, args&: LowestIdx, args&: LatestLoad);
4141}
4142
4143bool CombinerHelper::matchLoadOrCombine(
4144 MachineInstr &MI,
4145 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
4146 assert(MI.getOpcode() == TargetOpcode::G_OR);
4147 MachineFunction &MF = *MI.getMF();
4148 // Assuming a little-endian target, transform:
4149 // s8 *a = ...
4150 // s32 val = a[0] | (a[1] << 8) | (a[2] << 16) | (a[3] << 24)
4151 // =>
4152 // s32 val = *((i32)a)
4153 //
4154 // s8 *a = ...
4155 // s32 val = (a[0] << 24) | (a[1] << 16) | (a[2] << 8) | a[3]
4156 // =>
4157 // s32 val = BSWAP(*((s32)a))
4158 Register Dst = MI.getOperand(i: 0).getReg();
4159 LLT Ty = MRI.getType(Reg: Dst);
4160 if (Ty.isVector())
4161 return false;
4162
4163 // We need to combine at least two loads into this type. Since the smallest
4164 // possible load is into a byte, we need at least a 16-bit wide type.
4165 const unsigned WideMemSizeInBits = Ty.getSizeInBits();
4166 if (WideMemSizeInBits < 16 || WideMemSizeInBits % 8 != 0)
4167 return false;
4168
4169 // Match a collection of non-OR instructions in the pattern.
4170 auto RegsToVisit = findCandidatesForLoadOrCombine(Root: &MI);
4171 if (!RegsToVisit)
4172 return false;
4173
4174 // We have a collection of non-OR instructions. Figure out how wide each of
4175 // the small loads should be based off of the number of potential loads we
4176 // found.
4177 const unsigned NarrowMemSizeInBits = WideMemSizeInBits / RegsToVisit->size();
4178 if (NarrowMemSizeInBits % 8 != 0)
4179 return false;
4180
4181 // Check if each register feeding into each OR is a load from the same
4182 // base pointer + some arithmetic.
4183 //
4184 // e.g. a[0], a[1] << 8, a[2] << 16, etc.
4185 //
4186 // Also verify that each of these ends up putting a[i] into the same memory
4187 // offset as a load into a wide type would.
4188 SmallDenseMap<int64_t, int64_t, 8> MemOffset2Idx;
4189 GZExtLoad *LowestIdxLoad, *LatestLoad;
4190 int64_t LowestIdx;
4191 auto MaybeLoadInfo = findLoadOffsetsForLoadOrCombine(
4192 MemOffset2Idx, RegsToVisit: *RegsToVisit, MemSizeInBits: NarrowMemSizeInBits);
4193 if (!MaybeLoadInfo)
4194 return false;
4195 std::tie(args&: LowestIdxLoad, args&: LowestIdx, args&: LatestLoad) = *MaybeLoadInfo;
4196
4197 // We have a bunch of loads being OR'd together. Using the addresses + offsets
4198 // we found before, check if this corresponds to a big or little endian byte
4199 // pattern. If it does, then we can represent it using a load + possibly a
4200 // BSWAP.
4201 bool IsBigEndianTarget = MF.getDataLayout().isBigEndian();
4202 std::optional<bool> IsBigEndian = isBigEndian(MemOffset2Idx, LowestIdx);
4203 if (!IsBigEndian)
4204 return false;
4205 bool NeedsBSwap = IsBigEndianTarget != *IsBigEndian;
4206 if (NeedsBSwap && !isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_BSWAP, {Ty}}))
4207 return false;
4208
4209 // Make sure that the load from the lowest index produces offset 0 in the
4210 // final value.
4211 //
4212 // This ensures that we won't combine something like this:
4213 //
4214 // load x[i] -> byte 2
4215 // load x[i+1] -> byte 0 ---> wide_load x[i]
4216 // load x[i+2] -> byte 1
4217 const unsigned NumLoadsInTy = WideMemSizeInBits / NarrowMemSizeInBits;
4218 const unsigned ZeroByteOffset =
4219 *IsBigEndian
4220 ? bigEndianByteAt(ByteWidth: NumLoadsInTy, I: 0)
4221 : littleEndianByteAt(ByteWidth: NumLoadsInTy, I: 0);
4222 auto ZeroOffsetIdx = MemOffset2Idx.find(Val: ZeroByteOffset);
4223 if (ZeroOffsetIdx == MemOffset2Idx.end() ||
4224 ZeroOffsetIdx->second != LowestIdx)
4225 return false;
4226
4227 // We wil reuse the pointer from the load which ends up at byte offset 0. It
4228 // may not use index 0.
4229 Register Ptr = LowestIdxLoad->getPointerReg();
4230 const MachineMemOperand &MMO = LowestIdxLoad->getMMO();
4231 LegalityQuery::MemDesc MMDesc(MMO);
4232 MMDesc.MemoryTy = Ty;
4233 if (!isLegalOrBeforeLegalizer(
4234 Query: {TargetOpcode::G_LOAD, {Ty, MRI.getType(Reg: Ptr)}, {MMDesc}}))
4235 return false;
4236 auto PtrInfo = MMO.getPointerInfo();
4237 auto *NewMMO = MF.getMachineMemOperand(MMO: &MMO, PtrInfo, Size: WideMemSizeInBits / 8);
4238
4239 // Load must be allowed and fast on the target.
4240 LLVMContext &C = MF.getFunction().getContext();
4241 auto &DL = MF.getDataLayout();
4242 unsigned Fast = 0;
4243 if (!getTargetLowering().allowsMemoryAccess(Context&: C, DL, Ty, MMO: *NewMMO, Fast: &Fast) ||
4244 !Fast)
4245 return false;
4246
4247 MatchInfo = [=](MachineIRBuilder &MIB) {
4248 MIB.setInstrAndDebugLoc(*LatestLoad);
4249 Register LoadDst = NeedsBSwap ? MRI.cloneVirtualRegister(VReg: Dst) : Dst;
4250 MIB.buildLoad(Res: LoadDst, Addr: Ptr, MMO&: *NewMMO);
4251 if (NeedsBSwap)
4252 MIB.buildBSwap(Dst, Src0: LoadDst);
4253 };
4254 return true;
4255}
4256
4257bool CombinerHelper::matchExtendThroughPhis(MachineInstr &MI,
4258 MachineInstr *&ExtMI) const {
4259 auto &PHI = cast<GPhi>(Val&: MI);
4260 Register DstReg = PHI.getReg(Idx: 0);
4261
4262 // TODO: Extending a vector may be expensive, don't do this until heuristics
4263 // are better.
4264 if (MRI.getType(Reg: DstReg).isVector())
4265 return false;
4266
4267 // Try to match a phi, whose only use is an extend.
4268 if (!MRI.hasOneNonDBGUse(RegNo: DstReg))
4269 return false;
4270 ExtMI = &*MRI.use_instr_nodbg_begin(RegNo: DstReg);
4271 switch (ExtMI->getOpcode()) {
4272 case TargetOpcode::G_ANYEXT:
4273 return true; // G_ANYEXT is usually free.
4274 case TargetOpcode::G_ZEXT:
4275 case TargetOpcode::G_SEXT:
4276 break;
4277 default:
4278 return false;
4279 }
4280
4281 // If the target is likely to fold this extend away, don't propagate.
4282 if (Builder.getTII().isExtendLikelyToBeFolded(ExtMI&: *ExtMI, MRI))
4283 return false;
4284
4285 // We don't want to propagate the extends unless there's a good chance that
4286 // they'll be optimized in some way.
4287 // Collect the unique incoming values.
4288 SmallPtrSet<MachineInstr *, 4> InSrcs;
4289 for (unsigned I = 0; I < PHI.getNumIncomingValues(); ++I) {
4290 auto *DefMI = getDefIgnoringCopies(Reg: PHI.getIncomingValue(I), MRI);
4291 switch (DefMI->getOpcode()) {
4292 case TargetOpcode::G_LOAD:
4293 case TargetOpcode::G_TRUNC:
4294 case TargetOpcode::G_SEXT:
4295 case TargetOpcode::G_ZEXT:
4296 case TargetOpcode::G_ANYEXT:
4297 case TargetOpcode::G_CONSTANT:
4298 InSrcs.insert(Ptr: DefMI);
4299 // Don't try to propagate if there are too many places to create new
4300 // extends, chances are it'll increase code size.
4301 if (InSrcs.size() > 2)
4302 return false;
4303 break;
4304 default:
4305 return false;
4306 }
4307 }
4308 return true;
4309}
4310
4311void CombinerHelper::applyExtendThroughPhis(MachineInstr &MI,
4312 MachineInstr *&ExtMI) const {
4313 auto &PHI = cast<GPhi>(Val&: MI);
4314 Register DstReg = ExtMI->getOperand(i: 0).getReg();
4315 LLT ExtTy = MRI.getType(Reg: DstReg);
4316
4317 // Propagate the extension into the block of each incoming reg's block.
4318 // Use a SetVector here because PHIs can have duplicate edges, and we want
4319 // deterministic iteration order.
4320 SmallSetVector<MachineInstr *, 8> SrcMIs;
4321 SmallDenseMap<MachineInstr *, MachineInstr *, 8> OldToNewSrcMap;
4322 for (unsigned I = 0; I < PHI.getNumIncomingValues(); ++I) {
4323 auto SrcReg = PHI.getIncomingValue(I);
4324 MachineInstr *SrcMI;
4325 if (!mi_match(R: SrcReg, MRI, P: m_MInstr(MI&: SrcMI)))
4326 continue;
4327 if (!SrcMIs.insert(X: SrcMI))
4328 continue;
4329
4330 // Build an extend after each src inst.
4331 auto *MBB = SrcMI->getParent();
4332 MachineBasicBlock::iterator InsertPt = ++SrcMI->getIterator();
4333 if (InsertPt != MBB->end() && InsertPt->isPHI())
4334 InsertPt = MBB->getFirstNonPHI();
4335
4336 Builder.setInsertPt(MBB&: *SrcMI->getParent(), II: InsertPt);
4337 Builder.setDebugLoc(MI.getDebugLoc());
4338 auto NewExt = Builder.buildExtOrTrunc(ExtOpc: ExtMI->getOpcode(), Res: ExtTy, Op: SrcReg);
4339 OldToNewSrcMap[SrcMI] = NewExt;
4340 }
4341
4342 // Create a new phi with the extended inputs.
4343 Builder.setInstrAndDebugLoc(MI);
4344 auto NewPhi = Builder.buildInstrNoInsert(Opcode: TargetOpcode::G_PHI);
4345 NewPhi.addDef(RegNo: DstReg);
4346 for (const MachineOperand &MO : llvm::drop_begin(RangeOrContainer: MI.operands())) {
4347 if (!MO.isReg()) {
4348 NewPhi.addMBB(MBB: MO.getMBB());
4349 continue;
4350 }
4351 auto *NewSrc = OldToNewSrcMap[MRI.getVRegDef(Reg: MO.getReg())];
4352 NewPhi.addUse(RegNo: NewSrc->getOperand(i: 0).getReg());
4353 }
4354 Builder.insertInstr(MIB: NewPhi);
4355 ExtMI->eraseFromParent();
4356}
4357
4358bool CombinerHelper::matchExtractVecEltBuildVec(MachineInstr &MI,
4359 Register &Reg) const {
4360 assert(MI.getOpcode() == TargetOpcode::G_EXTRACT_VECTOR_ELT);
4361 // If we have a constant index, look for a G_BUILD_VECTOR source
4362 // and find the source register that the index maps to.
4363 Register SrcVec = MI.getOperand(i: 1).getReg();
4364 LLT SrcTy = MRI.getType(Reg: SrcVec);
4365 if (SrcTy.isScalableVector())
4366 return false;
4367
4368 auto Cst = getIConstantVRegValWithLookThrough(VReg: MI.getOperand(i: 2).getReg(), MRI);
4369 if (!Cst || Cst->Value.getZExtValue() >= SrcTy.getNumElements())
4370 return false;
4371
4372 unsigned VecIdx = Cst->Value.getZExtValue();
4373
4374 // Check if we have a build_vector or build_vector_trunc with an optional
4375 // trunc in front.
4376 MachineInstr *SrcVecMI;
4377 Register TruncSrc;
4378 if (mi_match(R: SrcVec, MRI, P: m_GTrunc(Src: m_Reg(R&: TruncSrc)))) {
4379 if (!mi_match(R: TruncSrc, MRI, P: m_MInstr(MI&: SrcVecMI)))
4380 return false;
4381 } else if (!mi_match(R: SrcVec, MRI, P: m_MInstr(MI&: SrcVecMI)))
4382 return false;
4383
4384 if (SrcVecMI->getOpcode() != TargetOpcode::G_BUILD_VECTOR &&
4385 SrcVecMI->getOpcode() != TargetOpcode::G_BUILD_VECTOR_TRUNC)
4386 return false;
4387
4388 EVT Ty(getMVTForLLT(Ty: SrcTy));
4389 if (!MRI.hasOneNonDBGUse(RegNo: SrcVec) &&
4390 !getTargetLowering().aggressivelyPreferBuildVectorSources(VecVT: Ty))
4391 return false;
4392
4393 Reg = SrcVecMI->getOperand(i: VecIdx + 1).getReg();
4394 return true;
4395}
4396
4397void CombinerHelper::applyExtractVecEltBuildVec(MachineInstr &MI,
4398 Register &Reg) const {
4399 // Check the type of the register, since it may have come from a
4400 // G_BUILD_VECTOR_TRUNC.
4401 LLT ScalarTy = MRI.getType(Reg);
4402 Register DstReg = MI.getOperand(i: 0).getReg();
4403 LLT DstTy = MRI.getType(Reg: DstReg);
4404
4405 if (ScalarTy != DstTy) {
4406 assert(ScalarTy.getSizeInBits() > DstTy.getSizeInBits());
4407 Builder.buildTrunc(Res: DstReg, Op: Reg);
4408 MI.eraseFromParent();
4409 return;
4410 }
4411 replaceSingleDefInstWithReg(MI, Replacement: Reg);
4412}
4413
4414bool CombinerHelper::matchExtractAllEltsFromBuildVector(
4415 MachineInstr &MI,
4416 SmallVectorImpl<std::pair<Register, MachineInstr *>> &SrcDstPairs) const {
4417 assert(MI.getOpcode() == TargetOpcode::G_BUILD_VECTOR);
4418 // This combine tries to find build_vector's which have every source element
4419 // extracted using G_EXTRACT_VECTOR_ELT. This can happen when transforms like
4420 // the masked load scalarization is run late in the pipeline. There's already
4421 // a combine for a similar pattern starting from the extract, but that
4422 // doesn't attempt to do it if there are multiple uses of the build_vector,
4423 // which in this case is true. Starting the combine from the build_vector
4424 // feels more natural than trying to find sibling nodes of extracts.
4425 // E.g.
4426 // %vec(<4 x s32>) = G_BUILD_VECTOR %s1(s32), %s2, %s3, %s4
4427 // %ext1 = G_EXTRACT_VECTOR_ELT %vec, 0
4428 // %ext2 = G_EXTRACT_VECTOR_ELT %vec, 1
4429 // %ext3 = G_EXTRACT_VECTOR_ELT %vec, 2
4430 // %ext4 = G_EXTRACT_VECTOR_ELT %vec, 3
4431 // ==>
4432 // replace ext{1,2,3,4} with %s{1,2,3,4}
4433
4434 Register DstReg = MI.getOperand(i: 0).getReg();
4435 LLT DstTy = MRI.getType(Reg: DstReg);
4436 unsigned NumElts = DstTy.getNumElements();
4437
4438 SmallBitVector ExtractedElts(NumElts);
4439 for (MachineInstr &II : MRI.use_nodbg_instructions(Reg: DstReg)) {
4440 if (II.getOpcode() != TargetOpcode::G_EXTRACT_VECTOR_ELT)
4441 return false;
4442 auto Cst = getIConstantVRegVal(VReg: II.getOperand(i: 2).getReg(), MRI);
4443 if (!Cst)
4444 return false;
4445 unsigned Idx = Cst->getZExtValue();
4446 if (Idx >= NumElts)
4447 return false; // Out of range.
4448 ExtractedElts.set(Idx);
4449 SrcDstPairs.emplace_back(
4450 Args: std::make_pair(x: MI.getOperand(i: Idx + 1).getReg(), y: &II));
4451 }
4452 // Match if every element was extracted.
4453 return ExtractedElts.all();
4454}
4455
4456void CombinerHelper::applyExtractAllEltsFromBuildVector(
4457 MachineInstr &MI,
4458 SmallVectorImpl<std::pair<Register, MachineInstr *>> &SrcDstPairs) const {
4459 assert(MI.getOpcode() == TargetOpcode::G_BUILD_VECTOR);
4460 for (auto &Pair : SrcDstPairs) {
4461 auto *ExtMI = Pair.second;
4462 replaceRegWith(MRI, FromReg: ExtMI->getOperand(i: 0).getReg(), ToReg: Pair.first);
4463 ExtMI->eraseFromParent();
4464 }
4465 MI.eraseFromParent();
4466}
4467
4468void CombinerHelper::applyBuildFn(
4469 MachineInstr &MI,
4470 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
4471 applyBuildFnNoErase(MI, MatchInfo);
4472 MI.eraseFromParent();
4473}
4474
4475void CombinerHelper::applyBuildFnNoErase(
4476 MachineInstr &MI,
4477 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
4478 MatchInfo(Builder);
4479}
4480
4481bool CombinerHelper::matchOrShiftToFunnelShift(MachineInstr &MI,
4482 bool AllowScalarConstants,
4483 BuildFnTy &MatchInfo) const {
4484 assert(MI.getOpcode() == TargetOpcode::G_OR);
4485
4486 Register Dst = MI.getOperand(i: 0).getReg();
4487 LLT Ty = MRI.getType(Reg: Dst);
4488 unsigned BitWidth = Ty.getScalarSizeInBits();
4489
4490 Register ShlSrc, ShlAmt, LShrSrc, LShrAmt, Amt;
4491 unsigned FshOpc = 0;
4492
4493 // Match (or (shl ...), (lshr ...)).
4494 if (!mi_match(R: Dst, MRI,
4495 // m_GOr() handles the commuted version as well.
4496 P: m_GOr(L: m_GShl(L: m_Reg(R&: ShlSrc), R: m_Reg(R&: ShlAmt)),
4497 R: m_GLShr(L: m_Reg(R&: LShrSrc), R: m_Reg(R&: LShrAmt)))))
4498 return false;
4499
4500 // Given constants C0 and C1 such that C0 + C1 is bit-width:
4501 // (or (shl x, C0), (lshr y, C1)) -> (fshl x, y, C0) or (fshr x, y, C1)
4502 int64_t CstShlAmt = 0, CstLShrAmt;
4503 if (mi_match(R: ShlAmt, MRI, P: m_ICstOrSplat(Cst&: CstShlAmt)) &&
4504 mi_match(R: LShrAmt, MRI, P: m_ICstOrSplat(Cst&: CstLShrAmt)) &&
4505 CstShlAmt + CstLShrAmt == BitWidth) {
4506 FshOpc = TargetOpcode::G_FSHR;
4507 Amt = LShrAmt;
4508 } else if (mi_match(R: LShrAmt, MRI,
4509 P: m_GSub(L: m_SpecificICstOrSplat(RequestedValue: BitWidth), R: m_Reg(R&: Amt))) &&
4510 ShlAmt == Amt) {
4511 // (or (shl x, amt), (lshr y, (sub bw, amt))) -> (fshl x, y, amt)
4512 FshOpc = TargetOpcode::G_FSHL;
4513 } else if (mi_match(R: ShlAmt, MRI,
4514 P: m_GSub(L: m_SpecificICstOrSplat(RequestedValue: BitWidth), R: m_Reg(R&: Amt))) &&
4515 LShrAmt == Amt) {
4516 // (or (shl x, (sub bw, amt)), (lshr y, amt)) -> (fshr x, y, amt)
4517 FshOpc = TargetOpcode::G_FSHR;
4518 } else {
4519 return false;
4520 }
4521
4522 LLT AmtTy = MRI.getType(Reg: Amt);
4523 if (!isLegalOrBeforeLegalizer(Query: {FshOpc, {Ty, AmtTy}}) &&
4524 (!AllowScalarConstants || CstShlAmt == 0 || !Ty.isScalar()))
4525 return false;
4526
4527 MatchInfo = [=](MachineIRBuilder &B) {
4528 B.buildInstr(Opc: FshOpc, DstOps: {Dst}, SrcOps: {ShlSrc, LShrSrc, Amt});
4529 };
4530 return true;
4531}
4532
4533/// Match an FSHL or FSHR that can be combined to a ROTR or ROTL rotate.
4534bool CombinerHelper::matchFunnelShiftToRotate(MachineInstr &MI) const {
4535 unsigned Opc = MI.getOpcode();
4536 assert(Opc == TargetOpcode::G_FSHL || Opc == TargetOpcode::G_FSHR);
4537 Register X = MI.getOperand(i: 1).getReg();
4538 Register Y = MI.getOperand(i: 2).getReg();
4539 if (X != Y)
4540 return false;
4541 unsigned RotateOpc =
4542 Opc == TargetOpcode::G_FSHL ? TargetOpcode::G_ROTL : TargetOpcode::G_ROTR;
4543 return isLegalOrBeforeLegalizer(Query: {RotateOpc, {MRI.getType(Reg: X), MRI.getType(Reg: Y)}});
4544}
4545
4546void CombinerHelper::applyFunnelShiftToRotate(MachineInstr &MI) const {
4547 unsigned Opc = MI.getOpcode();
4548 assert(Opc == TargetOpcode::G_FSHL || Opc == TargetOpcode::G_FSHR);
4549 bool IsFSHL = Opc == TargetOpcode::G_FSHL;
4550 Observer.changingInstr(MI);
4551 MI.setDesc(Builder.getTII().get(Opcode: IsFSHL ? TargetOpcode::G_ROTL
4552 : TargetOpcode::G_ROTR));
4553 MI.removeOperand(OpNo: 2);
4554 Observer.changedInstr(MI);
4555}
4556
4557// Fold (rot x, c) -> (rot x, c % BitSize)
4558bool CombinerHelper::matchRotateOutOfRange(MachineInstr &MI) const {
4559 assert(MI.getOpcode() == TargetOpcode::G_ROTL ||
4560 MI.getOpcode() == TargetOpcode::G_ROTR);
4561 unsigned Bitsize =
4562 MRI.getType(Reg: MI.getOperand(i: 0).getReg()).getScalarSizeInBits();
4563 Register AmtReg = MI.getOperand(i: 2).getReg();
4564 bool OutOfRange = false;
4565 auto MatchOutOfRange = [Bitsize, &OutOfRange](const Constant *C) {
4566 if (auto *CI = dyn_cast<ConstantInt>(Val: C))
4567 OutOfRange |= CI->getValue().uge(RHS: Bitsize);
4568 return true;
4569 };
4570 return matchUnaryPredicate(MRI, Reg: AmtReg, Match: MatchOutOfRange) && OutOfRange;
4571}
4572
4573void CombinerHelper::applyRotateOutOfRange(MachineInstr &MI) const {
4574 assert(MI.getOpcode() == TargetOpcode::G_ROTL ||
4575 MI.getOpcode() == TargetOpcode::G_ROTR);
4576 unsigned Bitsize =
4577 MRI.getType(Reg: MI.getOperand(i: 0).getReg()).getScalarSizeInBits();
4578 Register Amt = MI.getOperand(i: 2).getReg();
4579 LLT AmtTy = MRI.getType(Reg: Amt);
4580 auto Bits = Builder.buildConstant(Res: AmtTy, Val: Bitsize);
4581 Amt = Builder.buildURem(Dst: AmtTy, Src0: MI.getOperand(i: 2).getReg(), Src1: Bits).getReg(Idx: 0);
4582 Observer.changingInstr(MI);
4583 MI.getOperand(i: 2).setReg(Amt);
4584 Observer.changedInstr(MI);
4585}
4586
4587bool CombinerHelper::matchICmpToTrueFalseKnownBits(MachineInstr &MI,
4588 int64_t &MatchInfo) const {
4589 assert(MI.getOpcode() == TargetOpcode::G_ICMP);
4590 auto Pred = static_cast<CmpInst::Predicate>(MI.getOperand(i: 1).getPredicate());
4591
4592 // We want to avoid calling KnownBits on the LHS if possible, as this combine
4593 // has no filter and runs on every G_ICMP instruction. We can avoid calling
4594 // KnownBits on the LHS in two cases:
4595 //
4596 // - The RHS is unknown: Constants are always on RHS. If the RHS is unknown
4597 // we cannot do any transforms so we can safely bail out early.
4598 // - The RHS is zero: we don't need to know the LHS to do unsigned <0 and
4599 // >=0.
4600 auto KnownRHS = VT->getKnownBits(R: MI.getOperand(i: 3).getReg());
4601 if (KnownRHS.isUnknown())
4602 return false;
4603
4604 std::optional<bool> KnownVal;
4605 if (KnownRHS.isZero()) {
4606 // ? uge 0 -> always true
4607 // ? ult 0 -> always false
4608 if (Pred == CmpInst::ICMP_UGE)
4609 KnownVal = true;
4610 else if (Pred == CmpInst::ICMP_ULT)
4611 KnownVal = false;
4612 }
4613
4614 if (!KnownVal) {
4615 auto KnownLHS = VT->getKnownBits(R: MI.getOperand(i: 2).getReg());
4616 KnownVal = ICmpInst::compare(LHS: KnownLHS, RHS: KnownRHS, Pred);
4617 }
4618
4619 if (!KnownVal)
4620 return false;
4621 MatchInfo =
4622 *KnownVal
4623 ? getICmpTrueVal(TLI: getTargetLowering(),
4624 /*IsVector = */
4625 MRI.getType(Reg: MI.getOperand(i: 0).getReg()).isVector(),
4626 /* IsFP = */ false)
4627 : 0;
4628 return true;
4629}
4630
4631bool CombinerHelper::matchICmpToLHSKnownBits(
4632 MachineInstr &MI,
4633 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
4634 assert(MI.getOpcode() == TargetOpcode::G_ICMP);
4635 // Given:
4636 //
4637 // %x = G_WHATEVER (... x is known to be 0 or 1 ...)
4638 // %cmp = G_ICMP ne %x, 0
4639 //
4640 // Or:
4641 //
4642 // %x = G_WHATEVER (... x is known to be 0 or 1 ...)
4643 // %cmp = G_ICMP eq %x, 1
4644 //
4645 // We can replace %cmp with %x assuming true is 1 on the target.
4646 auto Pred = static_cast<CmpInst::Predicate>(MI.getOperand(i: 1).getPredicate());
4647 if (!CmpInst::isEquality(pred: Pred))
4648 return false;
4649 Register Dst = MI.getOperand(i: 0).getReg();
4650 LLT DstTy = MRI.getType(Reg: Dst);
4651 if (getICmpTrueVal(TLI: getTargetLowering(), IsVector: DstTy.isVector(),
4652 /* IsFP = */ false) != 1)
4653 return false;
4654 int64_t OneOrZero = Pred == CmpInst::ICMP_EQ;
4655 if (!mi_match(R: MI.getOperand(i: 3).getReg(), MRI, P: m_SpecificICst(RequestedValue: OneOrZero)))
4656 return false;
4657 Register LHS = MI.getOperand(i: 2).getReg();
4658 auto KnownLHS = VT->getKnownBits(R: LHS);
4659 if (KnownLHS.getMinValue() != 0 || KnownLHS.getMaxValue() != 1)
4660 return false;
4661 // Make sure replacing Dst with the LHS is a legal operation.
4662 LLT LHSTy = MRI.getType(Reg: LHS);
4663 unsigned LHSSize = LHSTy.getSizeInBits();
4664 unsigned DstSize = DstTy.getSizeInBits();
4665 unsigned Op = TargetOpcode::COPY;
4666 if (DstSize != LHSSize)
4667 Op = DstSize < LHSSize ? TargetOpcode::G_TRUNC : TargetOpcode::G_ZEXT;
4668 if (!isLegalOrBeforeLegalizer(Query: {Op, {DstTy, LHSTy}}))
4669 return false;
4670 MatchInfo = [=](MachineIRBuilder &B) { B.buildInstr(Opc: Op, DstOps: {Dst}, SrcOps: {LHS}); };
4671 return true;
4672}
4673
4674// Replace (and (or x, c1), c2) with (and x, c2) iff c1 & c2 == 0
4675bool CombinerHelper::matchAndOrDisjointMask(
4676 MachineInstr &MI,
4677 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
4678 assert(MI.getOpcode() == TargetOpcode::G_AND);
4679
4680 // Ignore vector types to simplify matching the two constants.
4681 // TODO: do this for vectors and scalars via a demanded bits analysis.
4682 LLT Ty = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
4683 if (Ty.isVector())
4684 return false;
4685
4686 Register Src;
4687 Register AndMaskReg;
4688 int64_t AndMaskBits;
4689 int64_t OrMaskBits;
4690 if (!mi_match(MI, MRI,
4691 P: m_GAnd(L: m_GOr(L: m_Reg(R&: Src), R: m_ICst(Cst&: OrMaskBits)),
4692 R: m_all_of(preds: m_ICst(Cst&: AndMaskBits), preds: m_Reg(R&: AndMaskReg)))))
4693 return false;
4694
4695 // Check if OrMask could turn on any bits in Src.
4696 if (AndMaskBits & OrMaskBits)
4697 return false;
4698
4699 MatchInfo = [=, &MI](MachineIRBuilder &B) {
4700 Observer.changingInstr(MI);
4701 // Canonicalize the result to have the constant on the RHS.
4702 if (MI.getOperand(i: 1).getReg() == AndMaskReg)
4703 MI.getOperand(i: 2).setReg(AndMaskReg);
4704 MI.getOperand(i: 1).setReg(Src);
4705 Observer.changedInstr(MI);
4706 };
4707 return true;
4708}
4709
4710/// Form a G_SBFX from a G_SEXT_INREG fed by a right shift.
4711bool CombinerHelper::matchBitfieldExtractFromSExtInReg(
4712 MachineInstr &MI,
4713 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
4714 assert(MI.getOpcode() == TargetOpcode::G_SEXT_INREG);
4715 Register Dst = MI.getOperand(i: 0).getReg();
4716 Register Src = MI.getOperand(i: 1).getReg();
4717 LLT Ty = MRI.getType(Reg: Src);
4718 LLT ExtractTy = getTargetLowering().getPreferredShiftAmountTy(ShiftValueTy: Ty);
4719 if (!LI || !LI->isLegalOrCustom(Query: {TargetOpcode::G_SBFX, {Ty, ExtractTy}}))
4720 return false;
4721 int64_t Width = MI.getOperand(i: 2).getImm();
4722 Register ShiftSrc;
4723 int64_t ShiftImm;
4724 if (!mi_match(
4725 R: Src, MRI,
4726 P: m_OneNonDBGUse(SP: m_any_of(preds: m_GAShr(L: m_Reg(R&: ShiftSrc), R: m_ICst(Cst&: ShiftImm)),
4727 preds: m_GLShr(L: m_Reg(R&: ShiftSrc), R: m_ICst(Cst&: ShiftImm))))))
4728 return false;
4729 if (ShiftImm < 0 || ShiftImm + Width > Ty.getScalarSizeInBits())
4730 return false;
4731
4732 MatchInfo = [=](MachineIRBuilder &B) {
4733 auto Cst1 = B.buildConstant(Res: ExtractTy, Val: ShiftImm);
4734 auto Cst2 = B.buildConstant(Res: ExtractTy, Val: Width);
4735 B.buildSbfx(Dst, Src: ShiftSrc, LSB: Cst1, Width: Cst2);
4736 };
4737 return true;
4738}
4739
4740/// Form a G_UBFX from "(a srl b) & mask", where b and mask are constants.
4741bool CombinerHelper::matchBitfieldExtractFromAnd(MachineInstr &MI,
4742 BuildFnTy &MatchInfo) const {
4743 GAnd *And = cast<GAnd>(Val: &MI);
4744 Register Dst = And->getReg(Idx: 0);
4745 LLT Ty = MRI.getType(Reg: Dst);
4746 LLT ExtractTy = getTargetLowering().getPreferredShiftAmountTy(ShiftValueTy: Ty);
4747 // Note that isLegalOrBeforeLegalizer is stricter and does not take custom
4748 // into account.
4749 if (LI && !LI->isLegalOrCustom(Query: {TargetOpcode::G_UBFX, {Ty, ExtractTy}}))
4750 return false;
4751
4752 int64_t AndImm, LSBImm;
4753 Register ShiftSrc;
4754 const unsigned Size = Ty.getScalarSizeInBits();
4755 if (!mi_match(R: And->getReg(Idx: 0), MRI,
4756 P: m_GAnd(L: m_OneNonDBGUse(SP: m_GLShr(L: m_Reg(R&: ShiftSrc), R: m_ICst(Cst&: LSBImm))),
4757 R: m_ICst(Cst&: AndImm))))
4758 return false;
4759
4760 // AndImm is sign-extended to 64 bits by m_ICst; restrict it to the operand
4761 // width so an all-ones mask (a redundant AND) is not misread as a wider mask.
4762 uint64_t MaybeMask = static_cast<uint64_t>(AndImm);
4763 if (Size < 64)
4764 MaybeMask &= maskTrailingOnes<uint64_t>(N: Size);
4765
4766 // The mask is a mask of the low bits iff imm & (imm+1) == 0.
4767 if (MaybeMask & (MaybeMask + 1))
4768 return false;
4769
4770 // LSB must fit within the register.
4771 if (static_cast<uint64_t>(LSBImm) >= Size)
4772 return false;
4773
4774 uint64_t Width = APInt(Size, MaybeMask).countr_one();
4775 // The extracted field [LSB, LSB+Width) must fit within the register.
4776 // Otherwise this is a redundant AND (e.g. an all-ones mask combined with a
4777 // non-zero shift) that is better handled by other combines, and would form
4778 // an out-of-range bitfield extract.
4779 if (static_cast<uint64_t>(LSBImm) + Width > Size)
4780 return false;
4781
4782 MatchInfo = [=](MachineIRBuilder &B) {
4783 auto WidthCst = B.buildConstant(Res: ExtractTy, Val: Width);
4784 auto LSBCst = B.buildConstant(Res: ExtractTy, Val: LSBImm);
4785 B.buildInstr(Opc: TargetOpcode::G_UBFX, DstOps: {Dst}, SrcOps: {ShiftSrc, LSBCst, WidthCst});
4786 };
4787 return true;
4788}
4789
4790bool CombinerHelper::matchBitfieldExtractFromShr(
4791 MachineInstr &MI,
4792 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
4793 const unsigned Opcode = MI.getOpcode();
4794 assert(Opcode == TargetOpcode::G_ASHR || Opcode == TargetOpcode::G_LSHR);
4795
4796 const Register Dst = MI.getOperand(i: 0).getReg();
4797
4798 const unsigned ExtrOpcode = Opcode == TargetOpcode::G_ASHR
4799 ? TargetOpcode::G_SBFX
4800 : TargetOpcode::G_UBFX;
4801
4802 // Check if the type we would use for the extract is legal
4803 LLT Ty = MRI.getType(Reg: Dst);
4804 LLT ExtractTy = getTargetLowering().getPreferredShiftAmountTy(ShiftValueTy: Ty);
4805 if (!LI || !LI->isLegalOrCustom(Query: {ExtrOpcode, {Ty, ExtractTy}}))
4806 return false;
4807
4808 Register ShlSrc;
4809 int64_t ShrAmt;
4810 int64_t ShlAmt;
4811 const unsigned Size = Ty.getScalarSizeInBits();
4812
4813 // Try to match shr (shl x, c1), c2
4814 if (!mi_match(R: Dst, MRI,
4815 P: m_BinOp(Opcode,
4816 L: m_OneNonDBGUse(SP: m_GShl(L: m_Reg(R&: ShlSrc), R: m_ICst(Cst&: ShlAmt))),
4817 R: m_ICst(Cst&: ShrAmt))))
4818 return false;
4819
4820 // Make sure that the shift sizes can fit a bitfield extract
4821 if (ShlAmt < 0 || ShlAmt > ShrAmt || ShrAmt >= Size)
4822 return false;
4823
4824 // Skip this combine if the G_SEXT_INREG combine could handle it
4825 if (Opcode == TargetOpcode::G_ASHR && ShlAmt == ShrAmt)
4826 return false;
4827
4828 // Calculate start position and width of the extract
4829 const int64_t Pos = ShrAmt - ShlAmt;
4830 const int64_t Width = Size - ShrAmt;
4831
4832 MatchInfo = [=](MachineIRBuilder &B) {
4833 auto WidthCst = B.buildConstant(Res: ExtractTy, Val: Width);
4834 auto PosCst = B.buildConstant(Res: ExtractTy, Val: Pos);
4835 B.buildInstr(Opc: ExtrOpcode, DstOps: {Dst}, SrcOps: {ShlSrc, PosCst, WidthCst});
4836 };
4837 return true;
4838}
4839
4840bool CombinerHelper::matchBitfieldExtractFromShrAnd(
4841 MachineInstr &MI,
4842 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
4843 const unsigned Opcode = MI.getOpcode();
4844 assert(Opcode == TargetOpcode::G_LSHR || Opcode == TargetOpcode::G_ASHR);
4845
4846 const Register Dst = MI.getOperand(i: 0).getReg();
4847 LLT Ty = MRI.getType(Reg: Dst);
4848 LLT ExtractTy = getTargetLowering().getPreferredShiftAmountTy(ShiftValueTy: Ty);
4849 if (LI && !LI->isLegalOrCustom(Query: {TargetOpcode::G_UBFX, {Ty, ExtractTy}}))
4850 return false;
4851
4852 // Try to match shr (and x, c1), c2
4853 Register AndSrc;
4854 int64_t ShrAmt;
4855 int64_t SMask;
4856 if (!mi_match(R: Dst, MRI,
4857 P: m_BinOp(Opcode,
4858 L: m_OneNonDBGUse(SP: m_GAnd(L: m_Reg(R&: AndSrc), R: m_ICst(Cst&: SMask))),
4859 R: m_ICst(Cst&: ShrAmt))))
4860 return false;
4861
4862 const unsigned Size = Ty.getScalarSizeInBits();
4863 if (ShrAmt < 0 || ShrAmt >= Size)
4864 return false;
4865
4866 // If the shift subsumes the mask, emit the 0 directly.
4867 if (0 == (SMask >> ShrAmt)) {
4868 MatchInfo = [=](MachineIRBuilder &B) {
4869 B.buildConstant(Res: Dst, Val: 0);
4870 };
4871 return true;
4872 }
4873
4874 // Check that ubfx can do the extraction, with no holes in the mask.
4875 uint64_t UMask = SMask;
4876 UMask |= maskTrailingOnes<uint64_t>(N: ShrAmt);
4877 UMask &= maskTrailingOnes<uint64_t>(N: Size);
4878 if (!isMask_64(Value: UMask))
4879 return false;
4880
4881 // Calculate start position and width of the extract.
4882 const int64_t Pos = ShrAmt;
4883 const int64_t Width = llvm::countr_one(Value: UMask) - ShrAmt;
4884
4885 // It's preferable to keep the shift, rather than form G_SBFX.
4886 // TODO: remove the G_AND via demanded bits analysis.
4887 if (Opcode == TargetOpcode::G_ASHR && Width + ShrAmt == Size)
4888 return false;
4889
4890 MatchInfo = [=](MachineIRBuilder &B) {
4891 auto WidthCst = B.buildConstant(Res: ExtractTy, Val: Width);
4892 auto PosCst = B.buildConstant(Res: ExtractTy, Val: Pos);
4893 B.buildInstr(Opc: TargetOpcode::G_UBFX, DstOps: {Dst}, SrcOps: {AndSrc, PosCst, WidthCst});
4894 };
4895 return true;
4896}
4897
4898bool CombinerHelper::reassociationCanBreakAddressingModePattern(
4899 MachineInstr &MI) const {
4900 auto &PtrAdd = cast<GPtrAdd>(Val&: MI);
4901
4902 Register Src1Reg = PtrAdd.getBaseReg();
4903 auto *Src1Def = getOpcodeDef<GPtrAdd>(Reg: Src1Reg, MRI);
4904 if (!Src1Def)
4905 return false;
4906
4907 Register Src2Reg = PtrAdd.getOffsetReg();
4908
4909 if (MRI.hasOneNonDBGUse(RegNo: Src1Reg))
4910 return false;
4911
4912 auto C1 = getIConstantVRegVal(VReg: Src1Def->getOffsetReg(), MRI);
4913 if (!C1)
4914 return false;
4915 auto C2 = getIConstantVRegVal(VReg: Src2Reg, MRI);
4916 if (!C2)
4917 return false;
4918
4919 const APInt &C1APIntVal = *C1;
4920 const APInt &C2APIntVal = *C2;
4921 const int64_t CombinedValue = (C1APIntVal + C2APIntVal).getSExtValue();
4922
4923 for (auto &UseMI : MRI.use_nodbg_instructions(Reg: PtrAdd.getReg(Idx: 0))) {
4924 // This combine may end up running before ptrtoint/inttoptr combines
4925 // manage to eliminate redundant conversions, so try to look through them.
4926 MachineInstr *ConvUseMI = &UseMI;
4927 unsigned ConvUseOpc = ConvUseMI->getOpcode();
4928 while (ConvUseOpc == TargetOpcode::G_INTTOPTR ||
4929 ConvUseOpc == TargetOpcode::G_PTRTOINT) {
4930 Register DefReg = ConvUseMI->getOperand(i: 0).getReg();
4931 if (!MRI.hasOneNonDBGUse(RegNo: DefReg))
4932 break;
4933 ConvUseMI = &*MRI.use_instr_nodbg_begin(RegNo: DefReg);
4934 ConvUseOpc = ConvUseMI->getOpcode();
4935 }
4936 auto *LdStMI = dyn_cast<GLoadStore>(Val: ConvUseMI);
4937 if (!LdStMI)
4938 continue;
4939 // Is x[offset2] already not a legal addressing mode? If so then
4940 // reassociating the constants breaks nothing (we test offset2 because
4941 // that's the one we hope to fold into the load or store).
4942 TargetLoweringBase::AddrMode AM;
4943 AM.HasBaseReg = true;
4944 AM.BaseOffs = C2APIntVal.getSExtValue();
4945 unsigned AS = MRI.getType(Reg: LdStMI->getPointerReg()).getAddressSpace();
4946 Type *AccessTy = getTypeForLLT(Ty: LdStMI->getMMO().getMemoryType(),
4947 C&: PtrAdd.getMF()->getFunction().getContext());
4948 const auto &TLI = *PtrAdd.getMF()->getSubtarget().getTargetLowering();
4949 if (!TLI.isLegalAddressingMode(DL: PtrAdd.getMF()->getDataLayout(), AM,
4950 Ty: AccessTy, AddrSpace: AS))
4951 continue;
4952
4953 // Would x[offset1+offset2] still be a legal addressing mode?
4954 AM.BaseOffs = CombinedValue;
4955 if (!TLI.isLegalAddressingMode(DL: PtrAdd.getMF()->getDataLayout(), AM,
4956 Ty: AccessTy, AddrSpace: AS))
4957 return true;
4958 }
4959
4960 return false;
4961}
4962
4963bool CombinerHelper::matchReassocConstantInnerRHS(GPtrAdd &MI,
4964 MachineInstr *RHS,
4965 BuildFnTy &MatchInfo) const {
4966 // G_PTR_ADD(BASE, G_ADD(X, C)) -> G_PTR_ADD(G_PTR_ADD(BASE, X), C)
4967 Register Src1Reg = MI.getOperand(i: 1).getReg();
4968 if (RHS->getOpcode() != TargetOpcode::G_ADD)
4969 return false;
4970 auto C2 = getIConstantVRegVal(VReg: RHS->getOperand(i: 2).getReg(), MRI);
4971 if (!C2)
4972 return false;
4973
4974 // If both additions are nuw, the reassociated additions are also nuw.
4975 // If the original G_PTR_ADD is additionally nusw, X and C are both not
4976 // negative, so BASE+X is between BASE and BASE+(X+C). The new G_PTR_ADDs are
4977 // therefore also nusw.
4978 // If the original G_PTR_ADD is additionally inbounds (which implies nusw),
4979 // the new G_PTR_ADDs are then also inbounds.
4980 unsigned PtrAddFlags = MI.getFlags();
4981 unsigned AddFlags = RHS->getFlags();
4982 bool IsNoUWrap = PtrAddFlags & AddFlags & MachineInstr::MIFlag::NoUWrap;
4983 bool IsNoUSWrap = IsNoUWrap && (PtrAddFlags & MachineInstr::MIFlag::NoUSWrap);
4984 bool IsInBounds = IsNoUWrap && (PtrAddFlags & MachineInstr::MIFlag::InBounds);
4985 unsigned Flags = 0;
4986 if (IsNoUWrap)
4987 Flags |= MachineInstr::MIFlag::NoUWrap;
4988 if (IsNoUSWrap)
4989 Flags |= MachineInstr::MIFlag::NoUSWrap;
4990 if (IsInBounds)
4991 Flags |= MachineInstr::MIFlag::InBounds;
4992
4993 MatchInfo = [=, &MI](MachineIRBuilder &B) {
4994 LLT PtrTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
4995
4996 auto NewBase =
4997 Builder.buildPtrAdd(Res: PtrTy, Op0: Src1Reg, Op1: RHS->getOperand(i: 1).getReg(), Flags);
4998 Observer.changingInstr(MI);
4999 MI.getOperand(i: 1).setReg(NewBase.getReg(Idx: 0));
5000 MI.getOperand(i: 2).setReg(RHS->getOperand(i: 2).getReg());
5001 MI.setFlags(Flags);
5002 Observer.changedInstr(MI);
5003 };
5004 return !reassociationCanBreakAddressingModePattern(MI);
5005}
5006
5007bool CombinerHelper::matchReassocConstantInnerLHS(GPtrAdd &MI,
5008 MachineInstr *LHS,
5009 MachineInstr *RHS,
5010 BuildFnTy &MatchInfo) const {
5011 // G_PTR_ADD (G_PTR_ADD X, C), Y) -> (G_PTR_ADD (G_PTR_ADD(X, Y), C)
5012 // if and only if (G_PTR_ADD X, C) has one use.
5013 Register LHSBase;
5014 std::optional<ValueAndVReg> LHSCstOff;
5015 if (!mi_match(R: MI.getBaseReg(), MRI,
5016 P: m_OneNonDBGUse(SP: m_GPtrAdd(L: m_Reg(R&: LHSBase), R: m_GCst(ValReg&: LHSCstOff)))))
5017 return false;
5018
5019 auto *LHSPtrAdd = cast<GPtrAdd>(Val: LHS);
5020
5021 // Reassociating nuw additions preserves nuw. If both original G_PTR_ADDs are
5022 // nuw and inbounds (which implies nusw), the offsets are both non-negative,
5023 // so the new G_PTR_ADDs are also inbounds.
5024 unsigned PtrAddFlags = MI.getFlags();
5025 unsigned LHSPtrAddFlags = LHSPtrAdd->getFlags();
5026 bool IsNoUWrap = PtrAddFlags & LHSPtrAddFlags & MachineInstr::MIFlag::NoUWrap;
5027 bool IsNoUSWrap = IsNoUWrap && (PtrAddFlags & LHSPtrAddFlags &
5028 MachineInstr::MIFlag::NoUSWrap);
5029 bool IsInBounds = IsNoUWrap && (PtrAddFlags & LHSPtrAddFlags &
5030 MachineInstr::MIFlag::InBounds);
5031 unsigned Flags = 0;
5032 if (IsNoUWrap)
5033 Flags |= MachineInstr::MIFlag::NoUWrap;
5034 if (IsNoUSWrap)
5035 Flags |= MachineInstr::MIFlag::NoUSWrap;
5036 if (IsInBounds)
5037 Flags |= MachineInstr::MIFlag::InBounds;
5038
5039 MatchInfo = [=, &MI](MachineIRBuilder &B) {
5040 // When we change LHSPtrAdd's offset register we might cause it to use a reg
5041 // before its def. Sink the instruction so the outer PTR_ADD to ensure this
5042 // doesn't happen.
5043 LHSPtrAdd->moveBefore(MovePos: &MI);
5044 Register RHSReg = MI.getOffsetReg();
5045 // set VReg will cause type mismatch if it comes from extend/trunc
5046 auto NewCst = B.buildConstant(Res: MRI.getType(Reg: RHSReg), Val: LHSCstOff->Value);
5047 Observer.changingInstr(MI);
5048 MI.getOperand(i: 2).setReg(NewCst.getReg(Idx: 0));
5049 MI.setFlags(Flags);
5050 Observer.changedInstr(MI);
5051 Observer.changingInstr(MI&: *LHSPtrAdd);
5052 LHSPtrAdd->getOperand(i: 2).setReg(RHSReg);
5053 LHSPtrAdd->setFlags(Flags);
5054 Observer.changedInstr(MI&: *LHSPtrAdd);
5055 };
5056 return !reassociationCanBreakAddressingModePattern(MI);
5057}
5058
5059bool CombinerHelper::matchReassocFoldConstantsInSubTree(
5060 GPtrAdd &MI, MachineInstr *LHS, MachineInstr *RHS,
5061 BuildFnTy &MatchInfo) const {
5062 // G_PTR_ADD(G_PTR_ADD(BASE, C1), C2) -> G_PTR_ADD(BASE, C1+C2)
5063 auto *LHSPtrAdd = dyn_cast<GPtrAdd>(Val: LHS);
5064 if (!LHSPtrAdd)
5065 return false;
5066
5067 Register Src2Reg = MI.getOperand(i: 2).getReg();
5068 Register LHSSrc1 = LHSPtrAdd->getBaseReg();
5069 Register LHSSrc2 = LHSPtrAdd->getOffsetReg();
5070 auto C1 = getIConstantVRegVal(VReg: LHSSrc2, MRI);
5071 if (!C1)
5072 return false;
5073 auto C2 = getIConstantVRegVal(VReg: Src2Reg, MRI);
5074 if (!C2)
5075 return false;
5076
5077 // Reassociating nuw additions preserves nuw. If both original G_PTR_ADDs are
5078 // inbounds, reaching the same result in one G_PTR_ADD is also inbounds.
5079 // The nusw constraints are satisfied because imm1+imm2 cannot exceed the
5080 // largest signed integer that fits into the index type, which is the maximum
5081 // size of allocated objects according to the IR Language Reference.
5082 unsigned PtrAddFlags = MI.getFlags();
5083 unsigned LHSPtrAddFlags = LHSPtrAdd->getFlags();
5084 bool IsNoUWrap = PtrAddFlags & LHSPtrAddFlags & MachineInstr::MIFlag::NoUWrap;
5085 bool IsInBounds =
5086 PtrAddFlags & LHSPtrAddFlags & MachineInstr::MIFlag::InBounds;
5087 unsigned Flags = 0;
5088 if (IsNoUWrap)
5089 Flags |= MachineInstr::MIFlag::NoUWrap;
5090 if (IsInBounds) {
5091 Flags |= MachineInstr::MIFlag::InBounds;
5092 Flags |= MachineInstr::MIFlag::NoUSWrap;
5093 }
5094
5095 MatchInfo = [=, &MI](MachineIRBuilder &B) {
5096 auto NewCst = B.buildConstant(Res: MRI.getType(Reg: Src2Reg), Val: *C1 + *C2);
5097 Observer.changingInstr(MI);
5098 MI.getOperand(i: 1).setReg(LHSSrc1);
5099 MI.getOperand(i: 2).setReg(NewCst.getReg(Idx: 0));
5100 MI.setFlags(Flags);
5101 Observer.changedInstr(MI);
5102 };
5103 return !reassociationCanBreakAddressingModePattern(MI);
5104}
5105
5106bool CombinerHelper::matchReassocPtrAdd(MachineInstr &MI,
5107 BuildFnTy &MatchInfo) const {
5108 auto &PtrAdd = cast<GPtrAdd>(Val&: MI);
5109 // We're trying to match a few pointer computation patterns here for
5110 // re-association opportunities.
5111 // 1) Isolating a constant operand to be on the RHS, e.g.:
5112 // G_PTR_ADD(BASE, G_ADD(X, C)) -> G_PTR_ADD(G_PTR_ADD(BASE, X), C)
5113 //
5114 // 2) Folding two constants in each sub-tree as long as such folding
5115 // doesn't break a legal addressing mode.
5116 // G_PTR_ADD(G_PTR_ADD(BASE, C1), C2) -> G_PTR_ADD(BASE, C1+C2)
5117 //
5118 // 3) Move a constant from the LHS of an inner op to the RHS of the outer.
5119 // G_PTR_ADD (G_PTR_ADD X, C), Y) -> G_PTR_ADD (G_PTR_ADD(X, Y), C)
5120 // iif (G_PTR_ADD X, C) has one use.
5121 MachineInstr *LHS, *RHS;
5122 if (!mi_match(R: PtrAdd.getBaseReg(), MRI, P: m_MInstr(MI&: LHS)) ||
5123 !mi_match(R: PtrAdd.getOffsetReg(), MRI, P: m_MInstr(MI&: RHS)))
5124 return false;
5125
5126 // Try to match example 2.
5127 if (matchReassocFoldConstantsInSubTree(MI&: PtrAdd, LHS, RHS, MatchInfo))
5128 return true;
5129
5130 // Try to match example 3.
5131 if (matchReassocConstantInnerLHS(MI&: PtrAdd, LHS, RHS, MatchInfo))
5132 return true;
5133
5134 // Try to match example 1.
5135 if (matchReassocConstantInnerRHS(MI&: PtrAdd, RHS, MatchInfo))
5136 return true;
5137
5138 return false;
5139}
5140bool CombinerHelper::tryReassocBinOp(unsigned Opc, Register DstReg,
5141 Register OpLHS, Register OpRHS,
5142 BuildFnTy &MatchInfo) const {
5143 LLT OpRHSTy = MRI.getType(Reg: OpRHS);
5144 MachineInstr *OpLHSDef;
5145 if (!mi_match(R: OpLHS, MRI, P: m_MInstr(MI&: OpLHSDef)) || OpLHSDef->getOpcode() != Opc)
5146 return false;
5147
5148 Register OpLHSLHS = OpLHSDef->getOperand(i: 1).getReg();
5149 Register OpLHSRHS = OpLHSDef->getOperand(i: 2).getReg();
5150
5151 // If the inner op is (X op C), pull the constant out so it can be folded with
5152 // other constants in the expression tree. Folding is not guaranteed so we
5153 // might have (C1 op C2). In that case do not pull a constant out because it
5154 // won't help and can lead to infinite loops.
5155 if (isConstantOrConstantSplatVector(Def: OpLHSRHS, MRI) &&
5156 !isConstantOrConstantSplatVector(Def: OpLHSLHS, MRI)) {
5157 if (isConstantOrConstantSplatVector(Def: OpRHS, MRI)) {
5158 // (Opc (Opc X, C1), C2) -> (Opc X, (Opc C1, C2))
5159 MatchInfo = [=](MachineIRBuilder &B) {
5160 auto NewCst = B.buildInstr(Opc, DstOps: {OpRHSTy}, SrcOps: {OpLHSRHS, OpRHS});
5161 B.buildInstr(Opc, DstOps: {DstReg}, SrcOps: {OpLHSLHS, NewCst});
5162 };
5163 return true;
5164 }
5165 if (getTargetLowering().isReassocProfitable(MRI, N0: OpLHS, N1: OpRHS)) {
5166 // Reassociate: (op (op x, c1), y) -> (op (op x, y), c1)
5167 // iff (op x, c1) has one use
5168 MatchInfo = [=](MachineIRBuilder &B) {
5169 auto NewLHSLHS = B.buildInstr(Opc, DstOps: {OpRHSTy}, SrcOps: {OpLHSLHS, OpRHS});
5170 B.buildInstr(Opc, DstOps: {DstReg}, SrcOps: {NewLHSLHS, OpLHSRHS});
5171 };
5172 return true;
5173 }
5174 }
5175
5176 return false;
5177}
5178
5179bool CombinerHelper::matchReassocCommBinOp(MachineInstr &MI,
5180 BuildFnTy &MatchInfo) const {
5181 // We don't check if the reassociation will break a legal addressing mode
5182 // here since pointer arithmetic is handled by G_PTR_ADD.
5183 unsigned Opc = MI.getOpcode();
5184 Register DstReg = MI.getOperand(i: 0).getReg();
5185 Register LHSReg = MI.getOperand(i: 1).getReg();
5186 Register RHSReg = MI.getOperand(i: 2).getReg();
5187
5188 if (tryReassocBinOp(Opc, DstReg, OpLHS: LHSReg, OpRHS: RHSReg, MatchInfo))
5189 return true;
5190 if (tryReassocBinOp(Opc, DstReg, OpLHS: RHSReg, OpRHS: LHSReg, MatchInfo))
5191 return true;
5192 return false;
5193}
5194
5195bool CombinerHelper::matchConstantFoldCastOp(MachineInstr &MI,
5196 APInt &MatchInfo) const {
5197 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
5198 Register SrcOp = MI.getOperand(i: 1).getReg();
5199
5200 if (auto MaybeCst = ConstantFoldCastOp(Opcode: MI.getOpcode(), DstTy, Op0: SrcOp, MRI)) {
5201 MatchInfo = *MaybeCst;
5202 return true;
5203 }
5204
5205 return false;
5206}
5207
5208bool CombinerHelper::matchConstantFoldUnaryIntOp(MachineInstr &MI,
5209 BuildFnTy &MatchInfo) const {
5210 Register Dst = MI.getOperand(i: 0).getReg();
5211 auto Csts = ConstantFoldUnaryIntOp(Opcode: MI.getOpcode(), DstTy: MRI.getType(Reg: Dst),
5212 Src: MI.getOperand(i: 1).getReg(), MRI);
5213 if (Csts.empty())
5214 return false;
5215
5216 MatchInfo = [Dst, Csts = std::move(Csts)](MachineIRBuilder &B) {
5217 if (Csts.size() == 1)
5218 B.buildConstant(Res: Dst, Val: Csts[0]);
5219 else
5220 B.buildBuildVectorConstant(Res: Dst, Ops: Csts);
5221 };
5222 return true;
5223}
5224
5225bool CombinerHelper::matchConstantFoldBinOp(MachineInstr &MI,
5226 APInt &MatchInfo) const {
5227 Register Op1 = MI.getOperand(i: 1).getReg();
5228 Register Op2 = MI.getOperand(i: 2).getReg();
5229 auto MaybeCst = ConstantFoldBinOp(Opcode: MI.getOpcode(), Op1, Op2, MRI);
5230 if (!MaybeCst)
5231 return false;
5232 MatchInfo = *MaybeCst;
5233 return true;
5234}
5235
5236bool CombinerHelper::matchConstantFoldFPBinOp(MachineInstr &MI,
5237 ConstantFP *&MatchInfo) const {
5238 Register Op1 = MI.getOperand(i: 1).getReg();
5239 Register Op2 = MI.getOperand(i: 2).getReg();
5240 auto MaybeCst = ConstantFoldFPBinOp(Opcode: MI.getOpcode(), Op1, Op2, MRI);
5241 if (!MaybeCst)
5242 return false;
5243 MatchInfo =
5244 ConstantFP::get(Context&: MI.getMF()->getFunction().getContext(), V: *MaybeCst);
5245 return true;
5246}
5247
5248bool CombinerHelper::matchConstantFoldFMA(MachineInstr &MI,
5249 ConstantFP *&MatchInfo) const {
5250 assert(MI.getOpcode() == TargetOpcode::G_FMA ||
5251 MI.getOpcode() == TargetOpcode::G_FMAD);
5252 auto [_, Op1, Op2, Op3] = MI.getFirst4Regs();
5253
5254 const ConstantFP *Op3Cst = getConstantFPVRegVal(VReg: Op3, MRI);
5255 if (!Op3Cst)
5256 return false;
5257
5258 const ConstantFP *Op2Cst = getConstantFPVRegVal(VReg: Op2, MRI);
5259 if (!Op2Cst)
5260 return false;
5261
5262 const ConstantFP *Op1Cst = getConstantFPVRegVal(VReg: Op1, MRI);
5263 if (!Op1Cst)
5264 return false;
5265
5266 APFloat Op1F = Op1Cst->getValueAPF();
5267 Op1F.fusedMultiplyAdd(Multiplicand: Op2Cst->getValueAPF(), Addend: Op3Cst->getValueAPF(),
5268 RM: APFloat::rmNearestTiesToEven);
5269 MatchInfo = ConstantFP::get(Context&: MI.getMF()->getFunction().getContext(), V: Op1F);
5270 return true;
5271}
5272
5273bool CombinerHelper::matchNarrowBinopFeedingAnd(
5274 MachineInstr &MI,
5275 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
5276 // Look for a binop feeding into an AND with a mask:
5277 //
5278 // %add = G_ADD %lhs, %rhs
5279 // %and = G_AND %add, 000...11111111
5280 //
5281 // Check if it's possible to perform the binop at a narrower width and zext
5282 // back to the original width like so:
5283 //
5284 // %narrow_lhs = G_TRUNC %lhs
5285 // %narrow_rhs = G_TRUNC %rhs
5286 // %narrow_add = G_ADD %narrow_lhs, %narrow_rhs
5287 // %new_add = G_ZEXT %narrow_add
5288 // %and = G_AND %new_add, 000...11111111
5289 //
5290 // This can allow later combines to eliminate the G_AND if it turns out
5291 // that the mask is irrelevant.
5292 assert(MI.getOpcode() == TargetOpcode::G_AND);
5293 Register Dst = MI.getOperand(i: 0).getReg();
5294 Register AndLHS = MI.getOperand(i: 1).getReg();
5295 Register AndRHS = MI.getOperand(i: 2).getReg();
5296 LLT WideTy = MRI.getType(Reg: Dst);
5297
5298 // If the potential binop has more than one use, then it's possible that one
5299 // of those uses will need its full width.
5300 if (!WideTy.isScalar() || !MRI.hasOneNonDBGUse(RegNo: AndLHS))
5301 return false;
5302
5303 // Check if the LHS feeding the AND is impacted by the high bits that we're
5304 // masking out.
5305 //
5306 // e.g. for 64-bit x, y:
5307 //
5308 // add_64(x, y) & 65535 == zext(add_16(trunc(x), trunc(y))) & 65535
5309 MachineInstr *LHSInst = getDefIgnoringCopies(Reg: AndLHS, MRI);
5310 if (!LHSInst)
5311 return false;
5312 unsigned LHSOpc = LHSInst->getOpcode();
5313 switch (LHSOpc) {
5314 default:
5315 return false;
5316 case TargetOpcode::G_ADD:
5317 case TargetOpcode::G_SUB:
5318 case TargetOpcode::G_MUL:
5319 case TargetOpcode::G_AND:
5320 case TargetOpcode::G_OR:
5321 case TargetOpcode::G_XOR:
5322 break;
5323 }
5324
5325 // Find the mask on the RHS.
5326 auto Cst = getIConstantVRegValWithLookThrough(VReg: AndRHS, MRI);
5327 if (!Cst)
5328 return false;
5329 auto Mask = Cst->Value;
5330 if (!Mask.isMask())
5331 return false;
5332
5333 // No point in combining if there's nothing to truncate.
5334 unsigned NarrowWidth = Mask.countr_one();
5335 if (NarrowWidth == WideTy.getSizeInBits())
5336 return false;
5337 LLT NarrowTy = LLT::integer(SizeInBits: NarrowWidth);
5338
5339 // Check if adding the zext + truncates could be harmful.
5340 auto &MF = *MI.getMF();
5341 const auto &TLI = getTargetLowering();
5342 LLVMContext &Ctx = MF.getFunction().getContext();
5343 if (!TLI.isTruncateFree(FromTy: WideTy, ToTy: NarrowTy, Ctx) ||
5344 !TLI.isZExtFree(FromTy: NarrowTy, ToTy: WideTy, Ctx))
5345 return false;
5346 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_TRUNC, {NarrowTy, WideTy}}) ||
5347 !isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_ZEXT, {WideTy, NarrowTy}}))
5348 return false;
5349 Register BinOpLHS = LHSInst->getOperand(i: 1).getReg();
5350 Register BinOpRHS = LHSInst->getOperand(i: 2).getReg();
5351 MatchInfo = [=, &MI](MachineIRBuilder &B) {
5352 auto NarrowLHS = Builder.buildTrunc(Res: NarrowTy, Op: BinOpLHS);
5353 auto NarrowRHS = Builder.buildTrunc(Res: NarrowTy, Op: BinOpRHS);
5354 auto NarrowBinOp =
5355 Builder.buildInstr(Opc: LHSOpc, DstOps: {NarrowTy}, SrcOps: {NarrowLHS, NarrowRHS});
5356 auto Ext = Builder.buildZExt(Res: WideTy, Op: NarrowBinOp);
5357 Observer.changingInstr(MI);
5358 MI.getOperand(i: 1).setReg(Ext.getReg(Idx: 0));
5359 Observer.changedInstr(MI);
5360 };
5361 return true;
5362}
5363
5364bool CombinerHelper::matchMulOBy2(MachineInstr &MI,
5365 BuildFnTy &MatchInfo) const {
5366 unsigned Opc = MI.getOpcode();
5367 assert(Opc == TargetOpcode::G_UMULO || Opc == TargetOpcode::G_SMULO);
5368
5369 if (!mi_match(R: MI.getOperand(i: 3).getReg(), MRI, P: m_SpecificICstOrSplat(RequestedValue: 2)))
5370 return false;
5371
5372 MatchInfo = [=, &MI](MachineIRBuilder &B) {
5373 Observer.changingInstr(MI);
5374 unsigned NewOpc = Opc == TargetOpcode::G_UMULO ? TargetOpcode::G_UADDO
5375 : TargetOpcode::G_SADDO;
5376 MI.setDesc(Builder.getTII().get(Opcode: NewOpc));
5377 MI.getOperand(i: 3).setReg(MI.getOperand(i: 2).getReg());
5378 Observer.changedInstr(MI);
5379 };
5380 return true;
5381}
5382
5383bool CombinerHelper::matchMulOBy0(MachineInstr &MI,
5384 BuildFnTy &MatchInfo) const {
5385 // (G_*MULO x, 0) -> 0 + no carry out
5386 assert(MI.getOpcode() == TargetOpcode::G_UMULO ||
5387 MI.getOpcode() == TargetOpcode::G_SMULO);
5388 if (!mi_match(R: MI.getOperand(i: 3).getReg(), MRI, P: m_SpecificICstOrSplat(RequestedValue: 0)))
5389 return false;
5390 Register Dst = MI.getOperand(i: 0).getReg();
5391 Register Carry = MI.getOperand(i: 1).getReg();
5392 if (!isConstantLegalOrBeforeLegalizer(Ty: MRI.getType(Reg: Dst)) ||
5393 !isConstantLegalOrBeforeLegalizer(Ty: MRI.getType(Reg: Carry)))
5394 return false;
5395 MatchInfo = [=](MachineIRBuilder &B) {
5396 B.buildConstant(Res: Dst, Val: 0);
5397 B.buildConstant(Res: Carry, Val: 0);
5398 };
5399 return true;
5400}
5401
5402bool CombinerHelper::matchAddEToAddO(MachineInstr &MI,
5403 BuildFnTy &MatchInfo) const {
5404 // (G_*ADDE x, y, 0) -> (G_*ADDO x, y)
5405 // (G_*SUBE x, y, 0) -> (G_*SUBO x, y)
5406 assert(MI.getOpcode() == TargetOpcode::G_UADDE ||
5407 MI.getOpcode() == TargetOpcode::G_SADDE ||
5408 MI.getOpcode() == TargetOpcode::G_USUBE ||
5409 MI.getOpcode() == TargetOpcode::G_SSUBE);
5410 if (!mi_match(R: MI.getOperand(i: 4).getReg(), MRI, P: m_SpecificICstOrSplat(RequestedValue: 0)))
5411 return false;
5412 MatchInfo = [&](MachineIRBuilder &B) {
5413 unsigned NewOpcode;
5414 switch (MI.getOpcode()) {
5415 case TargetOpcode::G_UADDE:
5416 NewOpcode = TargetOpcode::G_UADDO;
5417 break;
5418 case TargetOpcode::G_SADDE:
5419 NewOpcode = TargetOpcode::G_SADDO;
5420 break;
5421 case TargetOpcode::G_USUBE:
5422 NewOpcode = TargetOpcode::G_USUBO;
5423 break;
5424 case TargetOpcode::G_SSUBE:
5425 NewOpcode = TargetOpcode::G_SSUBO;
5426 break;
5427 }
5428 Observer.changingInstr(MI);
5429 MI.setDesc(B.getTII().get(Opcode: NewOpcode));
5430 MI.removeOperand(OpNo: 4);
5431 Observer.changedInstr(MI);
5432 };
5433 return true;
5434}
5435
5436bool CombinerHelper::matchSubAddSameReg(MachineInstr &MI,
5437 BuildFnTy &MatchInfo) const {
5438 assert(MI.getOpcode() == TargetOpcode::G_SUB);
5439 Register Dst = MI.getOperand(i: 0).getReg();
5440 // (x + y) - z -> x (if y == z)
5441 // (x + y) - z -> y (if x == z)
5442 Register X, Y, Z;
5443 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)))) {
5444 Register ReplaceReg;
5445 int64_t CstX, CstY;
5446 if (Y == Z || (mi_match(R: Y, MRI, P: m_ICstOrSplat(Cst&: CstY)) &&
5447 mi_match(R: Z, MRI, P: m_SpecificICstOrSplat(RequestedValue: CstY))))
5448 ReplaceReg = X;
5449 else if (X == Z || (mi_match(R: X, MRI, P: m_ICstOrSplat(Cst&: CstX)) &&
5450 mi_match(R: Z, MRI, P: m_SpecificICstOrSplat(RequestedValue: CstX))))
5451 ReplaceReg = Y;
5452 if (ReplaceReg) {
5453 MatchInfo = [=](MachineIRBuilder &B) { B.buildCopy(Res: Dst, Op: ReplaceReg); };
5454 return true;
5455 }
5456 }
5457
5458 // x - (y + z) -> 0 - y (if x == z)
5459 // x - (y + z) -> 0 - z (if x == y)
5460 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))))) {
5461 Register ReplaceReg;
5462 int64_t CstX;
5463 if (X == Z || (mi_match(R: X, MRI, P: m_ICstOrSplat(Cst&: CstX)) &&
5464 mi_match(R: Z, MRI, P: m_SpecificICstOrSplat(RequestedValue: CstX))))
5465 ReplaceReg = Y;
5466 else if (X == Y || (mi_match(R: X, MRI, P: m_ICstOrSplat(Cst&: CstX)) &&
5467 mi_match(R: Y, MRI, P: m_SpecificICstOrSplat(RequestedValue: CstX))))
5468 ReplaceReg = Z;
5469 if (ReplaceReg) {
5470 MatchInfo = [=](MachineIRBuilder &B) {
5471 auto Zero = B.buildConstant(Res: MRI.getType(Reg: Dst), Val: 0);
5472 B.buildSub(Dst, Src0: Zero, Src1: ReplaceReg);
5473 };
5474 return true;
5475 }
5476 }
5477 return false;
5478}
5479
5480MachineInstr *CombinerHelper::buildUDivOrURemUsingMul(MachineInstr &MI) const {
5481 unsigned Opcode = MI.getOpcode();
5482 assert(Opcode == TargetOpcode::G_UDIV || Opcode == TargetOpcode::G_UREM);
5483 auto &UDivorRem = cast<GenericMachineInstr>(Val&: MI);
5484 Register Dst = UDivorRem.getReg(Idx: 0);
5485 Register LHS = UDivorRem.getReg(Idx: 1);
5486 Register RHS = UDivorRem.getReg(Idx: 2);
5487 LLT Ty = MRI.getType(Reg: Dst);
5488 LLT ScalarTy = Ty.getScalarType();
5489 const unsigned EltBits = ScalarTy.getScalarSizeInBits();
5490 LLT ShiftAmtTy = getTargetLowering().getPreferredShiftAmountTy(ShiftValueTy: Ty);
5491 LLT ScalarShiftAmtTy = ShiftAmtTy.getScalarType();
5492
5493 auto &MIB = Builder;
5494
5495 bool UseSRL = false;
5496 SmallVector<Register, 16> Shifts, Factors;
5497 auto *RHSDefInstr = cast<GenericMachineInstr>(Val: getDefIgnoringCopies(Reg: RHS, MRI));
5498 bool IsSplat = getIConstantSplatVal(MI: *RHSDefInstr, MRI).has_value();
5499
5500 auto BuildExactUDIVPattern = [&](const Constant *C) {
5501 // Don't recompute inverses for each splat element.
5502 if (IsSplat && !Factors.empty()) {
5503 Shifts.push_back(Elt: Shifts[0]);
5504 Factors.push_back(Elt: Factors[0]);
5505 return true;
5506 }
5507
5508 auto *CI = cast<ConstantInt>(Val: C);
5509 APInt Divisor = CI->getValue();
5510 unsigned Shift = Divisor.countr_zero();
5511 if (Shift) {
5512 Divisor.lshrInPlace(ShiftAmt: Shift);
5513 UseSRL = true;
5514 }
5515
5516 // Calculate the multiplicative inverse modulo BW.
5517 APInt Factor = Divisor.multiplicativeInverse();
5518 Shifts.push_back(Elt: MIB.buildConstant(Res: ScalarShiftAmtTy, Val: Shift).getReg(Idx: 0));
5519 Factors.push_back(Elt: MIB.buildConstant(Res: ScalarTy, Val: Factor).getReg(Idx: 0));
5520 return true;
5521 };
5522
5523 if (MI.getFlag(Flag: MachineInstr::MIFlag::IsExact)) {
5524 // Collect all magic values from the build vector.
5525 if (!matchUnaryPredicate(MRI, Reg: RHS, Match: BuildExactUDIVPattern))
5526 llvm_unreachable("Expected unary predicate match to succeed");
5527
5528 Register Shift, Factor;
5529 if (Ty.isVector()) {
5530 Shift = MIB.buildBuildVector(Res: ShiftAmtTy, Ops: Shifts).getReg(Idx: 0);
5531 Factor = MIB.buildBuildVector(Res: Ty, Ops: Factors).getReg(Idx: 0);
5532 } else {
5533 Shift = Shifts[0];
5534 Factor = Factors[0];
5535 }
5536
5537 Register Res = LHS;
5538
5539 if (UseSRL)
5540 Res = MIB.buildLShr(Dst: Ty, Src0: Res, Src1: Shift, Flags: MachineInstr::IsExact).getReg(Idx: 0);
5541
5542 return MIB.buildMul(Dst: Ty, Src0: Res, Src1: Factor);
5543 }
5544
5545 unsigned KnownLeadingZeros =
5546 VT ? VT->getKnownBits(R: LHS).countMinLeadingZeros() : 0;
5547
5548 bool UseNPQ = false;
5549 SmallVector<Register, 16> PreShifts, PostShifts, MagicFactors, NPQFactors;
5550 auto BuildUDIVPattern = [&](const Constant *C) {
5551 auto *CI = cast<ConstantInt>(Val: C);
5552 const APInt &Divisor = CI->getValue();
5553
5554 bool SelNPQ = false;
5555 APInt Magic(Divisor.getBitWidth(), 0);
5556 unsigned PreShift = 0, PostShift = 0;
5557
5558 // Magic algorithm doesn't work for division by 1. We need to emit a select
5559 // at the end.
5560 // TODO: Use undef values for divisor of 1.
5561 if (!Divisor.isOne()) {
5562
5563 // UnsignedDivisionByConstantInfo doesn't work correctly if leading zeros
5564 // in the dividend exceeds the leading zeros for the divisor.
5565 UnsignedDivisionByConstantInfo magics =
5566 UnsignedDivisionByConstantInfo::get(
5567 D: Divisor, LeadingZeros: std::min(a: KnownLeadingZeros, b: Divisor.countl_zero()));
5568
5569 Magic = std::move(magics.Magic);
5570
5571 assert(magics.PreShift < Divisor.getBitWidth() &&
5572 "We shouldn't generate an undefined shift!");
5573 assert(magics.PostShift < Divisor.getBitWidth() &&
5574 "We shouldn't generate an undefined shift!");
5575 assert((!magics.IsAdd || magics.PreShift == 0) && "Unexpected pre-shift");
5576 PreShift = magics.PreShift;
5577 PostShift = magics.PostShift;
5578 SelNPQ = magics.IsAdd;
5579 }
5580
5581 PreShifts.push_back(
5582 Elt: MIB.buildConstant(Res: ScalarShiftAmtTy, Val: PreShift).getReg(Idx: 0));
5583 MagicFactors.push_back(Elt: MIB.buildConstant(Res: ScalarTy, Val: Magic).getReg(Idx: 0));
5584 NPQFactors.push_back(
5585 Elt: MIB.buildConstant(Res: ScalarTy,
5586 Val: SelNPQ ? APInt::getOneBitSet(numBits: EltBits, BitNo: EltBits - 1)
5587 : APInt::getZero(numBits: EltBits))
5588 .getReg(Idx: 0));
5589 PostShifts.push_back(
5590 Elt: MIB.buildConstant(Res: ScalarShiftAmtTy, Val: PostShift).getReg(Idx: 0));
5591 UseNPQ |= SelNPQ;
5592 return true;
5593 };
5594
5595 // Collect the shifts/magic values from each element.
5596 bool Matched = matchUnaryPredicate(MRI, Reg: RHS, Match: BuildUDIVPattern);
5597 (void)Matched;
5598 assert(Matched && "Expected unary predicate match to succeed");
5599
5600 Register PreShift, PostShift, MagicFactor, NPQFactor;
5601 auto *RHSDef = getOpcodeDef<GBuildVector>(Reg: RHS, MRI);
5602 if (RHSDef) {
5603 PreShift = MIB.buildBuildVector(Res: ShiftAmtTy, Ops: PreShifts).getReg(Idx: 0);
5604 MagicFactor = MIB.buildBuildVector(Res: Ty, Ops: MagicFactors).getReg(Idx: 0);
5605 NPQFactor = MIB.buildBuildVector(Res: Ty, Ops: NPQFactors).getReg(Idx: 0);
5606 PostShift = MIB.buildBuildVector(Res: ShiftAmtTy, Ops: PostShifts).getReg(Idx: 0);
5607 } else {
5608 assert(MRI.getType(RHS).isScalar() &&
5609 "Non-build_vector operation should have been a scalar");
5610 PreShift = PreShifts[0];
5611 MagicFactor = MagicFactors[0];
5612 PostShift = PostShifts[0];
5613 }
5614
5615 Register Q = LHS;
5616 Q = MIB.buildLShr(Dst: Ty, Src0: Q, Src1: PreShift).getReg(Idx: 0);
5617
5618 // Multiply the numerator (operand 0) by the magic value.
5619 Q = MIB.buildUMulH(Dst: Ty, Src0: Q, Src1: MagicFactor).getReg(Idx: 0);
5620
5621 if (UseNPQ) {
5622 Register NPQ = MIB.buildSub(Dst: Ty, Src0: LHS, Src1: Q).getReg(Idx: 0);
5623
5624 // For vectors we might have a mix of non-NPQ/NPQ paths, so use
5625 // G_UMULH to act as a SRL-by-1 for NPQ, else multiply by zero.
5626 if (Ty.isVector())
5627 NPQ = MIB.buildUMulH(Dst: Ty, Src0: NPQ, Src1: NPQFactor).getReg(Idx: 0);
5628 else
5629 NPQ = MIB.buildLShr(Dst: Ty, Src0: NPQ, Src1: MIB.buildConstant(Res: ShiftAmtTy, Val: 1)).getReg(Idx: 0);
5630
5631 Q = MIB.buildAdd(Dst: Ty, Src0: NPQ, Src1: Q).getReg(Idx: 0);
5632 }
5633
5634 Q = MIB.buildLShr(Dst: Ty, Src0: Q, Src1: PostShift).getReg(Idx: 0);
5635 auto One = MIB.buildConstant(Res: Ty, Val: 1);
5636 auto IsOne = MIB.buildICmp(
5637 Pred: CmpInst::Predicate::ICMP_EQ,
5638 Res: Ty.isScalar() ? LLT::integer(SizeInBits: 1) : Ty.changeElementType(NewEltTy: LLT::integer(SizeInBits: 1)),
5639 Op0: RHS, Op1: One);
5640 auto ret = MIB.buildSelect(Res: Ty, Tst: IsOne, Op0: LHS, Op1: Q);
5641
5642 if (Opcode == TargetOpcode::G_UREM) {
5643 auto Prod = MIB.buildMul(Dst: Ty, Src0: ret, Src1: RHS);
5644 return MIB.buildSub(Dst: Ty, Src0: LHS, Src1: Prod);
5645 }
5646 return ret;
5647}
5648
5649bool CombinerHelper::matchUDivOrURemByConst(MachineInstr &MI) const {
5650 unsigned Opcode = MI.getOpcode();
5651 assert(Opcode == TargetOpcode::G_UDIV || Opcode == TargetOpcode::G_UREM);
5652 Register Dst = MI.getOperand(i: 0).getReg();
5653 Register RHS = MI.getOperand(i: 2).getReg();
5654 LLT DstTy = MRI.getType(Reg: Dst);
5655
5656 auto &MF = *MI.getMF();
5657 AttributeList Attr = MF.getFunction().getAttributes();
5658 const auto &TLI = getTargetLowering();
5659 LLVMContext &Ctx = MF.getFunction().getContext();
5660 if (DstTy.getScalarSizeInBits() == 1 ||
5661 TLI.isIntDivCheap(VT: getApproximateEVTForLLT(Ty: DstTy, Ctx), Attr))
5662 return false;
5663
5664 // Don't do this for minsize because the instruction sequence is usually
5665 // larger.
5666 if (MF.getFunction().hasMinSize())
5667 return false;
5668
5669 if (Opcode == TargetOpcode::G_UDIV &&
5670 MI.getFlag(Flag: MachineInstr::MIFlag::IsExact)) {
5671 return matchUnaryPredicate(
5672 MRI, Reg: RHS, Match: [](const Constant *C) { return C && !C->isNullValue(); });
5673 }
5674
5675 MachineInstr *RHSDef;
5676 if (!mi_match(R: RHS, MRI, P: m_MInstr(MI&: RHSDef)) ||
5677 !isConstantOrConstantVector(MI&: *RHSDef, MRI))
5678 return false;
5679
5680 // Don't do this if the types are not going to be legal.
5681 if (LI) {
5682 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_MUL, {DstTy, DstTy}}))
5683 return false;
5684 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_UMULH, {DstTy}}))
5685 return false;
5686 if (!isLegalOrBeforeLegalizer(
5687 Query: {TargetOpcode::G_ICMP,
5688 {DstTy.isVector() ? DstTy.changeElementSize(NewEltSize: 1) : LLT::scalar(SizeInBits: 1),
5689 DstTy}}))
5690 return false;
5691 if (Opcode == TargetOpcode::G_UREM &&
5692 !isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_SUB, {DstTy, DstTy}}))
5693 return false;
5694 }
5695
5696 return matchUnaryPredicate(
5697 MRI, Reg: RHS, Match: [](const Constant *C) { return C && !C->isNullValue(); });
5698}
5699
5700void CombinerHelper::applyUDivOrURemByConst(MachineInstr &MI) const {
5701 auto *NewMI = buildUDivOrURemUsingMul(MI);
5702 replaceSingleDefInstWithReg(MI, Replacement: NewMI->getOperand(i: 0).getReg());
5703}
5704
5705bool CombinerHelper::matchSDivOrSRemByConst(MachineInstr &MI) const {
5706 unsigned Opcode = MI.getOpcode();
5707 assert(Opcode == TargetOpcode::G_SDIV || Opcode == TargetOpcode::G_SREM);
5708 Register Dst = MI.getOperand(i: 0).getReg();
5709 Register RHS = MI.getOperand(i: 2).getReg();
5710 LLT DstTy = MRI.getType(Reg: Dst);
5711 auto SizeInBits = DstTy.getScalarSizeInBits();
5712 LLT WideTy = DstTy.changeElementSize(NewEltSize: SizeInBits * 2);
5713
5714 auto &MF = *MI.getMF();
5715 AttributeList Attr = MF.getFunction().getAttributes();
5716 const auto &TLI = getTargetLowering();
5717 LLVMContext &Ctx = MF.getFunction().getContext();
5718 if (DstTy.getScalarSizeInBits() < 3 ||
5719 TLI.isIntDivCheap(VT: getApproximateEVTForLLT(Ty: DstTy, Ctx), Attr))
5720 return false;
5721
5722 // Don't do this for minsize because the instruction sequence is usually
5723 // larger.
5724 if (MF.getFunction().hasMinSize())
5725 return false;
5726
5727 // If the sdiv has an 'exact' flag we can use a simpler lowering.
5728 if (Opcode == TargetOpcode::G_SDIV &&
5729 MI.getFlag(Flag: MachineInstr::MIFlag::IsExact)) {
5730 return matchUnaryPredicate(
5731 MRI, Reg: RHS, Match: [](const Constant *C) { return C && !C->isNullValue(); });
5732 }
5733
5734 MachineInstr *RHSDef;
5735 if (!mi_match(R: RHS, MRI, P: m_MInstr(MI&: RHSDef)) ||
5736 !isConstantOrConstantVector(MI&: *RHSDef, MRI))
5737 return false;
5738
5739 // Don't do this if the types are not going to be legal.
5740 if (LI) {
5741 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_MUL, {DstTy, DstTy}}))
5742 return false;
5743 if (!isLegal(Query: {TargetOpcode::G_SMULH, {DstTy}}) &&
5744 !isLegalOrHasWidenScalar(Query: {TargetOpcode::G_MUL, {WideTy, WideTy}}))
5745 return false;
5746 if (Opcode == TargetOpcode::G_SREM &&
5747 !isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_SUB, {DstTy, DstTy}}))
5748 return false;
5749 }
5750
5751 return matchUnaryPredicate(
5752 MRI, Reg: RHS, Match: [](const Constant *C) { return C && !C->isNullValue(); });
5753}
5754
5755void CombinerHelper::applySDivOrSRemByConst(MachineInstr &MI) const {
5756 auto *NewMI = buildSDivOrSRemUsingMul(MI);
5757 replaceSingleDefInstWithReg(MI, Replacement: NewMI->getOperand(i: 0).getReg());
5758}
5759
5760MachineInstr *CombinerHelper::buildSDivOrSRemUsingMul(MachineInstr &MI) const {
5761 unsigned Opcode = MI.getOpcode();
5762 assert(MI.getOpcode() == TargetOpcode::G_SDIV ||
5763 Opcode == TargetOpcode::G_SREM);
5764 auto &SDivorRem = cast<GenericMachineInstr>(Val&: MI);
5765 Register Dst = SDivorRem.getReg(Idx: 0);
5766 Register LHS = SDivorRem.getReg(Idx: 1);
5767 Register RHS = SDivorRem.getReg(Idx: 2);
5768 LLT Ty = MRI.getType(Reg: Dst);
5769 LLT ScalarTy = Ty.getScalarType();
5770 const unsigned EltBits = ScalarTy.getScalarSizeInBits();
5771 LLT ShiftAmtTy = getTargetLowering().getPreferredShiftAmountTy(ShiftValueTy: Ty);
5772 LLT ScalarShiftAmtTy = ShiftAmtTy.getScalarType();
5773 auto &MIB = Builder;
5774
5775 bool UseSRA = false;
5776 SmallVector<Register, 16> ExactShifts, ExactFactors;
5777
5778 auto *RHSDefInstr = cast<GenericMachineInstr>(Val: getDefIgnoringCopies(Reg: RHS, MRI));
5779 bool IsSplat = getIConstantSplatVal(MI: *RHSDefInstr, MRI).has_value();
5780
5781 auto BuildExactSDIVPattern = [&](const Constant *C) {
5782 // Don't recompute inverses for each splat element.
5783 if (IsSplat && !ExactFactors.empty()) {
5784 ExactShifts.push_back(Elt: ExactShifts[0]);
5785 ExactFactors.push_back(Elt: ExactFactors[0]);
5786 return true;
5787 }
5788
5789 auto *CI = cast<ConstantInt>(Val: C);
5790 APInt Divisor = CI->getValue();
5791 unsigned Shift = Divisor.countr_zero();
5792 if (Shift) {
5793 Divisor.ashrInPlace(ShiftAmt: Shift);
5794 UseSRA = true;
5795 }
5796
5797 // Calculate the multiplicative inverse modulo BW.
5798 // 2^W requires W + 1 bits, so we have to extend and then truncate.
5799 APInt Factor = Divisor.multiplicativeInverse();
5800 ExactShifts.push_back(Elt: MIB.buildConstant(Res: ScalarShiftAmtTy, Val: Shift).getReg(Idx: 0));
5801 ExactFactors.push_back(Elt: MIB.buildConstant(Res: ScalarTy, Val: Factor).getReg(Idx: 0));
5802 return true;
5803 };
5804
5805 if (MI.getFlag(Flag: MachineInstr::MIFlag::IsExact)) {
5806 // Collect all magic values from the build vector.
5807 bool Matched = matchUnaryPredicate(MRI, Reg: RHS, Match: BuildExactSDIVPattern);
5808 (void)Matched;
5809 assert(Matched && "Expected unary predicate match to succeed");
5810
5811 Register Shift, Factor;
5812 if (Ty.isVector()) {
5813 Shift = MIB.buildBuildVector(Res: ShiftAmtTy, Ops: ExactShifts).getReg(Idx: 0);
5814 Factor = MIB.buildBuildVector(Res: Ty, Ops: ExactFactors).getReg(Idx: 0);
5815 } else {
5816 Shift = ExactShifts[0];
5817 Factor = ExactFactors[0];
5818 }
5819
5820 Register Res = LHS;
5821
5822 if (UseSRA)
5823 Res = MIB.buildAShr(Dst: Ty, Src0: Res, Src1: Shift, Flags: MachineInstr::IsExact).getReg(Idx: 0);
5824
5825 return MIB.buildMul(Dst: Ty, Src0: Res, Src1: Factor);
5826 }
5827
5828 SmallVector<Register, 16> MagicFactors, Factors, Shifts, ShiftMasks;
5829
5830 auto BuildSDIVPattern = [&](const Constant *C) {
5831 auto *CI = cast<ConstantInt>(Val: C);
5832 const APInt &Divisor = CI->getValue();
5833
5834 SignedDivisionByConstantInfo Magics =
5835 SignedDivisionByConstantInfo::get(D: Divisor);
5836 int NumeratorFactor = 0;
5837 int ShiftMask = -1;
5838
5839 if (Divisor.isOne() || Divisor.isAllOnes()) {
5840 // If d is +1/-1, we just multiply the numerator by +1/-1.
5841 NumeratorFactor = Divisor.getSExtValue();
5842 Magics.Magic = 0;
5843 Magics.ShiftAmount = 0;
5844 ShiftMask = 0;
5845 } else if (Divisor.isStrictlyPositive() && Magics.Magic.isNegative()) {
5846 // If d > 0 and m < 0, add the numerator.
5847 NumeratorFactor = 1;
5848 } else if (Divisor.isNegative() && Magics.Magic.isStrictlyPositive()) {
5849 // If d < 0 and m > 0, subtract the numerator.
5850 NumeratorFactor = -1;
5851 }
5852
5853 MagicFactors.push_back(Elt: MIB.buildConstant(Res: ScalarTy, Val: Magics.Magic).getReg(Idx: 0));
5854 Factors.push_back(Elt: MIB.buildConstant(Res: ScalarTy, Val: NumeratorFactor).getReg(Idx: 0));
5855 Shifts.push_back(
5856 Elt: MIB.buildConstant(Res: ScalarShiftAmtTy, Val: Magics.ShiftAmount).getReg(Idx: 0));
5857 ShiftMasks.push_back(Elt: MIB.buildConstant(Res: ScalarTy, Val: ShiftMask).getReg(Idx: 0));
5858
5859 return true;
5860 };
5861
5862 // Collect the shifts/magic values from each element.
5863 bool Matched = matchUnaryPredicate(MRI, Reg: RHS, Match: BuildSDIVPattern);
5864 (void)Matched;
5865 assert(Matched && "Expected unary predicate match to succeed");
5866
5867 Register MagicFactor, Factor, Shift, ShiftMask;
5868 auto *RHSDef = getOpcodeDef<GBuildVector>(Reg: RHS, MRI);
5869 if (RHSDef) {
5870 MagicFactor = MIB.buildBuildVector(Res: Ty, Ops: MagicFactors).getReg(Idx: 0);
5871 Factor = MIB.buildBuildVector(Res: Ty, Ops: Factors).getReg(Idx: 0);
5872 Shift = MIB.buildBuildVector(Res: ShiftAmtTy, Ops: Shifts).getReg(Idx: 0);
5873 ShiftMask = MIB.buildBuildVector(Res: Ty, Ops: ShiftMasks).getReg(Idx: 0);
5874 } else {
5875 assert(MRI.getType(RHS).isScalar() &&
5876 "Non-build_vector operation should have been a scalar");
5877 MagicFactor = MagicFactors[0];
5878 Factor = Factors[0];
5879 Shift = Shifts[0];
5880 ShiftMask = ShiftMasks[0];
5881 }
5882
5883 Register Q = LHS;
5884 Q = MIB.buildSMulH(Dst: Ty, Src0: LHS, Src1: MagicFactor).getReg(Idx: 0);
5885
5886 // (Optionally) Add/subtract the numerator using Factor.
5887 Factor = MIB.buildMul(Dst: Ty, Src0: LHS, Src1: Factor).getReg(Idx: 0);
5888 Q = MIB.buildAdd(Dst: Ty, Src0: Q, Src1: Factor).getReg(Idx: 0);
5889
5890 // Shift right algebraic by shift value.
5891 Q = MIB.buildAShr(Dst: Ty, Src0: Q, Src1: Shift).getReg(Idx: 0);
5892
5893 // Extract the sign bit, mask it and add it to the quotient.
5894 auto SignShift = MIB.buildConstant(Res: ShiftAmtTy, Val: EltBits - 1);
5895 auto T = MIB.buildLShr(Dst: Ty, Src0: Q, Src1: SignShift);
5896 T = MIB.buildAnd(Dst: Ty, Src0: T, Src1: ShiftMask);
5897 auto ret = MIB.buildAdd(Dst: Ty, Src0: Q, Src1: T);
5898
5899 if (Opcode == TargetOpcode::G_SREM) {
5900 auto Prod = MIB.buildMul(Dst: Ty, Src0: ret, Src1: RHS);
5901 return MIB.buildSub(Dst: Ty, Src0: LHS, Src1: Prod);
5902 }
5903 return ret;
5904}
5905
5906bool CombinerHelper::matchDivByPow2(MachineInstr &MI, bool IsSigned) const {
5907 assert((MI.getOpcode() == TargetOpcode::G_SDIV ||
5908 MI.getOpcode() == TargetOpcode::G_UDIV) &&
5909 "Expected SDIV or UDIV");
5910 auto &Div = cast<GenericMachineInstr>(Val&: MI);
5911 Register RHS = Div.getReg(Idx: 2);
5912 auto MatchPow2 = [&](const Constant *C) {
5913 auto *CI = dyn_cast<ConstantInt>(Val: C);
5914 return CI && (CI->getValue().isPowerOf2() ||
5915 (IsSigned && CI->getValue().isNegatedPowerOf2()));
5916 };
5917 return matchUnaryPredicate(MRI, Reg: RHS, Match: MatchPow2, /*AllowUndefs=*/false);
5918}
5919
5920void CombinerHelper::applySDivByPow2(MachineInstr &MI) const {
5921 assert(MI.getOpcode() == TargetOpcode::G_SDIV && "Expected SDIV");
5922 auto &SDiv = cast<GenericMachineInstr>(Val&: MI);
5923 Register Dst = SDiv.getReg(Idx: 0);
5924 Register LHS = SDiv.getReg(Idx: 1);
5925 Register RHS = SDiv.getReg(Idx: 2);
5926 LLT Ty = MRI.getType(Reg: Dst);
5927 LLT ShiftAmtTy = getTargetLowering().getPreferredShiftAmountTy(ShiftValueTy: Ty);
5928 LLT CCVT = Ty.isVector() ? LLT::vector(EC: Ty.getElementCount(), ScalarTy: LLT::integer(SizeInBits: 1))
5929 : LLT::integer(SizeInBits: 1);
5930
5931 // Effectively we want to lower G_SDIV %lhs, %rhs, where %rhs is a power of 2,
5932 // to the following version:
5933 //
5934 // %c1 = G_CTTZ %rhs
5935 // %inexact = G_SUB $bitwidth, %c1
5936 // %sign = %G_ASHR %lhs, $(bitwidth - 1)
5937 // %lshr = G_LSHR %sign, %inexact
5938 // %add = G_ADD %lhs, %lshr
5939 // %ashr = G_ASHR %add, %c1
5940 // %ashr = G_SELECT, %isoneorallones, %lhs, %ashr
5941 // %zero = G_CONSTANT $0
5942 // %neg = G_NEG %ashr
5943 // %isneg = G_ICMP SLT %rhs, %zero
5944 // %res = G_SELECT %isneg, %neg, %ashr
5945
5946 unsigned BitWidth = Ty.getScalarSizeInBits();
5947 auto Zero = Builder.buildConstant(Res: Ty, Val: 0);
5948
5949 auto Bits = Builder.buildConstant(Res: ShiftAmtTy, Val: BitWidth);
5950 auto C1 = Builder.buildCTTZ(Dst: ShiftAmtTy, Src0: RHS);
5951 auto Inexact = Builder.buildSub(Dst: ShiftAmtTy, Src0: Bits, Src1: C1);
5952 // Splat the sign bit into the register
5953 auto Sign = Builder.buildAShr(
5954 Dst: Ty, Src0: LHS, Src1: Builder.buildConstant(Res: ShiftAmtTy, Val: BitWidth - 1));
5955
5956 // Add (LHS < 0) ? abs2 - 1 : 0;
5957 auto LSrl = Builder.buildLShr(Dst: Ty, Src0: Sign, Src1: Inexact);
5958 auto Add = Builder.buildAdd(Dst: Ty, Src0: LHS, Src1: LSrl);
5959 auto AShr = Builder.buildAShr(Dst: Ty, Src0: Add, Src1: C1);
5960
5961 // Special case: (sdiv X, 1) -> X
5962 // Special Case: (sdiv X, -1) -> 0-X
5963 auto One = Builder.buildConstant(Res: Ty, Val: 1);
5964 auto MinusOne = Builder.buildConstant(Res: Ty, Val: -1);
5965 auto IsOne = Builder.buildICmp(Pred: CmpInst::Predicate::ICMP_EQ, Res: CCVT, Op0: RHS, Op1: One);
5966 auto IsMinusOne =
5967 Builder.buildICmp(Pred: CmpInst::Predicate::ICMP_EQ, Res: CCVT, Op0: RHS, Op1: MinusOne);
5968 auto IsOneOrMinusOne = Builder.buildOr(Dst: CCVT, Src0: IsOne, Src1: IsMinusOne);
5969 AShr = Builder.buildSelect(Res: Ty, Tst: IsOneOrMinusOne, Op0: LHS, Op1: AShr);
5970
5971 // If divided by a positive value, we're done. Otherwise, the result must be
5972 // negated.
5973 auto Neg = Builder.buildNeg(Dst: Ty, Src0: AShr);
5974 auto IsNeg = Builder.buildICmp(Pred: CmpInst::Predicate::ICMP_SLT, Res: CCVT, Op0: RHS, Op1: Zero);
5975 Builder.buildSelect(Res: MI.getOperand(i: 0).getReg(), Tst: IsNeg, Op0: Neg, Op1: AShr);
5976 MI.eraseFromParent();
5977}
5978
5979void CombinerHelper::applyUDivByPow2(MachineInstr &MI) const {
5980 assert(MI.getOpcode() == TargetOpcode::G_UDIV && "Expected UDIV");
5981 auto &UDiv = cast<GenericMachineInstr>(Val&: MI);
5982 Register Dst = UDiv.getReg(Idx: 0);
5983 Register LHS = UDiv.getReg(Idx: 1);
5984 Register RHS = UDiv.getReg(Idx: 2);
5985 LLT Ty = MRI.getType(Reg: Dst);
5986 LLT ShiftAmtTy = getTargetLowering().getPreferredShiftAmountTy(ShiftValueTy: Ty);
5987
5988 auto C1 = Builder.buildCTTZ(Dst: ShiftAmtTy, Src0: RHS);
5989 Builder.buildLShr(Dst: MI.getOperand(i: 0).getReg(), Src0: LHS, Src1: C1);
5990 MI.eraseFromParent();
5991}
5992
5993void CombinerHelper::applySimplifySRemByPow2(MachineInstr &MI) const {
5994 assert(MI.getOpcode() == TargetOpcode::G_SREM && "Expected SREM");
5995 auto &SRem = cast<GBinOp>(Val&: MI);
5996 Register Dst = SRem.getReg(Idx: 0);
5997 Register LHS = SRem.getLHSReg();
5998 Register RHS = SRem.getRHSReg();
5999 LLT Ty = MRI.getType(Reg: Dst);
6000 LLT ShiftAmtTy = getTargetLowering().getPreferredShiftAmountTy(ShiftValueTy: Ty);
6001
6002 // Effectively we want to lower G_SREM %lhs, %rhs, where %rhs is +/- a power
6003 // of 2, to the following branch-free bias-and-mask version:
6004 //
6005 // %abs = G_ABS %rhs
6006 // %mask = G_SUB %abs, 1
6007 // %sign = G_ASHR %lhs, $(bitwidth - 1)
6008 // %bias = G_AND %sign, %mask
6009 // %biased = G_ADD %lhs, %bias
6010 // %masked = G_AND %biased, %mask
6011 // %res = G_SUB %masked, %bias
6012 //
6013 // The bias adds (|%rhs| - 1) for negative %lhs, correcting rounding towards
6014 // zero (instead of towards -inf that a plain mask would give). Constant
6015 // divisors collapse %mask to a single G_CONSTANT via the CSEMIRBuilder folds
6016 // for G_ABS and G_SUB.
6017
6018 unsigned BitWidth = Ty.getScalarSizeInBits();
6019 auto AbsRHS = Builder.buildAbs(Dst: Ty, Src: RHS);
6020 auto Mask = Builder.buildSub(Dst: Ty, Src0: AbsRHS, Src1: Builder.buildConstant(Res: Ty, Val: 1));
6021 auto BWMinusOne = Builder.buildConstant(Res: ShiftAmtTy, Val: BitWidth - 1);
6022 auto Sign = Builder.buildAShr(Dst: Ty, Src0: LHS, Src1: BWMinusOne);
6023 auto Bias = Builder.buildAnd(Dst: Ty, Src0: Sign, Src1: Mask);
6024 auto Biased = Builder.buildAdd(Dst: Ty, Src0: LHS, Src1: Bias);
6025 auto Masked = Builder.buildAnd(Dst: Ty, Src0: Biased, Src1: Mask);
6026 Builder.buildSub(Dst, Src0: Masked, Src1: Bias);
6027 MI.eraseFromParent();
6028}
6029
6030bool CombinerHelper::matchUMulHToLShr(MachineInstr &MI) const {
6031 assert(MI.getOpcode() == TargetOpcode::G_UMULH);
6032 Register RHS = MI.getOperand(i: 2).getReg();
6033 Register Dst = MI.getOperand(i: 0).getReg();
6034 LLT Ty = MRI.getType(Reg: Dst);
6035 LLT RHSTy = MRI.getType(Reg: RHS);
6036 LLT ShiftAmtTy = getTargetLowering().getPreferredShiftAmountTy(ShiftValueTy: Ty);
6037 auto MatchPow2ExceptOne = [&](const Constant *C) {
6038 if (auto *CI = dyn_cast<ConstantInt>(Val: C))
6039 return CI->getValue().isPowerOf2() && !CI->getValue().isOne();
6040 return false;
6041 };
6042 if (!matchUnaryPredicate(MRI, Reg: RHS, Match: MatchPow2ExceptOne, AllowUndefs: false))
6043 return false;
6044 // We need to check both G_LSHR and G_CTLZ because the combine uses G_CTLZ to
6045 // get log base 2, and it is not always legal for on a target.
6046 return isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_LSHR, {Ty, ShiftAmtTy}}) &&
6047 isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_CTLZ, {RHSTy, RHSTy}});
6048}
6049
6050void CombinerHelper::applyUMulHToLShr(MachineInstr &MI) const {
6051 Register LHS = MI.getOperand(i: 1).getReg();
6052 Register RHS = MI.getOperand(i: 2).getReg();
6053 Register Dst = MI.getOperand(i: 0).getReg();
6054 LLT Ty = MRI.getType(Reg: Dst);
6055 LLT ShiftAmtTy = getTargetLowering().getPreferredShiftAmountTy(ShiftValueTy: Ty);
6056 unsigned NumEltBits = Ty.getScalarSizeInBits();
6057
6058 auto LogBase2 = buildLogBase2(V: RHS, MIB&: Builder);
6059 auto ShiftAmt =
6060 Builder.buildSub(Dst: Ty, Src0: Builder.buildConstant(Res: Ty, Val: NumEltBits), Src1: LogBase2);
6061 auto Trunc = Builder.buildZExtOrTrunc(Res: ShiftAmtTy, Op: ShiftAmt);
6062 Builder.buildLShr(Dst, Src0: LHS, Src1: Trunc);
6063 MI.eraseFromParent();
6064}
6065
6066bool CombinerHelper::matchTruncSSatS(MachineInstr &MI,
6067 Register &MatchInfo) const {
6068 Register Dst = MI.getOperand(i: 0).getReg();
6069 Register Src = MI.getOperand(i: 1).getReg();
6070 LLT DstTy = MRI.getType(Reg: Dst);
6071 LLT SrcTy = MRI.getType(Reg: Src);
6072 unsigned NumDstBits = DstTy.getScalarSizeInBits();
6073 unsigned NumSrcBits = SrcTy.getScalarSizeInBits();
6074 assert(NumSrcBits > NumDstBits && "Unexpected types for truncate operation");
6075
6076 if (!LI || !isLegalOrHasFewerElements(
6077 Query: {TargetOpcode::G_TRUNC_SSAT_S, {DstTy, SrcTy}}))
6078 return false;
6079
6080 APInt SignedMax = APInt::getSignedMaxValue(numBits: NumDstBits).sext(width: NumSrcBits);
6081 APInt SignedMin = APInt::getSignedMinValue(numBits: NumDstBits).sext(width: NumSrcBits);
6082 if (mi_match(
6083 R: Src, MRI,
6084 P: m_GSMin(L: m_GSMax(L: m_Reg(R&: MatchInfo), R: m_SpecificICstOrSplat(RequestedValue: SignedMin)),
6085 R: m_SpecificICstOrSplat(RequestedValue: SignedMax))))
6086 return true;
6087 if (mi_match(
6088 R: Src, MRI,
6089 P: m_GSMax(L: m_GSMin(L: m_Reg(R&: MatchInfo), R: m_SpecificICstOrSplat(RequestedValue: SignedMax)),
6090 R: m_SpecificICstOrSplat(RequestedValue: SignedMin))))
6091 return true;
6092
6093 // CVP in the midend will often transform trunc(smin(smax(..)) into
6094 // trunc nsw(smin(..)) as the smax against INT_MIN never saturates.
6095 if (MI.getFlag(Flag: MachineInstr::MIFlag::NoSWrap) &&
6096 mi_match(R: Src, MRI,
6097 P: m_GSMin(L: m_Reg(R&: MatchInfo), R: m_SpecificICstOrSplat(RequestedValue: SignedMax))))
6098 return true;
6099
6100 return false;
6101}
6102
6103void CombinerHelper::applyTruncSSatS(MachineInstr &MI,
6104 Register &MatchInfo) const {
6105 Register Dst = MI.getOperand(i: 0).getReg();
6106 Builder.buildTruncSSatS(Res: Dst, Op: MatchInfo);
6107 MI.eraseFromParent();
6108}
6109
6110bool CombinerHelper::matchTruncSSatU(MachineInstr &MI,
6111 Register &MatchInfo) const {
6112 Register Dst = MI.getOperand(i: 0).getReg();
6113 Register Src = MI.getOperand(i: 1).getReg();
6114 LLT DstTy = MRI.getType(Reg: Dst);
6115 LLT SrcTy = MRI.getType(Reg: Src);
6116 unsigned NumDstBits = DstTy.getScalarSizeInBits();
6117 unsigned NumSrcBits = SrcTy.getScalarSizeInBits();
6118 assert(NumSrcBits > NumDstBits && "Unexpected types for truncate operation");
6119
6120 if (!LI || !isLegalOrHasFewerElements(
6121 Query: {TargetOpcode::G_TRUNC_SSAT_U, {DstTy, SrcTy}}))
6122 return false;
6123 APInt UnsignedMax = APInt::getMaxValue(numBits: NumDstBits).zext(width: NumSrcBits);
6124 return mi_match(R: Src, MRI,
6125 P: m_GSMin(L: m_GSMax(L: m_Reg(R&: MatchInfo), R: m_SpecificICstOrSplat(RequestedValue: 0)),
6126 R: m_SpecificICstOrSplat(RequestedValue: UnsignedMax))) ||
6127 mi_match(R: Src, MRI,
6128 P: m_GSMax(L: m_GSMin(L: m_Reg(R&: MatchInfo),
6129 R: m_SpecificICstOrSplat(RequestedValue: UnsignedMax)),
6130 R: m_SpecificICstOrSplat(RequestedValue: 0))) ||
6131 mi_match(R: Src, MRI,
6132 P: m_GUMin(L: m_GSMax(L: m_Reg(R&: MatchInfo), R: m_SpecificICstOrSplat(RequestedValue: 0)),
6133 R: m_SpecificICstOrSplat(RequestedValue: UnsignedMax)));
6134}
6135
6136void CombinerHelper::applyTruncSSatU(MachineInstr &MI,
6137 Register &MatchInfo) const {
6138 Register Dst = MI.getOperand(i: 0).getReg();
6139 Builder.buildTruncSSatU(Res: Dst, Op: MatchInfo);
6140 MI.eraseFromParent();
6141}
6142
6143bool CombinerHelper::matchTruncUSatU(MachineInstr &MI,
6144 MachineInstr &MinMI) const {
6145 Register Min = MinMI.getOperand(i: 2).getReg();
6146 Register Val = MinMI.getOperand(i: 1).getReg();
6147 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6148 LLT SrcTy = MRI.getType(Reg: Val);
6149 unsigned NumDstBits = DstTy.getScalarSizeInBits();
6150 unsigned NumSrcBits = SrcTy.getScalarSizeInBits();
6151 assert(NumSrcBits > NumDstBits && "Unexpected types for truncate operation");
6152
6153 if (!LI || !isLegalOrHasFewerElements(
6154 Query: {TargetOpcode::G_TRUNC_SSAT_U, {DstTy, SrcTy}}))
6155 return false;
6156 APInt UnsignedMax = APInt::getMaxValue(numBits: NumDstBits).zext(width: NumSrcBits);
6157 return mi_match(R: Min, MRI, P: m_SpecificICstOrSplat(RequestedValue: UnsignedMax)) &&
6158 !mi_match(R: Val, MRI, P: m_GSMax(L: m_Reg(), R: m_Reg()));
6159}
6160
6161bool CombinerHelper::matchTruncUSatUToFPTOUISat(MachineInstr &MI,
6162 MachineInstr &SrcMI) const {
6163 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6164 LLT SrcTy = MRI.getType(Reg: SrcMI.getOperand(i: 1).getReg());
6165
6166 return LI &&
6167 isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_FPTOUI_SAT, {DstTy, SrcTy}});
6168}
6169
6170bool CombinerHelper::matchRedundantNegOperands(MachineInstr &MI,
6171 BuildFnTy &MatchInfo) const {
6172 unsigned Opc = MI.getOpcode();
6173 assert(Opc == TargetOpcode::G_FADD || Opc == TargetOpcode::G_FSUB);
6174
6175 Register Dst = MI.getOperand(i: 0).getReg();
6176 Register X = MI.getOperand(i: 1).getReg();
6177 Register Y = MI.getOperand(i: 2).getReg();
6178 LLT Type = MRI.getType(Reg: Dst);
6179
6180 // fold (fadd x, fneg(y)) -> (fsub x, y)
6181 // fold (fadd fneg(y), x) -> (fsub x, y)
6182 // G_ADD is commutative so both cases are checked by m_GFAdd
6183 if (mi_match(R: Dst, MRI, P: m_GFAdd(L: m_Reg(R&: X), R: m_GFNeg(Src: m_Reg(R&: Y)))) &&
6184 isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_FSUB, {Type}})) {
6185 Opc = TargetOpcode::G_FSUB;
6186 }
6187 /// fold (fsub x, fneg(y)) -> (fadd x, y)
6188 else if (mi_match(R: Dst, MRI, P: m_GFSub(L: m_Reg(R&: X), R: m_GFNeg(Src: m_Reg(R&: Y)))) &&
6189 isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_FADD, {Type}})) {
6190 Opc = TargetOpcode::G_FADD;
6191 } else
6192 return false;
6193
6194 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6195 Observer.changingInstr(MI);
6196 MI.setDesc(B.getTII().get(Opcode: Opc));
6197 MI.getOperand(i: 1).setReg(X);
6198 MI.getOperand(i: 2).setReg(Y);
6199 Observer.changedInstr(MI);
6200 };
6201 return true;
6202}
6203
6204bool CombinerHelper::matchFsubToFneg(MachineInstr &MI,
6205 Register &MatchInfo) const {
6206 assert(MI.getOpcode() == TargetOpcode::G_FSUB);
6207
6208 Register LHS = MI.getOperand(i: 1).getReg();
6209 MatchInfo = MI.getOperand(i: 2).getReg();
6210 LLT Ty = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6211
6212 const auto LHSCst = Ty.isVector()
6213 ? getFConstantSplat(VReg: LHS, MRI, /* allowUndef */ AllowUndef: true)
6214 : getFConstantVRegValWithLookThrough(VReg: LHS, MRI);
6215 if (!LHSCst)
6216 return false;
6217
6218 // -0.0 is always allowed
6219 if (LHSCst->Value.isNegZero())
6220 return true;
6221
6222 // +0.0 is only allowed if nsz is set.
6223 if (LHSCst->Value.isPosZero())
6224 return MI.getFlag(Flag: MachineInstr::FmNsz);
6225
6226 return false;
6227}
6228
6229void CombinerHelper::applyFsubToFneg(MachineInstr &MI,
6230 Register &MatchInfo) const {
6231 Register Dst = MI.getOperand(i: 0).getReg();
6232 Builder.buildFNeg(
6233 Dst, Src0: Builder.buildFCanonicalize(Dst: MRI.getType(Reg: Dst), Src0: MatchInfo).getReg(Idx: 0));
6234 eraseInst(MI);
6235}
6236
6237/// Checks if \p MI is TargetOpcode::G_FMUL and contractable either
6238/// due to global flags or MachineInstr flags.
6239static bool isContractableFMul(MachineInstr &MI, bool AllowFusionGlobally) {
6240 if (MI.getOpcode() != TargetOpcode::G_FMUL)
6241 return false;
6242 return AllowFusionGlobally || MI.getFlag(Flag: MachineInstr::MIFlag::FmContract);
6243}
6244
6245static bool hasMoreUses(const MachineInstr &MI0, const MachineInstr &MI1,
6246 const MachineRegisterInfo &MRI) {
6247 return std::distance(first: MRI.use_instr_nodbg_begin(RegNo: MI0.getOperand(i: 0).getReg()),
6248 last: MRI.use_instr_nodbg_end()) >
6249 std::distance(first: MRI.use_instr_nodbg_begin(RegNo: MI1.getOperand(i: 0).getReg()),
6250 last: MRI.use_instr_nodbg_end());
6251}
6252
6253bool CombinerHelper::canCombineFMadOrFMA(MachineInstr &MI,
6254 bool &AllowFusionGlobally,
6255 bool &HasFMAD, bool &Aggressive,
6256 bool CanReassociate) const {
6257
6258 auto *MF = MI.getMF();
6259 const auto &TLI = *MF->getSubtarget().getTargetLowering();
6260 LLT DstType = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6261
6262 if (CanReassociate && !MI.getFlag(Flag: MachineInstr::MIFlag::FmReassoc))
6263 return false;
6264
6265 // Floating-point multiply-add with intermediate rounding.
6266 HasFMAD = (!isPreLegalize() && TLI.isFMADLegal(MI, Ty: DstType));
6267 // Floating-point multiply-add without intermediate rounding.
6268 bool HasFMA = TLI.isFMAFasterThanFMulAndFAdd(MF: *MF, DstType) &&
6269 isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_FMA, {DstType}});
6270 // No valid opcode, do not combine.
6271 if (!HasFMAD && !HasFMA)
6272 return false;
6273
6274 // FMAD (with intermediate rounding) is always safe to form; FMA requires the
6275 // contract fast-math flag.
6276 AllowFusionGlobally = HasFMAD;
6277 // If the addition is not contractable, do not combine.
6278 if (!AllowFusionGlobally && !MI.getFlag(Flag: MachineInstr::MIFlag::FmContract))
6279 return false;
6280
6281 Aggressive = TLI.enableAggressiveFMAFusion(Ty: DstType);
6282 return true;
6283}
6284
6285bool CombinerHelper::matchCombineFAddFMulToFMadOrFMA(
6286 MachineInstr &MI,
6287 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
6288 assert(MI.getOpcode() == TargetOpcode::G_FADD);
6289
6290 bool AllowFusionGlobally, HasFMAD, Aggressive;
6291 if (!canCombineFMadOrFMA(MI, AllowFusionGlobally, HasFMAD, Aggressive))
6292 return false;
6293
6294 Register Op1 = MI.getOperand(i: 1).getReg();
6295 Register Op2 = MI.getOperand(i: 2).getReg();
6296 MachineInstr *Op1Def, *Op2Def;
6297 if (!mi_match(R: Op1, MRI, P: m_MInstr(MI&: Op1Def)) ||
6298 !mi_match(R: Op2, MRI, P: m_MInstr(MI&: Op2Def)))
6299 return false;
6300 DefinitionAndSourceRegister LHS = {.MI: Op1Def, .Reg: Op1};
6301 DefinitionAndSourceRegister RHS = {.MI: Op2Def, .Reg: Op2};
6302 unsigned PreferredFusedOpcode =
6303 HasFMAD ? TargetOpcode::G_FMAD : TargetOpcode::G_FMA;
6304
6305 // If we have two choices trying to fold (fadd (fmul u, v), (fmul x, y)),
6306 // prefer to fold the multiply with fewer uses.
6307 if (Aggressive && isContractableFMul(MI&: *LHS.MI, AllowFusionGlobally) &&
6308 isContractableFMul(MI&: *RHS.MI, AllowFusionGlobally)) {
6309 if (hasMoreUses(MI0: *LHS.MI, MI1: *RHS.MI, MRI))
6310 std::swap(a&: LHS, b&: RHS);
6311 }
6312
6313 // fold (fadd (fmul x, y), z) -> (fma x, y, z)
6314 if (isContractableFMul(MI&: *LHS.MI, AllowFusionGlobally) &&
6315 (Aggressive || MRI.hasOneNonDBGUse(RegNo: LHS.Reg))) {
6316 unsigned Flags = MI.getFlags() & LHS.MI->getFlags();
6317 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6318 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {MI.getOperand(i: 0).getReg()},
6319 SrcOps: {LHS.MI->getOperand(i: 1).getReg(),
6320 LHS.MI->getOperand(i: 2).getReg(), RHS.Reg},
6321 Flags);
6322 };
6323 return true;
6324 }
6325
6326 // fold (fadd x, (fmul y, z)) -> (fma y, z, x)
6327 if (isContractableFMul(MI&: *RHS.MI, AllowFusionGlobally) &&
6328 (Aggressive || MRI.hasOneNonDBGUse(RegNo: RHS.Reg))) {
6329 unsigned Flags = MI.getFlags() & RHS.MI->getFlags();
6330 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6331 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {MI.getOperand(i: 0).getReg()},
6332 SrcOps: {RHS.MI->getOperand(i: 1).getReg(),
6333 RHS.MI->getOperand(i: 2).getReg(), LHS.Reg},
6334 Flags);
6335 };
6336 return true;
6337 }
6338
6339 return false;
6340}
6341
6342bool CombinerHelper::matchCombineFAddFpExtFMulToFMadOrFMA(
6343 MachineInstr &MI,
6344 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
6345 assert(MI.getOpcode() == TargetOpcode::G_FADD);
6346
6347 bool AllowFusionGlobally, HasFMAD, Aggressive;
6348 if (!canCombineFMadOrFMA(MI, AllowFusionGlobally, HasFMAD, Aggressive))
6349 return false;
6350
6351 const auto &TLI = *MI.getMF()->getSubtarget().getTargetLowering();
6352 Register Op1 = MI.getOperand(i: 1).getReg();
6353 Register Op2 = MI.getOperand(i: 2).getReg();
6354 MachineInstr *Op1Def, *Op2Def;
6355 if (!mi_match(R: Op1, MRI, P: m_MInstr(MI&: Op1Def)) ||
6356 !mi_match(R: Op2, MRI, P: m_MInstr(MI&: Op2Def)))
6357 return false;
6358 DefinitionAndSourceRegister LHS = {.MI: Op1Def, .Reg: Op1};
6359 DefinitionAndSourceRegister RHS = {.MI: Op2Def, .Reg: Op2};
6360 LLT DstType = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6361
6362 unsigned PreferredFusedOpcode =
6363 HasFMAD ? TargetOpcode::G_FMAD : TargetOpcode::G_FMA;
6364
6365 MachineInstr *LHSFpExtSrc;
6366 bool LHSContractable =
6367 mi_match(R: LHS.Reg, MRI, P: m_GFPExt(Src: m_MInstr(MI&: LHSFpExtSrc))) &&
6368 isContractableFMul(MI&: *LHSFpExtSrc, AllowFusionGlobally) &&
6369 TLI.isFPExtFoldable(MI, Opcode: PreferredFusedOpcode, DestTy: DstType,
6370 SrcTy: MRI.getType(Reg: LHSFpExtSrc->getOperand(i: 1).getReg()));
6371 MachineInstr *RHSFpExtSrc;
6372 bool RHSContractable =
6373 mi_match(R: RHS.Reg, MRI, P: m_GFPExt(Src: m_MInstr(MI&: RHSFpExtSrc))) &&
6374 isContractableFMul(MI&: *RHSFpExtSrc, AllowFusionGlobally) &&
6375 TLI.isFPExtFoldable(MI, Opcode: PreferredFusedOpcode, DestTy: DstType,
6376 SrcTy: MRI.getType(Reg: RHSFpExtSrc->getOperand(i: 1).getReg()));
6377
6378 // fold (fadd (fpext (fmul x, y)), z) -> (fma (fpext x), (fpext y), z)
6379 if (LHSContractable || RHSContractable) {
6380 // Ensure that the contractable fmul with the fewest uses (if both are
6381 // contractable) is the LHS operand.
6382 if (!LHSContractable ||
6383 (RHSContractable && hasMoreUses(MI0: *LHSFpExtSrc, MI1: *RHSFpExtSrc, MRI))) {
6384 std::swap(a&: LHS, b&: RHS);
6385 LHSFpExtSrc = RHSFpExtSrc;
6386 }
6387
6388 unsigned Flags = MI.getFlags() & LHSFpExtSrc->getFlags();
6389 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6390 auto FpExtX = B.buildFPExt(Res: DstType, Op: LHSFpExtSrc->getOperand(i: 1).getReg());
6391 auto FpExtY = B.buildFPExt(Res: DstType, Op: LHSFpExtSrc->getOperand(i: 2).getReg());
6392 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {MI.getOperand(i: 0).getReg()},
6393 SrcOps: {FpExtX.getReg(Idx: 0), FpExtY.getReg(Idx: 0), RHS.Reg}, Flags);
6394 };
6395 return true;
6396 }
6397
6398 return false;
6399}
6400
6401bool CombinerHelper::matchCombineFAddFMAFMulToFMadOrFMA(
6402 MachineInstr &MI,
6403 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
6404 assert(MI.getOpcode() == TargetOpcode::G_FADD);
6405
6406 bool AllowFusionGlobally, HasFMAD, Aggressive;
6407 if (!canCombineFMadOrFMA(MI, AllowFusionGlobally, HasFMAD, Aggressive, CanReassociate: true))
6408 return false;
6409
6410 Register Op1 = MI.getOperand(i: 1).getReg();
6411 Register Op2 = MI.getOperand(i: 2).getReg();
6412 MachineInstr *Op1Def, *Op2Def;
6413 if (!mi_match(R: Op1, MRI, P: m_MInstr(MI&: Op1Def)) ||
6414 !mi_match(R: Op2, MRI, P: m_MInstr(MI&: Op2Def)))
6415 return false;
6416 DefinitionAndSourceRegister LHS = {.MI: Op1Def, .Reg: Op1};
6417 DefinitionAndSourceRegister RHS = {.MI: Op2Def, .Reg: Op2};
6418 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6419
6420 unsigned PreferredFusedOpcode =
6421 HasFMAD ? TargetOpcode::G_FMAD : TargetOpcode::G_FMA;
6422
6423 MachineInstr *FMA = nullptr;
6424 Register Z;
6425 // fold (fadd (fma x, y, (fmul u, v)), z) -> (fma x, y, (fma u, v, z))
6426 if (LHS.MI->getOpcode() == PreferredFusedOpcode &&
6427 mi_match(R: LHS.MI->getOperand(i: 3).getReg(), MRI,
6428 P: m_GFMul(L: m_Reg(), R: m_Reg())) &&
6429 MRI.hasOneNonDBGUse(RegNo: LHS.MI->getOperand(i: 0).getReg()) &&
6430 MRI.hasOneNonDBGUse(RegNo: LHS.MI->getOperand(i: 3).getReg())) {
6431 FMA = LHS.MI;
6432 Z = RHS.Reg;
6433 }
6434 // fold (fadd z, (fma x, y, (fmul u, v))) -> (fma x, y, (fma u, v, z))
6435 else if (RHS.MI->getOpcode() == PreferredFusedOpcode &&
6436 mi_match(R: RHS.MI->getOperand(i: 3).getReg(), MRI,
6437 P: m_GFMul(L: m_Reg(), R: m_Reg())) &&
6438 MRI.hasOneNonDBGUse(RegNo: RHS.MI->getOperand(i: 0).getReg()) &&
6439 MRI.hasOneNonDBGUse(RegNo: RHS.MI->getOperand(i: 3).getReg())) {
6440 Z = LHS.Reg;
6441 FMA = RHS.MI;
6442 }
6443
6444 if (FMA) {
6445 MachineInstr *FMulMI;
6446 if (!mi_match(R: FMA->getOperand(i: 3).getReg(), MRI, P: m_MInstr(MI&: FMulMI)))
6447 return false;
6448 Register X = FMA->getOperand(i: 1).getReg();
6449 Register Y = FMA->getOperand(i: 2).getReg();
6450 Register U = FMulMI->getOperand(i: 1).getReg();
6451 Register V = FMulMI->getOperand(i: 2).getReg();
6452 unsigned InnerFlags = MI.getFlags() & FMulMI->getFlags();
6453 unsigned OuterFlags = MI.getFlags() & FMA->getFlags();
6454
6455 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6456 Register InnerFMA = MRI.createGenericVirtualRegister(Ty: DstTy);
6457 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {InnerFMA}, SrcOps: {U, V, Z}, Flags: InnerFlags);
6458 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {MI.getOperand(i: 0).getReg()},
6459 SrcOps: {X, Y, InnerFMA}, Flags: OuterFlags);
6460 };
6461 return true;
6462 }
6463
6464 return false;
6465}
6466
6467bool CombinerHelper::matchCombineFAddFpExtFMulToFMadOrFMAAggressive(
6468 MachineInstr &MI,
6469 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
6470 assert(MI.getOpcode() == TargetOpcode::G_FADD);
6471
6472 bool AllowFusionGlobally, HasFMAD, Aggressive;
6473 if (!canCombineFMadOrFMA(MI, AllowFusionGlobally, HasFMAD, Aggressive))
6474 return false;
6475
6476 if (!Aggressive)
6477 return false;
6478
6479 const auto &TLI = *MI.getMF()->getSubtarget().getTargetLowering();
6480 LLT DstType = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6481 Register Op1 = MI.getOperand(i: 1).getReg();
6482 Register Op2 = MI.getOperand(i: 2).getReg();
6483 MachineInstr *Op1Def, *Op2Def;
6484 if (!mi_match(R: Op1, MRI, P: m_MInstr(MI&: Op1Def)) ||
6485 !mi_match(R: Op2, MRI, P: m_MInstr(MI&: Op2Def)))
6486 return false;
6487 DefinitionAndSourceRegister LHS = {.MI: Op1Def, .Reg: Op1};
6488 DefinitionAndSourceRegister RHS = {.MI: Op2Def, .Reg: Op2};
6489
6490 unsigned PreferredFusedOpcode =
6491 HasFMAD ? TargetOpcode::G_FMAD : TargetOpcode::G_FMA;
6492
6493 // If we have two choices trying to fold (fadd (fmul u, v), (fmul x, y)),
6494 // prefer to fold the multiply with fewer uses.
6495 if (Aggressive && isContractableFMul(MI&: *LHS.MI, AllowFusionGlobally) &&
6496 isContractableFMul(MI&: *RHS.MI, AllowFusionGlobally)) {
6497 if (hasMoreUses(MI0: *LHS.MI, MI1: *RHS.MI, MRI))
6498 std::swap(a&: LHS, b&: RHS);
6499 }
6500
6501 // Builds: (fma x, y, (fma (fpext u), (fpext v), z))
6502 auto buildMatchInfo = [=, &MI](Register U, Register V, Register Z, Register X,
6503 Register Y, unsigned InnerFlags,
6504 unsigned OuterFlags, MachineIRBuilder &B) {
6505 Register FpExtU = B.buildFPExt(Res: DstType, Op: U).getReg(Idx: 0);
6506 Register FpExtV = B.buildFPExt(Res: DstType, Op: V).getReg(Idx: 0);
6507 Register InnerFMA = B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {DstType},
6508 SrcOps: {FpExtU, FpExtV, Z}, Flags: InnerFlags)
6509 .getReg(Idx: 0);
6510 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {MI.getOperand(i: 0).getReg()},
6511 SrcOps: {X, Y, InnerFMA}, Flags: OuterFlags);
6512 };
6513
6514 MachineInstr *FMulMI, *FMAMI;
6515 // fold (fadd (fma x, y, (fpext (fmul u, v))), z)
6516 // -> (fma x, y, (fma (fpext u), (fpext v), z))
6517 if (LHS.MI->getOpcode() == PreferredFusedOpcode &&
6518 mi_match(R: LHS.MI->getOperand(i: 3).getReg(), MRI,
6519 P: m_GFPExt(Src: m_MInstr(MI&: FMulMI))) &&
6520 isContractableFMul(MI&: *FMulMI, AllowFusionGlobally) &&
6521 TLI.isFPExtFoldable(MI, Opcode: PreferredFusedOpcode, DestTy: DstType,
6522 SrcTy: MRI.getType(Reg: FMulMI->getOperand(i: 0).getReg()))) {
6523 unsigned InnerFlags = MI.getFlags() & FMulMI->getFlags();
6524 unsigned OuterFlags = MI.getFlags() & LHS.MI->getFlags();
6525 MatchInfo = [=](MachineIRBuilder &B) {
6526 buildMatchInfo(FMulMI->getOperand(i: 1).getReg(),
6527 FMulMI->getOperand(i: 2).getReg(), RHS.Reg,
6528 LHS.MI->getOperand(i: 1).getReg(),
6529 LHS.MI->getOperand(i: 2).getReg(), InnerFlags, OuterFlags, B);
6530 };
6531 return true;
6532 }
6533
6534 // fold (fadd (fpext (fma x, y, (fmul u, v))), z)
6535 // -> (fma (fpext x), (fpext y), (fma (fpext u), (fpext v), z))
6536 // FIXME: This turns two single-precision and one double-precision
6537 // operation into two double-precision operations, which might not be
6538 // interesting for all targets, especially GPUs.
6539 if (mi_match(R: LHS.Reg, MRI, P: m_GFPExt(Src: m_MInstr(MI&: FMAMI))) &&
6540 FMAMI->getOpcode() == PreferredFusedOpcode) {
6541 MachineInstr *FMulMI;
6542 if (!mi_match(R: FMAMI->getOperand(i: 3).getReg(), MRI, P: m_MInstr(MI&: FMulMI)))
6543 return false;
6544 if (isContractableFMul(MI&: *FMulMI, AllowFusionGlobally) &&
6545 TLI.isFPExtFoldable(MI, Opcode: PreferredFusedOpcode, DestTy: DstType,
6546 SrcTy: MRI.getType(Reg: FMAMI->getOperand(i: 0).getReg()))) {
6547 unsigned InnerFlags = MI.getFlags() & FMulMI->getFlags();
6548 unsigned OuterFlags = MI.getFlags() & FMAMI->getFlags();
6549 MatchInfo = [=](MachineIRBuilder &B) {
6550 Register X = FMAMI->getOperand(i: 1).getReg();
6551 Register Y = FMAMI->getOperand(i: 2).getReg();
6552 X = B.buildFPExt(Res: DstType, Op: X).getReg(Idx: 0);
6553 Y = B.buildFPExt(Res: DstType, Op: Y).getReg(Idx: 0);
6554 buildMatchInfo(FMulMI->getOperand(i: 1).getReg(),
6555 FMulMI->getOperand(i: 2).getReg(), RHS.Reg, X, Y,
6556 InnerFlags, OuterFlags, B);
6557 };
6558
6559 return true;
6560 }
6561 }
6562
6563 // fold (fadd z, (fma x, y, (fpext (fmul u, v)))
6564 // -> (fma x, y, (fma (fpext u), (fpext v), z))
6565 if (RHS.MI->getOpcode() == PreferredFusedOpcode &&
6566 mi_match(R: RHS.MI->getOperand(i: 3).getReg(), MRI,
6567 P: m_GFPExt(Src: m_MInstr(MI&: FMulMI))) &&
6568 isContractableFMul(MI&: *FMulMI, AllowFusionGlobally) &&
6569 TLI.isFPExtFoldable(MI, Opcode: PreferredFusedOpcode, DestTy: DstType,
6570 SrcTy: MRI.getType(Reg: FMulMI->getOperand(i: 0).getReg()))) {
6571 unsigned InnerFlags = MI.getFlags() & FMulMI->getFlags();
6572 unsigned OuterFlags = MI.getFlags() & RHS.MI->getFlags();
6573 MatchInfo = [=](MachineIRBuilder &B) {
6574 buildMatchInfo(FMulMI->getOperand(i: 1).getReg(),
6575 FMulMI->getOperand(i: 2).getReg(), LHS.Reg,
6576 RHS.MI->getOperand(i: 1).getReg(),
6577 RHS.MI->getOperand(i: 2).getReg(), InnerFlags, OuterFlags, B);
6578 };
6579 return true;
6580 }
6581
6582 // fold (fadd z, (fpext (fma x, y, (fmul u, v)))
6583 // -> (fma (fpext x), (fpext y), (fma (fpext u), (fpext v), z))
6584 // FIXME: This turns two single-precision and one double-precision
6585 // operation into two double-precision operations, which might not be
6586 // interesting for all targets, especially GPUs.
6587 if (mi_match(R: RHS.Reg, MRI, P: m_GFPExt(Src: m_MInstr(MI&: FMAMI))) &&
6588 FMAMI->getOpcode() == PreferredFusedOpcode) {
6589 MachineInstr *FMulMI;
6590 if (!mi_match(R: FMAMI->getOperand(i: 3).getReg(), MRI, P: m_MInstr(MI&: FMulMI)))
6591 return false;
6592 if (isContractableFMul(MI&: *FMulMI, AllowFusionGlobally) &&
6593 TLI.isFPExtFoldable(MI, Opcode: PreferredFusedOpcode, DestTy: DstType,
6594 SrcTy: MRI.getType(Reg: FMAMI->getOperand(i: 0).getReg()))) {
6595 unsigned InnerFlags = MI.getFlags() & FMulMI->getFlags();
6596 unsigned OuterFlags = MI.getFlags() & FMAMI->getFlags();
6597 MatchInfo = [=](MachineIRBuilder &B) {
6598 Register X = FMAMI->getOperand(i: 1).getReg();
6599 Register Y = FMAMI->getOperand(i: 2).getReg();
6600 X = B.buildFPExt(Res: DstType, Op: X).getReg(Idx: 0);
6601 Y = B.buildFPExt(Res: DstType, Op: Y).getReg(Idx: 0);
6602 buildMatchInfo(FMulMI->getOperand(i: 1).getReg(),
6603 FMulMI->getOperand(i: 2).getReg(), LHS.Reg, X, Y,
6604 InnerFlags, OuterFlags, B);
6605 };
6606 return true;
6607 }
6608 }
6609
6610 return false;
6611}
6612
6613bool CombinerHelper::matchCombineFSubFMulToFMadOrFMA(
6614 MachineInstr &MI,
6615 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
6616 assert(MI.getOpcode() == TargetOpcode::G_FSUB);
6617
6618 bool AllowFusionGlobally, HasFMAD, Aggressive;
6619 if (!canCombineFMadOrFMA(MI, AllowFusionGlobally, HasFMAD, Aggressive))
6620 return false;
6621
6622 Register Op1 = MI.getOperand(i: 1).getReg();
6623 Register Op2 = MI.getOperand(i: 2).getReg();
6624 MachineInstr *Op1Def, *Op2Def;
6625 if (!mi_match(R: Op1, MRI, P: m_MInstr(MI&: Op1Def)) ||
6626 !mi_match(R: Op2, MRI, P: m_MInstr(MI&: Op2Def)))
6627 return false;
6628 DefinitionAndSourceRegister LHS = {.MI: Op1Def, .Reg: Op1};
6629 DefinitionAndSourceRegister RHS = {.MI: Op2Def, .Reg: Op2};
6630 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6631
6632 // If we have two choices trying to fold (fsub (fmul u, v), (fmul x, y)),
6633 // prefer to fold the multiply with fewer uses.
6634 int FirstMulHasFewerUses = true;
6635 if (isContractableFMul(MI&: *LHS.MI, AllowFusionGlobally) &&
6636 isContractableFMul(MI&: *RHS.MI, AllowFusionGlobally) &&
6637 hasMoreUses(MI0: *LHS.MI, MI1: *RHS.MI, MRI))
6638 FirstMulHasFewerUses = false;
6639
6640 unsigned PreferredFusedOpcode =
6641 HasFMAD ? TargetOpcode::G_FMAD : TargetOpcode::G_FMA;
6642
6643 // fold (fsub (fmul x, y), z) -> (fma x, y, -z)
6644 if (FirstMulHasFewerUses &&
6645 (isContractableFMul(MI&: *LHS.MI, AllowFusionGlobally) &&
6646 (Aggressive || MRI.hasOneNonDBGUse(RegNo: LHS.Reg)))) {
6647 unsigned Flags = MI.getFlags() & LHS.MI->getFlags();
6648 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6649 Register NegZ = B.buildFNeg(Dst: DstTy, Src0: RHS.Reg).getReg(Idx: 0);
6650 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {MI.getOperand(i: 0).getReg()},
6651 SrcOps: {LHS.MI->getOperand(i: 1).getReg(),
6652 LHS.MI->getOperand(i: 2).getReg(), NegZ},
6653 Flags);
6654 };
6655 return true;
6656 }
6657 // fold (fsub x, (fmul y, z)) -> (fma -y, z, x)
6658 else if ((isContractableFMul(MI&: *RHS.MI, AllowFusionGlobally) &&
6659 (Aggressive || MRI.hasOneNonDBGUse(RegNo: RHS.Reg)))) {
6660 unsigned Flags = MI.getFlags() & RHS.MI->getFlags();
6661 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6662 Register NegY =
6663 B.buildFNeg(Dst: DstTy, Src0: RHS.MI->getOperand(i: 1).getReg()).getReg(Idx: 0);
6664 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {MI.getOperand(i: 0).getReg()},
6665 SrcOps: {NegY, RHS.MI->getOperand(i: 2).getReg(), LHS.Reg}, Flags);
6666 };
6667 return true;
6668 }
6669
6670 return false;
6671}
6672
6673bool CombinerHelper::matchCombineFSubFNegFMulToFMadOrFMA(
6674 MachineInstr &MI,
6675 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
6676 assert(MI.getOpcode() == TargetOpcode::G_FSUB);
6677
6678 bool AllowFusionGlobally, HasFMAD, Aggressive;
6679 if (!canCombineFMadOrFMA(MI, AllowFusionGlobally, HasFMAD, Aggressive))
6680 return false;
6681
6682 Register LHSReg = MI.getOperand(i: 1).getReg();
6683 Register RHSReg = MI.getOperand(i: 2).getReg();
6684 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6685
6686 unsigned PreferredFusedOpcode =
6687 HasFMAD ? TargetOpcode::G_FMAD : TargetOpcode::G_FMA;
6688
6689 MachineInstr *FMulMI;
6690 // fold (fsub (fneg (fmul x, y)), z) -> (fma (fneg x), y, (fneg z))
6691 if (mi_match(R: LHSReg, MRI, P: m_GFNeg(Src: m_MInstr(MI&: FMulMI))) &&
6692 (Aggressive || (MRI.hasOneNonDBGUse(RegNo: LHSReg) &&
6693 MRI.hasOneNonDBGUse(RegNo: FMulMI->getOperand(i: 0).getReg()))) &&
6694 isContractableFMul(MI&: *FMulMI, AllowFusionGlobally)) {
6695 unsigned Flags = MI.getFlags() & FMulMI->getFlags();
6696 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6697 Register NegX =
6698 B.buildFNeg(Dst: DstTy, Src0: FMulMI->getOperand(i: 1).getReg()).getReg(Idx: 0);
6699 Register NegZ = B.buildFNeg(Dst: DstTy, Src0: RHSReg).getReg(Idx: 0);
6700 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {MI.getOperand(i: 0).getReg()},
6701 SrcOps: {NegX, FMulMI->getOperand(i: 2).getReg(), NegZ}, Flags);
6702 };
6703 return true;
6704 }
6705
6706 // fold (fsub x, (fneg (fmul, y, z))) -> (fma y, z, x)
6707 if (mi_match(R: RHSReg, MRI, P: m_GFNeg(Src: m_MInstr(MI&: FMulMI))) &&
6708 (Aggressive || (MRI.hasOneNonDBGUse(RegNo: RHSReg) &&
6709 MRI.hasOneNonDBGUse(RegNo: FMulMI->getOperand(i: 0).getReg()))) &&
6710 isContractableFMul(MI&: *FMulMI, AllowFusionGlobally)) {
6711 unsigned Flags = MI.getFlags() & FMulMI->getFlags();
6712 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6713 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {MI.getOperand(i: 0).getReg()},
6714 SrcOps: {FMulMI->getOperand(i: 1).getReg(),
6715 FMulMI->getOperand(i: 2).getReg(), LHSReg},
6716 Flags);
6717 };
6718 return true;
6719 }
6720
6721 return false;
6722}
6723
6724bool CombinerHelper::matchCombineFSubFpExtFMulToFMadOrFMA(
6725 MachineInstr &MI,
6726 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
6727 assert(MI.getOpcode() == TargetOpcode::G_FSUB);
6728
6729 bool AllowFusionGlobally, HasFMAD, Aggressive;
6730 if (!canCombineFMadOrFMA(MI, AllowFusionGlobally, HasFMAD, Aggressive))
6731 return false;
6732
6733 Register LHSReg = MI.getOperand(i: 1).getReg();
6734 Register RHSReg = MI.getOperand(i: 2).getReg();
6735 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6736
6737 unsigned PreferredFusedOpcode =
6738 HasFMAD ? TargetOpcode::G_FMAD : TargetOpcode::G_FMA;
6739
6740 MachineInstr *FMulMI;
6741 // fold (fsub (fpext (fmul x, y)), z) -> (fma (fpext x), (fpext y), (fneg z))
6742 if (mi_match(R: LHSReg, MRI, P: m_GFPExt(Src: m_MInstr(MI&: FMulMI))) &&
6743 isContractableFMul(MI&: *FMulMI, AllowFusionGlobally) &&
6744 (Aggressive || MRI.hasOneNonDBGUse(RegNo: LHSReg))) {
6745 unsigned Flags = MI.getFlags() & FMulMI->getFlags();
6746 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6747 Register FpExtX =
6748 B.buildFPExt(Res: DstTy, Op: FMulMI->getOperand(i: 1).getReg()).getReg(Idx: 0);
6749 Register FpExtY =
6750 B.buildFPExt(Res: DstTy, Op: FMulMI->getOperand(i: 2).getReg()).getReg(Idx: 0);
6751 Register NegZ = B.buildFNeg(Dst: DstTy, Src0: RHSReg).getReg(Idx: 0);
6752 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {MI.getOperand(i: 0).getReg()},
6753 SrcOps: {FpExtX, FpExtY, NegZ}, Flags);
6754 };
6755 return true;
6756 }
6757
6758 // fold (fsub x, (fpext (fmul y, z))) -> (fma (fneg (fpext y)), (fpext z), x)
6759 if (mi_match(R: RHSReg, MRI, P: m_GFPExt(Src: m_MInstr(MI&: FMulMI))) &&
6760 isContractableFMul(MI&: *FMulMI, AllowFusionGlobally) &&
6761 (Aggressive || MRI.hasOneNonDBGUse(RegNo: RHSReg))) {
6762 unsigned Flags = MI.getFlags() & FMulMI->getFlags();
6763 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6764 Register FpExtY =
6765 B.buildFPExt(Res: DstTy, Op: FMulMI->getOperand(i: 1).getReg()).getReg(Idx: 0);
6766 Register NegY = B.buildFNeg(Dst: DstTy, Src0: FpExtY).getReg(Idx: 0);
6767 Register FpExtZ =
6768 B.buildFPExt(Res: DstTy, Op: FMulMI->getOperand(i: 2).getReg()).getReg(Idx: 0);
6769 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {MI.getOperand(i: 0).getReg()},
6770 SrcOps: {NegY, FpExtZ, LHSReg}, Flags);
6771 };
6772 return true;
6773 }
6774
6775 return false;
6776}
6777
6778bool CombinerHelper::matchCombineFSubFpExtFNegFMulToFMadOrFMA(
6779 MachineInstr &MI,
6780 std::function<void(MachineIRBuilder &)> &MatchInfo) const {
6781 assert(MI.getOpcode() == TargetOpcode::G_FSUB);
6782
6783 bool AllowFusionGlobally, HasFMAD, Aggressive;
6784 if (!canCombineFMadOrFMA(MI, AllowFusionGlobally, HasFMAD, Aggressive))
6785 return false;
6786
6787 const auto &TLI = *MI.getMF()->getSubtarget().getTargetLowering();
6788 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6789 Register LHSReg = MI.getOperand(i: 1).getReg();
6790 Register RHSReg = MI.getOperand(i: 2).getReg();
6791
6792 unsigned PreferredFusedOpcode =
6793 HasFMAD ? TargetOpcode::G_FMAD : TargetOpcode::G_FMA;
6794
6795 auto buildMatchInfo = [=](Register Dst, Register X, Register Y, Register Z,
6796 unsigned Flags, MachineIRBuilder &B) {
6797 Register FpExtX = B.buildFPExt(Res: DstTy, Op: X).getReg(Idx: 0);
6798 Register FpExtY = B.buildFPExt(Res: DstTy, Op: Y).getReg(Idx: 0);
6799 B.buildInstr(Opc: PreferredFusedOpcode, DstOps: {Dst}, SrcOps: {FpExtX, FpExtY, Z}, Flags);
6800 };
6801
6802 MachineInstr *FMulMI;
6803 // fold (fsub (fpext (fneg (fmul x, y))), z) ->
6804 // (fneg (fma (fpext x), (fpext y), z))
6805 // fold (fsub (fneg (fpext (fmul x, y))), z) ->
6806 // (fneg (fma (fpext x), (fpext y), z))
6807 if ((mi_match(R: LHSReg, MRI, P: m_GFPExt(Src: m_GFNeg(Src: m_MInstr(MI&: FMulMI)))) ||
6808 mi_match(R: LHSReg, MRI, P: m_GFNeg(Src: m_GFPExt(Src: m_MInstr(MI&: FMulMI))))) &&
6809 isContractableFMul(MI&: *FMulMI, AllowFusionGlobally) &&
6810 TLI.isFPExtFoldable(MI, Opcode: PreferredFusedOpcode, DestTy: DstTy,
6811 SrcTy: MRI.getType(Reg: FMulMI->getOperand(i: 0).getReg()))) {
6812 unsigned Flags = MI.getFlags() & FMulMI->getFlags();
6813 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6814 Register FMAReg = MRI.createGenericVirtualRegister(Ty: DstTy);
6815 buildMatchInfo(FMAReg, FMulMI->getOperand(i: 1).getReg(),
6816 FMulMI->getOperand(i: 2).getReg(), RHSReg, Flags, B);
6817 B.buildFNeg(Dst: MI.getOperand(i: 0).getReg(), Src0: FMAReg);
6818 };
6819 return true;
6820 }
6821
6822 // fold (fsub x, (fpext (fneg (fmul y, z)))) -> (fma (fpext y), (fpext z), x)
6823 // fold (fsub x, (fneg (fpext (fmul y, z)))) -> (fma (fpext y), (fpext z), x)
6824 if ((mi_match(R: RHSReg, MRI, P: m_GFPExt(Src: m_GFNeg(Src: m_MInstr(MI&: FMulMI)))) ||
6825 mi_match(R: RHSReg, MRI, P: m_GFNeg(Src: m_GFPExt(Src: m_MInstr(MI&: FMulMI))))) &&
6826 isContractableFMul(MI&: *FMulMI, AllowFusionGlobally) &&
6827 TLI.isFPExtFoldable(MI, Opcode: PreferredFusedOpcode, DestTy: DstTy,
6828 SrcTy: MRI.getType(Reg: FMulMI->getOperand(i: 0).getReg()))) {
6829 unsigned Flags = MI.getFlags() & FMulMI->getFlags();
6830 MatchInfo = [=, &MI](MachineIRBuilder &B) {
6831 buildMatchInfo(MI.getOperand(i: 0).getReg(), FMulMI->getOperand(i: 1).getReg(),
6832 FMulMI->getOperand(i: 2).getReg(), LHSReg, Flags, B);
6833 };
6834 return true;
6835 }
6836
6837 return false;
6838}
6839
6840bool CombinerHelper::matchCombineFMinMaxNaN(MachineInstr &MI,
6841 unsigned &IdxToPropagate) const {
6842 bool PropagateNaN;
6843 switch (MI.getOpcode()) {
6844 default:
6845 return false;
6846 case TargetOpcode::G_FMINNUM:
6847 case TargetOpcode::G_FMAXNUM:
6848 PropagateNaN = false;
6849 break;
6850 case TargetOpcode::G_FMINIMUM:
6851 case TargetOpcode::G_FMAXIMUM:
6852 PropagateNaN = true;
6853 break;
6854 }
6855
6856 auto MatchNaN = [&](unsigned Idx) {
6857 Register MaybeNaNReg = MI.getOperand(i: Idx).getReg();
6858 const ConstantFP *MaybeCst = getConstantFPVRegVal(VReg: MaybeNaNReg, MRI);
6859 if (!MaybeCst || !MaybeCst->getValueAPF().isNaN())
6860 return false;
6861 IdxToPropagate = PropagateNaN ? Idx : (Idx == 1 ? 2 : 1);
6862 return true;
6863 };
6864
6865 return MatchNaN(1) || MatchNaN(2);
6866}
6867
6868// Combine multiple FDIVs with the same divisor into multiple FMULs by the
6869// reciprocal.
6870// E.g., (a / Y; b / Y;) -> (recip = 1.0 / Y; a * recip; b * recip)
6871bool CombinerHelper::matchRepeatedFPDivisor(
6872 MachineInstr &MI, SmallVector<MachineInstr *> &MatchInfo) const {
6873 assert(MI.getOpcode() == TargetOpcode::G_FDIV);
6874
6875 Register X = MI.getOperand(i: 1).getReg();
6876 Register Y = MI.getOperand(i: 2).getReg();
6877
6878 if (!MI.getFlag(Flag: MachineInstr::MIFlag::FmArcp))
6879 return false;
6880
6881 auto IsOne = [this](Register X) {
6882 auto N0CFP = isConstantOrConstantSplatVectorFP(Def: X, MRI);
6883 return N0CFP && (N0CFP->isOne() || N0CFP->isMinusOne());
6884 };
6885
6886 // Skip if current node is a reciprocal/fneg-reciprocal.
6887 if (IsOne(X))
6888 return false;
6889
6890 // Exit early if the target does not want this transform or if there can't
6891 // possibly be enough uses of the divisor to make the transform worthwhile.
6892 unsigned MinUses = getTargetLowering().combineRepeatedFPDivisors();
6893 if (!MinUses)
6894 return false;
6895
6896 // Find all FDIV users of the same divisor. For the moment we limit all
6897 // instructions to a single BB and use the first Instr in MatchInfo as the
6898 // dominating position.
6899 MatchInfo.push_back(Elt: &MI);
6900 for (auto &U : MRI.use_nodbg_instructions(Reg: Y)) {
6901 if (&U == &MI || U.getParent() != MI.getParent())
6902 continue;
6903 if (U.getOpcode() == TargetOpcode::G_FDIV &&
6904 U.getOperand(i: 2).getReg() == Y && U.getOperand(i: 1).getReg() != Y &&
6905 !IsOne(U.getOperand(i: 1).getReg())) {
6906 // This division is eligible for optimization only if global unsafe math
6907 // is enabled or if this division allows reciprocal formation.
6908 if (U.getFlag(Flag: MachineInstr::MIFlag::FmArcp)) {
6909 MatchInfo.push_back(Elt: &U);
6910 if (dominates(DefMI: U, UseMI: *MatchInfo[0]))
6911 std::swap(a&: MatchInfo[0], b&: MatchInfo.back());
6912 }
6913 }
6914 }
6915
6916 // Now that we have the actual number of divisor uses, make sure it meets
6917 // the minimum threshold specified by the target.
6918 return MatchInfo.size() >= MinUses;
6919}
6920
6921void CombinerHelper::applyRepeatedFPDivisor(
6922 SmallVector<MachineInstr *> &MatchInfo) const {
6923 // Generate the new div at the position of the first instruction, that we have
6924 // ensured will dominate all other instructions.
6925 Builder.setInsertPt(MBB&: *MatchInfo[0]->getParent(), II: MatchInfo[0]);
6926 LLT Ty = MRI.getType(Reg: MatchInfo[0]->getOperand(i: 0).getReg());
6927 auto Div = Builder.buildFDiv(Dst: Ty, Src0: Builder.buildFConstant(Res: Ty, Val: 1.0),
6928 Src1: MatchInfo[0]->getOperand(i: 2).getReg(),
6929 Flags: MatchInfo[0]->getFlags());
6930
6931 // Replace all found div's with fmul instructions.
6932 for (MachineInstr *MI : MatchInfo) {
6933 Builder.setInsertPt(MBB&: *MI->getParent(), II: MI);
6934 Builder.buildFMul(Dst: MI->getOperand(i: 0).getReg(), Src0: MI->getOperand(i: 1).getReg(),
6935 Src1: Div->getOperand(i: 0).getReg(), Flags: MI->getFlags());
6936 MI->eraseFromParent();
6937 }
6938}
6939
6940bool CombinerHelper::matchBuildVectorIdentityFold(MachineInstr &MI,
6941 Register &MatchInfo) const {
6942 // This combine folds the following patterns:
6943 //
6944 // G_BUILD_VECTOR_TRUNC (G_BITCAST(x), G_LSHR(G_BITCAST(x), k))
6945 // G_BUILD_VECTOR(G_TRUNC(G_BITCAST(x)), G_TRUNC(G_LSHR(G_BITCAST(x), k)))
6946 // into
6947 // x
6948 // if
6949 // k == sizeof(VecEltTy)/2
6950 // type(x) == type(dst)
6951 //
6952 // G_BUILD_VECTOR(G_TRUNC(G_BITCAST(x)), undef)
6953 // into
6954 // x
6955 // if
6956 // type(x) == type(dst)
6957
6958 LLT DstVecTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6959 LLT DstEltTy = DstVecTy.getElementType();
6960
6961 Register Lo, Hi;
6962
6963 if (mi_match(
6964 MI, MRI,
6965 P: m_GBuildVector(L: m_GTrunc(Src: m_GBitcast(Src: m_Reg(R&: Lo))), R: m_GImplicitDef()))) {
6966 MatchInfo = Lo;
6967 return MRI.getType(Reg: MatchInfo) == DstVecTy;
6968 }
6969
6970 std::optional<ValueAndVReg> ShiftAmount;
6971 const auto LoPattern = m_GBitcast(Src: m_Reg(R&: Lo));
6972 const auto HiPattern = m_GLShr(L: m_GBitcast(Src: m_Reg(R&: Hi)), R: m_GCst(ValReg&: ShiftAmount));
6973 if (mi_match(
6974 MI, MRI,
6975 P: m_any_of(preds: m_GBuildVectorTrunc(L: LoPattern, R: HiPattern),
6976 preds: m_GBuildVector(L: m_GTrunc(Src: LoPattern), R: m_GTrunc(Src: HiPattern))))) {
6977 if (Lo == Hi && ShiftAmount->Value == DstEltTy.getSizeInBits()) {
6978 MatchInfo = Lo;
6979 return MRI.getType(Reg: MatchInfo) == DstVecTy;
6980 }
6981 }
6982
6983 return false;
6984}
6985
6986bool CombinerHelper::matchTruncBuildVectorFold(MachineInstr &MI,
6987 Register &MatchInfo) const {
6988 // Replace (G_TRUNC (G_BITCAST (G_BUILD_VECTOR x, y)) with just x
6989 // if type(x) == type(G_TRUNC)
6990 if (!mi_match(R: MI.getOperand(i: 1).getReg(), MRI,
6991 P: m_GBitcast(Src: m_GBuildVector(L: m_Reg(R&: MatchInfo), R: m_Reg()))))
6992 return false;
6993
6994 return MRI.getType(Reg: MatchInfo) == MRI.getType(Reg: MI.getOperand(i: 0).getReg());
6995}
6996
6997bool CombinerHelper::matchTruncLshrBuildVectorFold(MachineInstr &MI,
6998 Register &MatchInfo) const {
6999 // Replace (G_TRUNC (G_LSHR (G_BITCAST (G_BUILD_VECTOR x, y)), K)) with
7000 // y if K == size of vector element type
7001 std::optional<ValueAndVReg> ShiftAmt;
7002 if (!mi_match(R: MI.getOperand(i: 1).getReg(), MRI,
7003 P: m_GLShr(L: m_GBitcast(Src: m_GBuildVector(L: m_Reg(), R: m_Reg(R&: MatchInfo))),
7004 R: m_GCst(ValReg&: ShiftAmt))))
7005 return false;
7006
7007 LLT MatchTy = MRI.getType(Reg: MatchInfo);
7008 return ShiftAmt->Value.getZExtValue() == MatchTy.getSizeInBits() &&
7009 MatchTy == MRI.getType(Reg: MI.getOperand(i: 0).getReg());
7010}
7011
7012unsigned CombinerHelper::getFPMinMaxOpcForSelect(
7013 CmpInst::Predicate Pred, LLT DstTy,
7014 SelectPatternNaNBehaviour VsNaNRetVal) const {
7015 assert(VsNaNRetVal != SelectPatternNaNBehaviour::NOT_APPLICABLE &&
7016 "Expected a NaN behaviour?");
7017 // Choose an opcode based off of legality or the behaviour when one of the
7018 // LHS/RHS may be NaN.
7019 switch (Pred) {
7020 default:
7021 return 0;
7022 case CmpInst::FCMP_UGT:
7023 case CmpInst::FCMP_UGE:
7024 case CmpInst::FCMP_OGT:
7025 case CmpInst::FCMP_OGE:
7026 if (VsNaNRetVal == SelectPatternNaNBehaviour::RETURNS_OTHER)
7027 return TargetOpcode::G_FMAXNUM;
7028 if (VsNaNRetVal == SelectPatternNaNBehaviour::RETURNS_NAN)
7029 return TargetOpcode::G_FMAXIMUM;
7030 if (isLegal(Query: {TargetOpcode::G_FMAXNUM, {DstTy}}))
7031 return TargetOpcode::G_FMAXNUM;
7032 if (isLegal(Query: {TargetOpcode::G_FMAXIMUM, {DstTy}}))
7033 return TargetOpcode::G_FMAXIMUM;
7034 return 0;
7035 case CmpInst::FCMP_ULT:
7036 case CmpInst::FCMP_ULE:
7037 case CmpInst::FCMP_OLT:
7038 case CmpInst::FCMP_OLE:
7039 if (VsNaNRetVal == SelectPatternNaNBehaviour::RETURNS_OTHER)
7040 return TargetOpcode::G_FMINNUM;
7041 if (VsNaNRetVal == SelectPatternNaNBehaviour::RETURNS_NAN)
7042 return TargetOpcode::G_FMINIMUM;
7043 if (isLegal(Query: {TargetOpcode::G_FMINNUM, {DstTy}}))
7044 return TargetOpcode::G_FMINNUM;
7045 if (!isLegal(Query: {TargetOpcode::G_FMINIMUM, {DstTy}}))
7046 return 0;
7047 return TargetOpcode::G_FMINIMUM;
7048 }
7049}
7050
7051CombinerHelper::SelectPatternNaNBehaviour
7052CombinerHelper::computeRetValAgainstNaN(Register LHS, Register RHS,
7053 bool IsOrderedComparison) const {
7054 bool LHSSafe = VT->isKnownNeverNaN(Val: LHS);
7055 bool RHSSafe = VT->isKnownNeverNaN(Val: RHS);
7056 // Completely unsafe.
7057 if (!LHSSafe && !RHSSafe)
7058 return SelectPatternNaNBehaviour::NOT_APPLICABLE;
7059 if (LHSSafe && RHSSafe)
7060 return SelectPatternNaNBehaviour::RETURNS_ANY;
7061 // An ordered comparison will return false when given a NaN, so it
7062 // returns the RHS.
7063 if (IsOrderedComparison)
7064 return LHSSafe ? SelectPatternNaNBehaviour::RETURNS_NAN
7065 : SelectPatternNaNBehaviour::RETURNS_OTHER;
7066 // An unordered comparison will return true when given a NaN, so it
7067 // returns the LHS.
7068 return LHSSafe ? SelectPatternNaNBehaviour::RETURNS_OTHER
7069 : SelectPatternNaNBehaviour::RETURNS_NAN;
7070}
7071
7072bool CombinerHelper::matchFPSelectToMinMax(Register Dst, Register Cond,
7073 Register TrueVal, Register FalseVal,
7074 BuildFnTy &MatchInfo) const {
7075 // Match: select (fcmp cond x, y) x, y
7076 // select (fcmp cond x, y) y, x
7077 // And turn it into fminnum/fmaxnum or fmin/fmax based off of the condition.
7078 LLT DstTy = MRI.getType(Reg: Dst);
7079 // Bail out early on pointers, since we'll never want to fold to a min/max.
7080 if (DstTy.isPointer())
7081 return false;
7082 // Match a floating point compare with a less-than/greater-than predicate.
7083 // TODO: Allow multiple users of the compare if they are all selects.
7084 CmpInst::Predicate Pred;
7085 Register CmpLHS, CmpRHS;
7086 if (!mi_match(R: Cond, MRI,
7087 P: m_OneNonDBGUse(
7088 SP: m_GFCmp(P: m_Pred(P&: Pred), L: m_Reg(R&: CmpLHS), R: m_Reg(R&: CmpRHS)))) ||
7089 CmpInst::isEquality(pred: Pred))
7090 return false;
7091 SelectPatternNaNBehaviour ResWithKnownNaNInfo =
7092 computeRetValAgainstNaN(LHS: CmpLHS, RHS: CmpRHS, IsOrderedComparison: CmpInst::isOrdered(predicate: Pred));
7093 if (ResWithKnownNaNInfo == SelectPatternNaNBehaviour::NOT_APPLICABLE)
7094 return false;
7095 if (TrueVal == CmpRHS && FalseVal == CmpLHS) {
7096 std::swap(a&: CmpLHS, b&: CmpRHS);
7097 Pred = CmpInst::getSwappedPredicate(pred: Pred);
7098 if (ResWithKnownNaNInfo == SelectPatternNaNBehaviour::RETURNS_NAN)
7099 ResWithKnownNaNInfo = SelectPatternNaNBehaviour::RETURNS_OTHER;
7100 else if (ResWithKnownNaNInfo == SelectPatternNaNBehaviour::RETURNS_OTHER)
7101 ResWithKnownNaNInfo = SelectPatternNaNBehaviour::RETURNS_NAN;
7102 }
7103 if (TrueVal != CmpLHS || FalseVal != CmpRHS)
7104 return false;
7105 // Decide what type of max/min this should be based off of the predicate.
7106 unsigned Opc = getFPMinMaxOpcForSelect(Pred, DstTy, VsNaNRetVal: ResWithKnownNaNInfo);
7107 if (!Opc || !isLegal(Query: {Opc, {DstTy}}))
7108 return false;
7109 // Comparisons between signed zero and zero may have different results...
7110 // unless we have fmaximum/fminimum. In that case, we know -0 < 0.
7111 if (Opc != TargetOpcode::G_FMAXIMUM && Opc != TargetOpcode::G_FMINIMUM) {
7112 // We don't know if a comparison between two 0s will give us a consistent
7113 // result. Be conservative and only proceed if at least one side is
7114 // non-zero.
7115 auto KnownNonZeroSide = getFConstantVRegValWithLookThrough(VReg: CmpLHS, MRI);
7116 if (!KnownNonZeroSide || !KnownNonZeroSide->Value.isNonZero()) {
7117 KnownNonZeroSide = getFConstantVRegValWithLookThrough(VReg: CmpRHS, MRI);
7118 if (!KnownNonZeroSide || !KnownNonZeroSide->Value.isNonZero())
7119 return false;
7120 }
7121 }
7122 MatchInfo = [=](MachineIRBuilder &B) {
7123 B.buildInstr(Opc, DstOps: {Dst}, SrcOps: {CmpLHS, CmpRHS});
7124 };
7125 return true;
7126}
7127
7128bool CombinerHelper::matchSimplifySelectToMinMax(MachineInstr &MI,
7129 BuildFnTy &MatchInfo) const {
7130 // TODO: Handle integer cases.
7131 assert(MI.getOpcode() == TargetOpcode::G_SELECT);
7132 // Condition may be fed by a truncated compare.
7133 Register Cond = MI.getOperand(i: 1).getReg();
7134 Register MaybeTrunc;
7135 if (mi_match(R: Cond, MRI, P: m_OneNonDBGUse(SP: m_GTrunc(Src: m_Reg(R&: MaybeTrunc)))))
7136 Cond = MaybeTrunc;
7137 Register Dst = MI.getOperand(i: 0).getReg();
7138 Register TrueVal = MI.getOperand(i: 2).getReg();
7139 Register FalseVal = MI.getOperand(i: 3).getReg();
7140 return matchFPSelectToMinMax(Dst, Cond, TrueVal, FalseVal, MatchInfo);
7141}
7142
7143bool CombinerHelper::matchRedundantBinOpInEquality(MachineInstr &MI,
7144 BuildFnTy &MatchInfo) const {
7145 assert(MI.getOpcode() == TargetOpcode::G_ICMP);
7146 // (X + Y) == X --> Y == 0
7147 // (X + Y) != X --> Y != 0
7148 // (X - Y) == X --> Y == 0
7149 // (X - Y) != X --> Y != 0
7150 // (X ^ Y) == X --> Y == 0
7151 // (X ^ Y) != X --> Y != 0
7152 Register Dst = MI.getOperand(i: 0).getReg();
7153 CmpInst::Predicate Pred;
7154 Register X, Y, OpLHS, OpRHS;
7155 bool MatchedSub = mi_match(
7156 R: Dst, MRI,
7157 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))));
7158 if (MatchedSub && X != OpLHS)
7159 return false;
7160 if (!MatchedSub) {
7161 if (!mi_match(R: Dst, MRI,
7162 P: m_c_GICmp(P: m_Pred(P&: Pred), L: m_Reg(R&: X),
7163 R: m_any_of(preds: m_GAdd(L: m_Reg(R&: OpLHS), R: m_Reg(R&: OpRHS)),
7164 preds: m_GXor(L: m_Reg(R&: OpLHS), R: m_Reg(R&: OpRHS))))))
7165 return false;
7166 Y = X == OpLHS ? OpRHS : X == OpRHS ? OpLHS : Register();
7167 }
7168 MatchInfo = [=](MachineIRBuilder &B) {
7169 auto Zero = B.buildConstant(Res: MRI.getType(Reg: Y), Val: 0);
7170 B.buildICmp(Pred, Res: Dst, Op0: Y, Op1: Zero);
7171 };
7172 return CmpInst::isEquality(pred: Pred) && Y.isValid();
7173}
7174
7175/// Return the minimum useless shift amount that results in complete loss of the
7176/// source value. Return std::nullopt when it cannot determine a value.
7177static std::optional<unsigned>
7178getMinUselessShift(KnownBits ValueKB, unsigned Opcode,
7179 std::optional<int64_t> &Result) {
7180 assert((Opcode == TargetOpcode::G_SHL || Opcode == TargetOpcode::G_LSHR ||
7181 Opcode == TargetOpcode::G_ASHR) &&
7182 "Expect G_SHL, G_LSHR or G_ASHR.");
7183 auto SignificantBits = 0;
7184 switch (Opcode) {
7185 case TargetOpcode::G_SHL:
7186 SignificantBits = ValueKB.countMinTrailingZeros();
7187 Result = 0;
7188 break;
7189 case TargetOpcode::G_LSHR:
7190 Result = 0;
7191 SignificantBits = ValueKB.countMinLeadingZeros();
7192 break;
7193 case TargetOpcode::G_ASHR:
7194 if (ValueKB.isNonNegative()) {
7195 SignificantBits = ValueKB.countMinLeadingZeros();
7196 Result = 0;
7197 } else if (ValueKB.isNegative()) {
7198 SignificantBits = ValueKB.countMinLeadingOnes();
7199 Result = -1;
7200 } else {
7201 // Cannot determine shift result.
7202 Result = std::nullopt;
7203 }
7204 break;
7205 default:
7206 break;
7207 }
7208 return ValueKB.getBitWidth() - SignificantBits;
7209}
7210
7211bool CombinerHelper::matchShiftsTooBig(
7212 MachineInstr &MI, std::optional<int64_t> &MatchInfo) const {
7213 Register ShiftVal = MI.getOperand(i: 1).getReg();
7214 Register ShiftReg = MI.getOperand(i: 2).getReg();
7215 LLT ResTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
7216 auto IsShiftTooBig = [&](const Constant *C) {
7217 auto *CI = dyn_cast<ConstantInt>(Val: C);
7218 if (!CI)
7219 return false;
7220 if (CI->uge(Num: ResTy.getScalarSizeInBits())) {
7221 MatchInfo = std::nullopt;
7222 return true;
7223 }
7224 auto OptMaxUsefulShift = getMinUselessShift(ValueKB: VT->getKnownBits(R: ShiftVal),
7225 Opcode: MI.getOpcode(), Result&: MatchInfo);
7226 return OptMaxUsefulShift && CI->uge(Num: *OptMaxUsefulShift);
7227 };
7228 return matchUnaryPredicate(MRI, Reg: ShiftReg, Match: IsShiftTooBig);
7229}
7230
7231bool CombinerHelper::matchCommuteConstantToRHS(MachineInstr &MI) const {
7232 unsigned LHSOpndIdx = 1;
7233 unsigned RHSOpndIdx = 2;
7234 switch (MI.getOpcode()) {
7235 case TargetOpcode::G_UADDO:
7236 case TargetOpcode::G_SADDO:
7237 case TargetOpcode::G_UMULO:
7238 case TargetOpcode::G_SMULO:
7239 LHSOpndIdx = 2;
7240 RHSOpndIdx = 3;
7241 break;
7242 default:
7243 break;
7244 }
7245 Register LHS = MI.getOperand(i: LHSOpndIdx).getReg();
7246 Register RHS = MI.getOperand(i: RHSOpndIdx).getReg();
7247 MachineInstr *LHSDef, *RHSDef;
7248 if (!mi_match(R: LHS, MRI, P: m_MInstr(MI&: LHSDef)) ||
7249 !mi_match(R: RHS, MRI, P: m_MInstr(MI&: RHSDef)))
7250 return false;
7251
7252 if (!getIConstantVRegVal(VReg: LHS, MRI)) {
7253 // Skip commuting if LHS is not a constant. But, LHS may be a
7254 // G_CONSTANT_FOLD_BARRIER. If so we commute as long as we don't already
7255 // have a constant on the RHS.
7256 if (LHSDef->getOpcode() != TargetOpcode::G_CONSTANT_FOLD_BARRIER)
7257 return false;
7258 }
7259 // Commute as long as RHS is not a constant or G_CONSTANT_FOLD_BARRIER.
7260 return RHSDef->getOpcode() != TargetOpcode::G_CONSTANT_FOLD_BARRIER &&
7261 !getIConstantVRegVal(VReg: RHS, MRI);
7262}
7263
7264bool CombinerHelper::matchCommuteFPConstantToRHS(MachineInstr &MI) const {
7265 Register LHS = MI.getOperand(i: 1).getReg();
7266 Register RHS = MI.getOperand(i: 2).getReg();
7267 std::optional<FPValueAndVReg> ValAndVReg;
7268 if (!mi_match(R: LHS, MRI, P: m_GFCstOrSplat(FPValReg&: ValAndVReg)))
7269 return false;
7270 return !mi_match(R: RHS, MRI, P: m_GFCstOrSplat(FPValReg&: ValAndVReg));
7271}
7272
7273void CombinerHelper::applyCommuteBinOpOperands(MachineInstr &MI) const {
7274 Observer.changingInstr(MI);
7275 unsigned LHSOpndIdx = 1;
7276 unsigned RHSOpndIdx = 2;
7277 switch (MI.getOpcode()) {
7278 case TargetOpcode::G_UADDO:
7279 case TargetOpcode::G_SADDO:
7280 case TargetOpcode::G_UMULO:
7281 case TargetOpcode::G_SMULO:
7282 LHSOpndIdx = 2;
7283 RHSOpndIdx = 3;
7284 break;
7285 default:
7286 break;
7287 }
7288 Register LHSReg = MI.getOperand(i: LHSOpndIdx).getReg();
7289 Register RHSReg = MI.getOperand(i: RHSOpndIdx).getReg();
7290 MI.getOperand(i: LHSOpndIdx).setReg(RHSReg);
7291 MI.getOperand(i: RHSOpndIdx).setReg(LHSReg);
7292 Observer.changedInstr(MI);
7293}
7294
7295bool CombinerHelper::isOneOrOneSplat(Register Src, bool AllowUndefs) const {
7296 LLT SrcTy = MRI.getType(Reg: Src);
7297 if (SrcTy.isFixedVector())
7298 return isConstantSplatVector(Src, SplatValue: 1, AllowUndefs);
7299 if (SrcTy.isScalar()) {
7300 if (AllowUndefs && getOpcodeDef<GImplicitDef>(Reg: Src, MRI) != nullptr)
7301 return true;
7302 auto IConstant = getIConstantVRegValWithLookThrough(VReg: Src, MRI);
7303 return IConstant && IConstant->Value == 1;
7304 }
7305 return false; // scalable vector
7306}
7307
7308bool CombinerHelper::isZeroOrZeroSplat(Register Src, bool AllowUndefs) const {
7309 LLT SrcTy = MRI.getType(Reg: Src);
7310 if (SrcTy.isFixedVector())
7311 return isConstantSplatVector(Src, SplatValue: 0, AllowUndefs);
7312 if (SrcTy.isScalar()) {
7313 if (AllowUndefs && getOpcodeDef<GImplicitDef>(Reg: Src, MRI) != nullptr)
7314 return true;
7315 auto IConstant = getIConstantVRegValWithLookThrough(VReg: Src, MRI);
7316 return IConstant && IConstant->Value == 0;
7317 }
7318 return false; // scalable vector
7319}
7320
7321// Ignores COPYs during conformance checks.
7322// FIXME scalable vectors.
7323bool CombinerHelper::isConstantSplatVector(Register Src, int64_t SplatValue,
7324 bool AllowUndefs) const {
7325 GBuildVector *BuildVector = getOpcodeDef<GBuildVector>(Reg: Src, MRI);
7326 if (!BuildVector)
7327 return false;
7328 unsigned NumSources = BuildVector->getNumSources();
7329
7330 for (unsigned I = 0; I < NumSources; ++I) {
7331 GImplicitDef *ImplicitDef =
7332 getOpcodeDef<GImplicitDef>(Reg: BuildVector->getSourceReg(I), MRI);
7333 if (ImplicitDef && AllowUndefs)
7334 continue;
7335 if (ImplicitDef && !AllowUndefs)
7336 return false;
7337 std::optional<ValueAndVReg> IConstant =
7338 getIConstantVRegValWithLookThrough(VReg: BuildVector->getSourceReg(I), MRI);
7339 if (IConstant && IConstant->Value == SplatValue)
7340 continue;
7341 return false;
7342 }
7343 return true;
7344}
7345
7346// Ignores COPYs during lookups.
7347// FIXME scalable vectors
7348std::optional<APInt>
7349CombinerHelper::getConstantOrConstantSplatVector(Register Src) const {
7350 auto IConstant = getIConstantVRegValWithLookThrough(VReg: Src, MRI);
7351 if (IConstant)
7352 return IConstant->Value;
7353
7354 GBuildVector *BuildVector = getOpcodeDef<GBuildVector>(Reg: Src, MRI);
7355 if (!BuildVector)
7356 return std::nullopt;
7357 unsigned NumSources = BuildVector->getNumSources();
7358
7359 std::optional<APInt> Value = std::nullopt;
7360 for (unsigned I = 0; I < NumSources; ++I) {
7361 std::optional<ValueAndVReg> IConstant =
7362 getIConstantVRegValWithLookThrough(VReg: BuildVector->getSourceReg(I), MRI);
7363 if (!IConstant)
7364 return std::nullopt;
7365 if (!Value)
7366 Value = IConstant->Value;
7367 else if (*Value != IConstant->Value)
7368 return std::nullopt;
7369 }
7370 return Value;
7371}
7372
7373// FIXME G_SPLAT_VECTOR
7374bool CombinerHelper::isConstantOrConstantVectorI(Register Src) const {
7375 auto IConstant = getIConstantVRegValWithLookThrough(VReg: Src, MRI);
7376 if (IConstant)
7377 return true;
7378
7379 GBuildVector *BuildVector = getOpcodeDef<GBuildVector>(Reg: Src, MRI);
7380 if (!BuildVector)
7381 return false;
7382
7383 unsigned NumSources = BuildVector->getNumSources();
7384 for (unsigned I = 0; I < NumSources; ++I) {
7385 std::optional<ValueAndVReg> IConstant =
7386 getIConstantVRegValWithLookThrough(VReg: BuildVector->getSourceReg(I), MRI);
7387 if (!IConstant)
7388 return false;
7389 }
7390 return true;
7391}
7392
7393// TODO: use knownbits to determine zeros
7394bool CombinerHelper::tryFoldSelectOfConstants(GSelect *Select,
7395 BuildFnTy &MatchInfo) const {
7396 uint32_t Flags = Select->getFlags();
7397 Register Dest = Select->getReg(Idx: 0);
7398 Register Cond = Select->getCondReg();
7399 Register True = Select->getTrueReg();
7400 Register False = Select->getFalseReg();
7401 LLT CondTy = MRI.getType(Reg: Select->getCondReg());
7402 LLT TrueTy = MRI.getType(Reg: Select->getTrueReg());
7403
7404 // We only do this combine for scalar boolean conditions.
7405 if (CondTy != LLT::scalar(SizeInBits: 1))
7406 return false;
7407
7408 if (TrueTy.isPointer())
7409 return false;
7410
7411 // Both are scalars.
7412 std::optional<ValueAndVReg> TrueOpt =
7413 getIConstantVRegValWithLookThrough(VReg: True, MRI);
7414 std::optional<ValueAndVReg> FalseOpt =
7415 getIConstantVRegValWithLookThrough(VReg: False, MRI);
7416
7417 if (!TrueOpt || !FalseOpt)
7418 return false;
7419
7420 APInt TrueValue = TrueOpt->Value;
7421 APInt FalseValue = FalseOpt->Value;
7422
7423 // select Cond, 1, 0 --> zext (Cond)
7424 if (TrueValue.isOne() && FalseValue.isZero()) {
7425 MatchInfo = [=](MachineIRBuilder &B) {
7426 B.setInstrAndDebugLoc(*Select);
7427 B.buildZExtOrTrunc(Res: Dest, Op: Cond);
7428 };
7429 return true;
7430 }
7431
7432 // select Cond, -1, 0 --> sext (Cond)
7433 if (TrueValue.isAllOnes() && FalseValue.isZero()) {
7434 MatchInfo = [=](MachineIRBuilder &B) {
7435 B.setInstrAndDebugLoc(*Select);
7436 B.buildSExtOrTrunc(Res: Dest, Op: Cond);
7437 };
7438 return true;
7439 }
7440
7441 // select Cond, 0, 1 --> zext (!Cond)
7442 if (TrueValue.isZero() && FalseValue.isOne()) {
7443 MatchInfo = [=](MachineIRBuilder &B) {
7444 B.setInstrAndDebugLoc(*Select);
7445 Register Inner = MRI.createGenericVirtualRegister(Ty: CondTy);
7446 B.buildNot(Dst: Inner, Src0: Cond);
7447 B.buildZExtOrTrunc(Res: Dest, Op: Inner);
7448 };
7449 return true;
7450 }
7451
7452 // select Cond, 0, -1 --> sext (!Cond)
7453 if (TrueValue.isZero() && FalseValue.isAllOnes()) {
7454 MatchInfo = [=](MachineIRBuilder &B) {
7455 B.setInstrAndDebugLoc(*Select);
7456 Register Inner = MRI.createGenericVirtualRegister(Ty: CondTy);
7457 B.buildNot(Dst: Inner, Src0: Cond);
7458 B.buildSExtOrTrunc(Res: Dest, Op: Inner);
7459 };
7460 return true;
7461 }
7462
7463 // select Cond, C1, C1-1 --> add (zext Cond), C1-1
7464 if (TrueValue - 1 == FalseValue) {
7465 MatchInfo = [=](MachineIRBuilder &B) {
7466 B.setInstrAndDebugLoc(*Select);
7467 Register Inner = MRI.createGenericVirtualRegister(Ty: TrueTy);
7468 B.buildZExtOrTrunc(Res: Inner, Op: Cond);
7469 B.buildAdd(Dst: Dest, Src0: Inner, Src1: False);
7470 };
7471 return true;
7472 }
7473
7474 // select Cond, C1, C1+1 --> add (sext Cond), C1+1
7475 if (TrueValue + 1 == FalseValue) {
7476 MatchInfo = [=](MachineIRBuilder &B) {
7477 B.setInstrAndDebugLoc(*Select);
7478 Register Inner = MRI.createGenericVirtualRegister(Ty: TrueTy);
7479 B.buildSExtOrTrunc(Res: Inner, Op: Cond);
7480 B.buildAdd(Dst: Dest, Src0: Inner, Src1: False);
7481 };
7482 return true;
7483 }
7484
7485 // select Cond, Pow2, 0 --> (zext Cond) << log2(Pow2)
7486 if (TrueValue.isPowerOf2() && FalseValue.isZero()) {
7487 MatchInfo = [=](MachineIRBuilder &B) {
7488 B.setInstrAndDebugLoc(*Select);
7489 Register Inner = MRI.createGenericVirtualRegister(Ty: TrueTy);
7490 B.buildZExtOrTrunc(Res: Inner, Op: Cond);
7491 // The shift amount must be scalar.
7492 LLT ShiftTy = TrueTy.isVector() ? TrueTy.getElementType() : TrueTy;
7493 auto ShAmtC = B.buildConstant(Res: ShiftTy, Val: TrueValue.exactLogBase2());
7494 B.buildShl(Dst: Dest, Src0: Inner, Src1: ShAmtC, Flags);
7495 };
7496 return true;
7497 }
7498
7499 // select Cond, 0, Pow2 --> (zext (!Cond)) << log2(Pow2)
7500 if (FalseValue.isPowerOf2() && TrueValue.isZero()) {
7501 MatchInfo = [=](MachineIRBuilder &B) {
7502 B.setInstrAndDebugLoc(*Select);
7503 Register Not = MRI.createGenericVirtualRegister(Ty: CondTy);
7504 B.buildNot(Dst: Not, Src0: Cond);
7505 Register Inner = MRI.createGenericVirtualRegister(Ty: TrueTy);
7506 B.buildZExtOrTrunc(Res: Inner, Op: Not);
7507 // The shift amount must be scalar.
7508 LLT ShiftTy = TrueTy.isVector() ? TrueTy.getElementType() : TrueTy;
7509 auto ShAmtC = B.buildConstant(Res: ShiftTy, Val: FalseValue.exactLogBase2());
7510 B.buildShl(Dst: Dest, Src0: Inner, Src1: ShAmtC, Flags);
7511 };
7512 return true;
7513 }
7514
7515 // select Cond, -1, C --> or (sext Cond), C
7516 if (TrueValue.isAllOnes()) {
7517 MatchInfo = [=](MachineIRBuilder &B) {
7518 B.setInstrAndDebugLoc(*Select);
7519 Register Inner = MRI.createGenericVirtualRegister(Ty: TrueTy);
7520 B.buildSExtOrTrunc(Res: Inner, Op: Cond);
7521 B.buildOr(Dst: Dest, Src0: Inner, Src1: False, Flags);
7522 };
7523 return true;
7524 }
7525
7526 // select Cond, C, -1 --> or (sext (not Cond)), C
7527 if (FalseValue.isAllOnes()) {
7528 MatchInfo = [=](MachineIRBuilder &B) {
7529 B.setInstrAndDebugLoc(*Select);
7530 Register Not = MRI.createGenericVirtualRegister(Ty: CondTy);
7531 B.buildNot(Dst: Not, Src0: Cond);
7532 Register Inner = MRI.createGenericVirtualRegister(Ty: TrueTy);
7533 B.buildSExtOrTrunc(Res: Inner, Op: Not);
7534 B.buildOr(Dst: Dest, Src0: Inner, Src1: True, Flags);
7535 };
7536 return true;
7537 }
7538
7539 return false;
7540}
7541
7542// TODO: use knownbits to determine zeros
7543bool CombinerHelper::tryFoldBoolSelectToLogic(GSelect *Select,
7544 BuildFnTy &MatchInfo) const {
7545 uint32_t Flags = Select->getFlags();
7546 Register DstReg = Select->getReg(Idx: 0);
7547 Register Cond = Select->getCondReg();
7548 Register True = Select->getTrueReg();
7549 Register False = Select->getFalseReg();
7550 LLT CondTy = MRI.getType(Reg: Select->getCondReg());
7551 LLT TrueTy = MRI.getType(Reg: Select->getTrueReg());
7552
7553 // Boolean or fixed vector of booleans.
7554 if (CondTy.isScalableVector() ||
7555 (CondTy.isFixedVector() &&
7556 CondTy.getElementType().getScalarSizeInBits() != 1) ||
7557 CondTy.getScalarSizeInBits() != 1)
7558 return false;
7559
7560 if (CondTy != TrueTy)
7561 return false;
7562
7563 // select Cond, Cond, F --> or Cond, F
7564 // select Cond, 1, F --> or Cond, F
7565 if ((Cond == True) || isOneOrOneSplat(Src: True, /* AllowUndefs */ true)) {
7566 MatchInfo = [=](MachineIRBuilder &B) {
7567 B.setInstrAndDebugLoc(*Select);
7568 Register Ext = MRI.createGenericVirtualRegister(Ty: TrueTy);
7569 B.buildZExtOrTrunc(Res: Ext, Op: Cond);
7570 auto FreezeFalse = B.buildFreeze(Dst: TrueTy, Src: False);
7571 B.buildOr(Dst: DstReg, Src0: Ext, Src1: FreezeFalse, Flags);
7572 };
7573 return true;
7574 }
7575
7576 // select Cond, T, Cond --> and Cond, T
7577 // select Cond, T, 0 --> and Cond, T
7578 if ((Cond == False) || isZeroOrZeroSplat(Src: False, /* AllowUndefs */ true)) {
7579 MatchInfo = [=](MachineIRBuilder &B) {
7580 B.setInstrAndDebugLoc(*Select);
7581 Register Ext = MRI.createGenericVirtualRegister(Ty: TrueTy);
7582 B.buildZExtOrTrunc(Res: Ext, Op: Cond);
7583 auto FreezeTrue = B.buildFreeze(Dst: TrueTy, Src: True);
7584 B.buildAnd(Dst: DstReg, Src0: Ext, Src1: FreezeTrue);
7585 };
7586 return true;
7587 }
7588
7589 // select Cond, T, 1 --> or (not Cond), T
7590 if (isOneOrOneSplat(Src: False, /* AllowUndefs */ true)) {
7591 MatchInfo = [=](MachineIRBuilder &B) {
7592 B.setInstrAndDebugLoc(*Select);
7593 // First the not.
7594 Register Inner = MRI.createGenericVirtualRegister(Ty: CondTy);
7595 B.buildNot(Dst: Inner, Src0: Cond);
7596 // Then an ext to match the destination register.
7597 Register Ext = MRI.createGenericVirtualRegister(Ty: TrueTy);
7598 B.buildZExtOrTrunc(Res: Ext, Op: Inner);
7599 auto FreezeTrue = B.buildFreeze(Dst: TrueTy, Src: True);
7600 B.buildOr(Dst: DstReg, Src0: Ext, Src1: FreezeTrue, Flags);
7601 };
7602 return true;
7603 }
7604
7605 // select Cond, 0, F --> and (not Cond), F
7606 if (isZeroOrZeroSplat(Src: True, /* AllowUndefs */ true)) {
7607 MatchInfo = [=](MachineIRBuilder &B) {
7608 B.setInstrAndDebugLoc(*Select);
7609 // First the not.
7610 Register Inner = MRI.createGenericVirtualRegister(Ty: CondTy);
7611 B.buildNot(Dst: Inner, Src0: Cond);
7612 // Then an ext to match the destination register.
7613 Register Ext = MRI.createGenericVirtualRegister(Ty: TrueTy);
7614 B.buildZExtOrTrunc(Res: Ext, Op: Inner);
7615 auto FreezeFalse = B.buildFreeze(Dst: TrueTy, Src: False);
7616 B.buildAnd(Dst: DstReg, Src0: Ext, Src1: FreezeFalse);
7617 };
7618 return true;
7619 }
7620
7621 return false;
7622}
7623
7624bool CombinerHelper::matchSelectIMinMax(const MachineOperand &MO,
7625 BuildFnTy &MatchInfo) const {
7626 Register DstReg = MO.getReg();
7627 Register CondReg, True, False;
7628 if (!mi_match(R: DstReg, MRI,
7629 P: m_GISelect(Src0: m_Reg(R&: CondReg), Src1: m_Reg(R&: True), Src2: m_Reg(R&: False))))
7630 return false;
7631
7632 CmpInst::Predicate Pred;
7633 Register CmpLHS, CmpRHS;
7634 if (!mi_match(R: CondReg, MRI,
7635 P: m_GICmp(P: m_Pred(P&: Pred), L: m_Reg(R&: CmpLHS), R: m_Reg(R&: CmpRHS))))
7636 return false;
7637
7638 LLT DstTy = MRI.getType(Reg: DstReg);
7639 if (DstTy.isPointerOrPointerVector())
7640 return false;
7641
7642 // We want to fold the icmp and replace the select.
7643 if (!MRI.hasOneNonDBGUse(RegNo: CondReg))
7644 return false;
7645
7646 // We need a larger or smaller predicate for
7647 // canonicalization.
7648 if (CmpInst::isEquality(pred: Pred))
7649 return false;
7650
7651 // We can swap CmpLHS and CmpRHS for higher hitrate.
7652 if (True == CmpRHS && False == CmpLHS) {
7653 std::swap(a&: CmpLHS, b&: CmpRHS);
7654 Pred = CmpInst::getSwappedPredicate(pred: Pred);
7655 }
7656
7657 // (icmp X, Y) ? X : Y -> integer minmax.
7658 // see matchSelectPattern in ValueTracking.
7659 // Legality between G_SELECT and integer minmax can differ.
7660 if (True != CmpLHS || False != CmpRHS)
7661 return false;
7662
7663 switch (Pred) {
7664 case ICmpInst::ICMP_UGT:
7665 case ICmpInst::ICMP_UGE: {
7666 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_UMAX, DstTy}))
7667 return false;
7668 MatchInfo = [=](MachineIRBuilder &B) { B.buildUMax(Dst: DstReg, Src0: True, Src1: False); };
7669 return true;
7670 }
7671 case ICmpInst::ICMP_SGT:
7672 case ICmpInst::ICMP_SGE: {
7673 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_SMAX, DstTy}))
7674 return false;
7675 MatchInfo = [=](MachineIRBuilder &B) { B.buildSMax(Dst: DstReg, Src0: True, Src1: False); };
7676 return true;
7677 }
7678 case ICmpInst::ICMP_ULT:
7679 case ICmpInst::ICMP_ULE: {
7680 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_UMIN, DstTy}))
7681 return false;
7682 MatchInfo = [=](MachineIRBuilder &B) { B.buildUMin(Dst: DstReg, Src0: True, Src1: False); };
7683 return true;
7684 }
7685 case ICmpInst::ICMP_SLT:
7686 case ICmpInst::ICMP_SLE: {
7687 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_SMIN, DstTy}))
7688 return false;
7689 MatchInfo = [=](MachineIRBuilder &B) { B.buildSMin(Dst: DstReg, Src0: True, Src1: False); };
7690 return true;
7691 }
7692 default:
7693 return false;
7694 }
7695}
7696
7697// (neg (min/max x, (neg x))) --> (max/min x, (neg x))
7698bool CombinerHelper::matchSimplifyNegMinMax(MachineInstr &MI,
7699 BuildFnTy &MatchInfo) const {
7700 assert(MI.getOpcode() == TargetOpcode::G_SUB);
7701 Register DestReg = MI.getOperand(i: 0).getReg();
7702 LLT DestTy = MRI.getType(Reg: DestReg);
7703
7704 Register X;
7705 Register Sub0;
7706 auto NegPattern = m_all_of(preds: m_Neg(Src: m_DeferredReg(R&: X)), preds: m_Reg(R&: Sub0));
7707 if (mi_match(R: DestReg, MRI,
7708 P: m_Neg(Src: m_OneUse(SP: m_any_of(preds: m_GSMin(L: m_Reg(R&: X), R: NegPattern),
7709 preds: m_GSMax(L: m_Reg(R&: X), R: NegPattern),
7710 preds: m_GUMin(L: m_Reg(R&: X), R: NegPattern),
7711 preds: m_GUMax(L: m_Reg(R&: X), R: NegPattern)))))) {
7712 MachineInstr *MinMaxMI;
7713 if (!mi_match(R: MI.getOperand(i: 2).getReg(), MRI, P: m_MInstr(MI&: MinMaxMI)))
7714 return false;
7715 unsigned NewOpc = getInverseGMinMaxOpcode(MinMaxOpc: MinMaxMI->getOpcode());
7716 if (isLegal(Query: {NewOpc, {DestTy}})) {
7717 MatchInfo = [=](MachineIRBuilder &B) {
7718 B.buildInstr(Opc: NewOpc, DstOps: {DestReg}, SrcOps: {X, Sub0});
7719 };
7720 return true;
7721 }
7722 }
7723
7724 return false;
7725}
7726
7727bool CombinerHelper::matchSelect(MachineInstr &MI, BuildFnTy &MatchInfo) const {
7728 GSelect *Select = cast<GSelect>(Val: &MI);
7729
7730 if (tryFoldSelectOfConstants(Select, MatchInfo))
7731 return true;
7732
7733 if (tryFoldBoolSelectToLogic(Select, MatchInfo))
7734 return true;
7735
7736 return false;
7737}
7738
7739/// Fold (icmp Pred1 V1, C1) && (icmp Pred2 V2, C2)
7740/// or (icmp Pred1 V1, C1) || (icmp Pred2 V2, C2)
7741/// into a single comparison using range-based reasoning.
7742/// see InstCombinerImpl::foldAndOrOfICmpsUsingRanges.
7743bool CombinerHelper::tryFoldAndOrOrICmpsUsingRanges(
7744 GLogicalBinOp *Logic, BuildFnTy &MatchInfo) const {
7745 assert(Logic->getOpcode() != TargetOpcode::G_XOR && "unexpected xor");
7746 bool IsAnd = Logic->getOpcode() == TargetOpcode::G_AND;
7747 Register DstReg = Logic->getReg(Idx: 0);
7748 Register LHS = Logic->getLHSReg();
7749 Register RHS = Logic->getRHSReg();
7750 unsigned Flags = Logic->getFlags();
7751
7752 // We need an G_ICMP on the LHS register.
7753 GICmp *Cmp1 = getOpcodeDef<GICmp>(Reg: LHS, MRI);
7754 if (!Cmp1)
7755 return false;
7756
7757 // We need an G_ICMP on the RHS register.
7758 GICmp *Cmp2 = getOpcodeDef<GICmp>(Reg: RHS, MRI);
7759 if (!Cmp2)
7760 return false;
7761
7762 // We want to fold the icmps.
7763 if (!MRI.hasOneNonDBGUse(RegNo: Cmp1->getReg(Idx: 0)) ||
7764 !MRI.hasOneNonDBGUse(RegNo: Cmp2->getReg(Idx: 0)))
7765 return false;
7766
7767 APInt C1;
7768 APInt C2;
7769 std::optional<ValueAndVReg> MaybeC1 =
7770 getIConstantVRegValWithLookThrough(VReg: Cmp1->getRHSReg(), MRI);
7771 if (!MaybeC1)
7772 return false;
7773 C1 = MaybeC1->Value;
7774
7775 std::optional<ValueAndVReg> MaybeC2 =
7776 getIConstantVRegValWithLookThrough(VReg: Cmp2->getRHSReg(), MRI);
7777 if (!MaybeC2)
7778 return false;
7779 C2 = MaybeC2->Value;
7780
7781 Register R1 = Cmp1->getLHSReg();
7782 Register R2 = Cmp2->getLHSReg();
7783 CmpInst::Predicate Pred1 = Cmp1->getCond();
7784 CmpInst::Predicate Pred2 = Cmp2->getCond();
7785 LLT CmpTy = MRI.getType(Reg: Cmp1->getReg(Idx: 0));
7786 LLT CmpOperandTy = MRI.getType(Reg: R1);
7787
7788 if (CmpOperandTy.isPointer())
7789 return false;
7790
7791 // We build ands, adds, and constants of type CmpOperandTy.
7792 // They must be legal to build.
7793 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_AND, CmpOperandTy}) ||
7794 !isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_ADD, CmpOperandTy}) ||
7795 !isConstantLegalOrBeforeLegalizer(Ty: CmpOperandTy))
7796 return false;
7797
7798 // Look through add of a constant offset on R1, R2, or both operands. This
7799 // allows us to interpret the R + C' < C'' range idiom into a proper range.
7800 std::optional<APInt> Offset1;
7801 std::optional<APInt> Offset2;
7802 if (R1 != R2) {
7803 if (GAdd *Add = getOpcodeDef<GAdd>(Reg: R1, MRI)) {
7804 std::optional<ValueAndVReg> MaybeOffset1 =
7805 getIConstantVRegValWithLookThrough(VReg: Add->getRHSReg(), MRI);
7806 if (MaybeOffset1) {
7807 R1 = Add->getLHSReg();
7808 Offset1 = MaybeOffset1->Value;
7809 }
7810 }
7811 if (GAdd *Add = getOpcodeDef<GAdd>(Reg: R2, MRI)) {
7812 std::optional<ValueAndVReg> MaybeOffset2 =
7813 getIConstantVRegValWithLookThrough(VReg: Add->getRHSReg(), MRI);
7814 if (MaybeOffset2) {
7815 R2 = Add->getLHSReg();
7816 Offset2 = MaybeOffset2->Value;
7817 }
7818 }
7819 }
7820
7821 if (R1 != R2)
7822 return false;
7823
7824 // We calculate the icmp ranges including maybe offsets.
7825 ConstantRange CR1 = ConstantRange::makeExactICmpRegion(
7826 Pred: IsAnd ? ICmpInst::getInversePredicate(pred: Pred1) : Pred1, Other: C1);
7827 if (Offset1)
7828 CR1 = CR1.subtract(CI: *Offset1);
7829
7830 ConstantRange CR2 = ConstantRange::makeExactICmpRegion(
7831 Pred: IsAnd ? ICmpInst::getInversePredicate(pred: Pred2) : Pred2, Other: C2);
7832 if (Offset2)
7833 CR2 = CR2.subtract(CI: *Offset2);
7834
7835 bool CreateMask = false;
7836 APInt LowerDiff;
7837 std::optional<ConstantRange> CR = CR1.exactUnionWith(CR: CR2);
7838 if (!CR) {
7839 // We need non-wrapping ranges.
7840 if (CR1.isWrappedSet() || CR2.isWrappedSet())
7841 return false;
7842
7843 // Check whether we have equal-size ranges that only differ by one bit.
7844 // In that case we can apply a mask to map one range onto the other.
7845 LowerDiff = CR1.getLower() ^ CR2.getLower();
7846 APInt UpperDiff = (CR1.getUpper() - 1) ^ (CR2.getUpper() - 1);
7847 APInt CR1Size = CR1.getUpper() - CR1.getLower();
7848 if (!LowerDiff.isPowerOf2() || LowerDiff != UpperDiff ||
7849 CR1Size != CR2.getUpper() - CR2.getLower())
7850 return false;
7851
7852 CR = CR1.getLower().ult(RHS: CR2.getLower()) ? CR1 : CR2;
7853 CreateMask = true;
7854 }
7855
7856 if (IsAnd)
7857 CR = CR->inverse();
7858
7859 CmpInst::Predicate NewPred;
7860 APInt NewC, Offset;
7861 CR->getEquivalentICmp(Pred&: NewPred, RHS&: NewC, Offset);
7862
7863 // We take the result type of one of the original icmps, CmpTy, for
7864 // the to be build icmp. The operand type, CmpOperandTy, is used for
7865 // the other instructions and constants to be build. The types of
7866 // the parameters and output are the same for add and and. CmpTy
7867 // and the type of DstReg might differ. That is why we zext or trunc
7868 // the icmp into the destination register.
7869
7870 MatchInfo = [=](MachineIRBuilder &B) {
7871 if (CreateMask && Offset != 0) {
7872 auto TildeLowerDiff = B.buildConstant(Res: CmpOperandTy, Val: ~LowerDiff);
7873 auto And = B.buildAnd(Dst: CmpOperandTy, Src0: R1, Src1: TildeLowerDiff); // the mask.
7874 auto OffsetC = B.buildConstant(Res: CmpOperandTy, Val: Offset);
7875 auto Add = B.buildAdd(Dst: CmpOperandTy, Src0: And, Src1: OffsetC, Flags);
7876 auto NewCon = B.buildConstant(Res: CmpOperandTy, Val: NewC);
7877 auto ICmp = B.buildICmp(Pred: NewPred, Res: CmpTy, Op0: Add, Op1: NewCon);
7878 B.buildZExtOrTrunc(Res: DstReg, Op: ICmp);
7879 } else if (CreateMask && Offset == 0) {
7880 auto TildeLowerDiff = B.buildConstant(Res: CmpOperandTy, Val: ~LowerDiff);
7881 auto And = B.buildAnd(Dst: CmpOperandTy, Src0: R1, Src1: TildeLowerDiff); // the mask.
7882 auto NewCon = B.buildConstant(Res: CmpOperandTy, Val: NewC);
7883 auto ICmp = B.buildICmp(Pred: NewPred, Res: CmpTy, Op0: And, Op1: NewCon);
7884 B.buildZExtOrTrunc(Res: DstReg, Op: ICmp);
7885 } else if (!CreateMask && Offset != 0) {
7886 auto OffsetC = B.buildConstant(Res: CmpOperandTy, Val: Offset);
7887 auto Add = B.buildAdd(Dst: CmpOperandTy, Src0: R1, Src1: OffsetC, Flags);
7888 auto NewCon = B.buildConstant(Res: CmpOperandTy, Val: NewC);
7889 auto ICmp = B.buildICmp(Pred: NewPred, Res: CmpTy, Op0: Add, Op1: NewCon);
7890 B.buildZExtOrTrunc(Res: DstReg, Op: ICmp);
7891 } else if (!CreateMask && Offset == 0) {
7892 auto NewCon = B.buildConstant(Res: CmpOperandTy, Val: NewC);
7893 auto ICmp = B.buildICmp(Pred: NewPred, Res: CmpTy, Op0: R1, Op1: NewCon);
7894 B.buildZExtOrTrunc(Res: DstReg, Op: ICmp);
7895 } else {
7896 llvm_unreachable("unexpected configuration of CreateMask and Offset");
7897 }
7898 };
7899 return true;
7900}
7901
7902bool CombinerHelper::tryFoldLogicOfFCmps(GLogicalBinOp *Logic,
7903 BuildFnTy &MatchInfo) const {
7904 assert(Logic->getOpcode() != TargetOpcode::G_XOR && "unexpecte xor");
7905 Register DestReg = Logic->getReg(Idx: 0);
7906 Register LHS = Logic->getLHSReg();
7907 Register RHS = Logic->getRHSReg();
7908 bool IsAnd = Logic->getOpcode() == TargetOpcode::G_AND;
7909
7910 // We need a compare on the LHS register.
7911 GFCmp *Cmp1 = getOpcodeDef<GFCmp>(Reg: LHS, MRI);
7912 if (!Cmp1)
7913 return false;
7914
7915 // We need a compare on the RHS register.
7916 GFCmp *Cmp2 = getOpcodeDef<GFCmp>(Reg: RHS, MRI);
7917 if (!Cmp2)
7918 return false;
7919
7920 LLT CmpTy = MRI.getType(Reg: Cmp1->getReg(Idx: 0));
7921 LLT CmpOperandTy = MRI.getType(Reg: Cmp1->getLHSReg());
7922
7923 // We build one fcmp, want to fold the fcmps, replace the logic op,
7924 // and the fcmps must have the same shape.
7925 if (!isLegalOrBeforeLegalizer(
7926 Query: {TargetOpcode::G_FCMP, {CmpTy, CmpOperandTy}}) ||
7927 !MRI.hasOneNonDBGUse(RegNo: Logic->getReg(Idx: 0)) ||
7928 !MRI.hasOneNonDBGUse(RegNo: Cmp1->getReg(Idx: 0)) ||
7929 !MRI.hasOneNonDBGUse(RegNo: Cmp2->getReg(Idx: 0)) ||
7930 MRI.getType(Reg: Cmp1->getLHSReg()) != MRI.getType(Reg: Cmp2->getLHSReg()))
7931 return false;
7932
7933 CmpInst::Predicate PredL = Cmp1->getCond();
7934 CmpInst::Predicate PredR = Cmp2->getCond();
7935 Register LHS0 = Cmp1->getLHSReg();
7936 Register LHS1 = Cmp1->getRHSReg();
7937 Register RHS0 = Cmp2->getLHSReg();
7938 Register RHS1 = Cmp2->getRHSReg();
7939
7940 if (LHS0 == RHS1 && LHS1 == RHS0) {
7941 // Swap RHS operands to match LHS.
7942 PredR = CmpInst::getSwappedPredicate(pred: PredR);
7943 std::swap(a&: RHS0, b&: RHS1);
7944 }
7945
7946 if (LHS0 == RHS0 && LHS1 == RHS1) {
7947 // We determine the new predicate.
7948 unsigned CmpCodeL = getFCmpCode(CC: PredL);
7949 unsigned CmpCodeR = getFCmpCode(CC: PredR);
7950 unsigned NewPred = IsAnd ? CmpCodeL & CmpCodeR : CmpCodeL | CmpCodeR;
7951 unsigned Flags = Cmp1->getFlags() | Cmp2->getFlags();
7952 MatchInfo = [=](MachineIRBuilder &B) {
7953 // The fcmp predicates fill the lower part of the enum.
7954 FCmpInst::Predicate Pred = static_cast<FCmpInst::Predicate>(NewPred);
7955 if (Pred == FCmpInst::FCMP_FALSE &&
7956 isConstantLegalOrBeforeLegalizer(Ty: CmpTy)) {
7957 auto False = B.buildConstant(Res: CmpTy, Val: 0);
7958 B.buildZExtOrTrunc(Res: DestReg, Op: False);
7959 } else if (Pred == FCmpInst::FCMP_TRUE &&
7960 isConstantLegalOrBeforeLegalizer(Ty: CmpTy)) {
7961 auto True =
7962 B.buildConstant(Res: CmpTy, Val: getICmpTrueVal(TLI: getTargetLowering(),
7963 IsVector: CmpTy.isVector() /*isVector*/,
7964 IsFP: true /*isFP*/));
7965 B.buildZExtOrTrunc(Res: DestReg, Op: True);
7966 } else { // We take the predicate without predicate optimizations.
7967 auto Cmp = B.buildFCmp(Pred, Res: CmpTy, Op0: LHS0, Op1: LHS1, Flags);
7968 B.buildZExtOrTrunc(Res: DestReg, Op: Cmp);
7969 }
7970 };
7971 return true;
7972 }
7973
7974 return false;
7975}
7976
7977bool CombinerHelper::matchAnd(MachineInstr &MI, BuildFnTy &MatchInfo) const {
7978 GAnd *And = cast<GAnd>(Val: &MI);
7979
7980 if (tryFoldAndOrOrICmpsUsingRanges(Logic: And, MatchInfo))
7981 return true;
7982
7983 if (tryFoldLogicOfFCmps(Logic: And, MatchInfo))
7984 return true;
7985
7986 return false;
7987}
7988
7989bool CombinerHelper::matchOr(MachineInstr &MI, BuildFnTy &MatchInfo) const {
7990 GOr *Or = cast<GOr>(Val: &MI);
7991
7992 if (tryFoldAndOrOrICmpsUsingRanges(Logic: Or, MatchInfo))
7993 return true;
7994
7995 if (tryFoldLogicOfFCmps(Logic: Or, MatchInfo))
7996 return true;
7997
7998 return false;
7999}
8000
8001bool CombinerHelper::matchAddOverflow(MachineInstr &MI,
8002 BuildFnTy &MatchInfo) const {
8003 GAddCarryOut *Add = cast<GAddCarryOut>(Val: &MI);
8004
8005 // Addo has no flags
8006 Register Dst = Add->getReg(Idx: 0);
8007 Register Carry = Add->getReg(Idx: 1);
8008 Register LHS = Add->getLHSReg();
8009 Register RHS = Add->getRHSReg();
8010 bool IsSigned = Add->isSigned();
8011 LLT DstTy = MRI.getType(Reg: Dst);
8012 LLT CarryTy = MRI.getType(Reg: Carry);
8013
8014 // Fold addo, if the carry is dead -> add, undef.
8015 if (MRI.use_nodbg_empty(RegNo: Carry) &&
8016 isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_ADD, {DstTy}})) {
8017 MatchInfo = [=](MachineIRBuilder &B) {
8018 B.buildAdd(Dst, Src0: LHS, Src1: RHS);
8019 B.buildUndef(Res: Carry);
8020 };
8021 return true;
8022 }
8023
8024 // Canonicalize constant to RHS.
8025 if (isConstantOrConstantVectorI(Src: LHS) && !isConstantOrConstantVectorI(Src: RHS)) {
8026 if (IsSigned) {
8027 MatchInfo = [=](MachineIRBuilder &B) {
8028 B.buildSAddo(Res: Dst, CarryOut: Carry, Op0: RHS, Op1: LHS);
8029 };
8030 return true;
8031 }
8032 // !IsSigned
8033 MatchInfo = [=](MachineIRBuilder &B) {
8034 B.buildUAddo(Res: Dst, CarryOut: Carry, Op0: RHS, Op1: LHS);
8035 };
8036 return true;
8037 }
8038
8039 std::optional<APInt> MaybeLHS = getConstantOrConstantSplatVector(Src: LHS);
8040 std::optional<APInt> MaybeRHS = getConstantOrConstantSplatVector(Src: RHS);
8041
8042 // Fold addo(c1, c2) -> c3, carry.
8043 if (MaybeLHS && MaybeRHS && isConstantLegalOrBeforeLegalizer(Ty: DstTy) &&
8044 isConstantLegalOrBeforeLegalizer(Ty: CarryTy)) {
8045 bool Overflow;
8046 APInt Result = IsSigned ? MaybeLHS->sadd_ov(RHS: *MaybeRHS, Overflow)
8047 : MaybeLHS->uadd_ov(RHS: *MaybeRHS, Overflow);
8048 MatchInfo = [=](MachineIRBuilder &B) {
8049 B.buildConstant(Res: Dst, Val: Result);
8050 B.buildConstant(Res: Carry, Val: Overflow);
8051 };
8052 return true;
8053 }
8054
8055 // Fold (addo x, 0) -> x, no carry
8056 if (MaybeRHS && *MaybeRHS == 0 && isConstantLegalOrBeforeLegalizer(Ty: CarryTy)) {
8057 MatchInfo = [=](MachineIRBuilder &B) {
8058 B.buildCopy(Res: Dst, Op: LHS);
8059 B.buildConstant(Res: Carry, Val: 0);
8060 };
8061 return true;
8062 }
8063
8064 // Given 2 constant operands whose sum does not overflow:
8065 // uaddo (X +nuw C0), C1 -> uaddo X, C0 + C1
8066 // saddo (X +nsw C0), C1 -> saddo X, C0 + C1
8067 GAdd *AddLHS = getOpcodeDef<GAdd>(Reg: LHS, MRI);
8068 if (MaybeRHS && AddLHS && MRI.hasOneNonDBGUse(RegNo: Add->getReg(Idx: 0)) &&
8069 ((IsSigned && AddLHS->getFlag(Flag: MachineInstr::MIFlag::NoSWrap)) ||
8070 (!IsSigned && AddLHS->getFlag(Flag: MachineInstr::MIFlag::NoUWrap)))) {
8071 std::optional<APInt> MaybeAddRHS =
8072 getConstantOrConstantSplatVector(Src: AddLHS->getRHSReg());
8073 if (MaybeAddRHS) {
8074 bool Overflow;
8075 APInt NewC = IsSigned ? MaybeAddRHS->sadd_ov(RHS: *MaybeRHS, Overflow)
8076 : MaybeAddRHS->uadd_ov(RHS: *MaybeRHS, Overflow);
8077 if (!Overflow && isConstantLegalOrBeforeLegalizer(Ty: DstTy)) {
8078 if (IsSigned) {
8079 MatchInfo = [=](MachineIRBuilder &B) {
8080 auto ConstRHS = B.buildConstant(Res: DstTy, Val: NewC);
8081 B.buildSAddo(Res: Dst, CarryOut: Carry, Op0: AddLHS->getLHSReg(), Op1: ConstRHS);
8082 };
8083 return true;
8084 }
8085 // !IsSigned
8086 MatchInfo = [=](MachineIRBuilder &B) {
8087 auto ConstRHS = B.buildConstant(Res: DstTy, Val: NewC);
8088 B.buildUAddo(Res: Dst, CarryOut: Carry, Op0: AddLHS->getLHSReg(), Op1: ConstRHS);
8089 };
8090 return true;
8091 }
8092 }
8093 };
8094
8095 // We try to combine addo to non-overflowing add.
8096 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_ADD, {DstTy}}) ||
8097 !isConstantLegalOrBeforeLegalizer(Ty: CarryTy))
8098 return false;
8099
8100 // We try to combine uaddo to non-overflowing add.
8101 if (!IsSigned) {
8102 ConstantRange CRLHS =
8103 ConstantRange::fromKnownBits(Known: VT->getKnownBits(R: LHS), /*IsSigned=*/false);
8104 ConstantRange CRRHS =
8105 ConstantRange::fromKnownBits(Known: VT->getKnownBits(R: RHS), /*IsSigned=*/false);
8106
8107 switch (CRLHS.unsignedAddMayOverflow(Other: CRRHS)) {
8108 case ConstantRange::OverflowResult::MayOverflow:
8109 return false;
8110 case ConstantRange::OverflowResult::NeverOverflows: {
8111 MatchInfo = [=](MachineIRBuilder &B) {
8112 B.buildAdd(Dst, Src0: LHS, Src1: RHS, Flags: MachineInstr::MIFlag::NoUWrap);
8113 B.buildConstant(Res: Carry, Val: 0);
8114 };
8115 return true;
8116 }
8117 case ConstantRange::OverflowResult::AlwaysOverflowsLow:
8118 case ConstantRange::OverflowResult::AlwaysOverflowsHigh: {
8119 MatchInfo = [=](MachineIRBuilder &B) {
8120 B.buildAdd(Dst, Src0: LHS, Src1: RHS);
8121 B.buildConstant(Res: Carry, Val: 1);
8122 };
8123 return true;
8124 }
8125 }
8126 return false;
8127 }
8128
8129 // We try to combine saddo to non-overflowing add.
8130
8131 // If LHS and RHS each have at least two sign bits, then there is no signed
8132 // overflow.
8133 if (VT->computeNumSignBits(R: RHS) > 1 && VT->computeNumSignBits(R: LHS) > 1) {
8134 MatchInfo = [=](MachineIRBuilder &B) {
8135 B.buildAdd(Dst, Src0: LHS, Src1: RHS, Flags: MachineInstr::MIFlag::NoSWrap);
8136 B.buildConstant(Res: Carry, Val: 0);
8137 };
8138 return true;
8139 }
8140
8141 ConstantRange CRLHS =
8142 ConstantRange::fromKnownBits(Known: VT->getKnownBits(R: LHS), /*IsSigned=*/true);
8143 ConstantRange CRRHS =
8144 ConstantRange::fromKnownBits(Known: VT->getKnownBits(R: RHS), /*IsSigned=*/true);
8145
8146 switch (CRLHS.signedAddMayOverflow(Other: CRRHS)) {
8147 case ConstantRange::OverflowResult::MayOverflow:
8148 return false;
8149 case ConstantRange::OverflowResult::NeverOverflows: {
8150 MatchInfo = [=](MachineIRBuilder &B) {
8151 B.buildAdd(Dst, Src0: LHS, Src1: RHS, Flags: MachineInstr::MIFlag::NoSWrap);
8152 B.buildConstant(Res: Carry, Val: 0);
8153 };
8154 return true;
8155 }
8156 case ConstantRange::OverflowResult::AlwaysOverflowsLow:
8157 case ConstantRange::OverflowResult::AlwaysOverflowsHigh: {
8158 MatchInfo = [=](MachineIRBuilder &B) {
8159 B.buildAdd(Dst, Src0: LHS, Src1: RHS);
8160 B.buildConstant(Res: Carry, Val: 1);
8161 };
8162 return true;
8163 }
8164 }
8165
8166 return false;
8167}
8168
8169void CombinerHelper::applyBuildFnMO(const MachineOperand &MO,
8170 BuildFnTy &MatchInfo) const {
8171 MachineInstr *Root = getDefIgnoringCopies(Reg: MO.getReg(), MRI);
8172 MatchInfo(Builder);
8173 Root->eraseFromParent();
8174}
8175
8176bool CombinerHelper::matchFPowIExpansion(MachineInstr &MI,
8177 int64_t Exponent) const {
8178 bool OptForSize = MI.getMF()->getFunction().hasOptSize();
8179 return getTargetLowering().isBeneficialToExpandPowI(Exponent, OptForSize);
8180}
8181
8182void CombinerHelper::applyExpandFPowI(MachineInstr &MI,
8183 int64_t Exponent) const {
8184 auto [Dst, Base] = MI.getFirst2Regs();
8185 LLT Ty = MRI.getType(Reg: Dst);
8186 int64_t ExpVal = Exponent;
8187
8188 if (ExpVal == 0) {
8189 Builder.buildFConstant(Res: Dst, Val: 1.0);
8190 MI.removeFromParent();
8191 return;
8192 }
8193
8194 if (ExpVal < 0)
8195 ExpVal = -ExpVal;
8196
8197 // We use the simple binary decomposition method from SelectionDAG ExpandPowI
8198 // to generate the multiply sequence. There are more optimal ways to do this
8199 // (for example, powi(x,15) generates one more multiply than it should), but
8200 // this has the benefit of being both really simple and much better than a
8201 // libcall.
8202 std::optional<SrcOp> Res;
8203 SrcOp CurSquare = Base;
8204 while (ExpVal > 0) {
8205 if (ExpVal & 1) {
8206 if (!Res)
8207 Res = CurSquare;
8208 else
8209 Res = Builder.buildFMul(Dst: Ty, Src0: *Res, Src1: CurSquare);
8210 }
8211
8212 CurSquare = Builder.buildFMul(Dst: Ty, Src0: CurSquare, Src1: CurSquare);
8213 ExpVal >>= 1;
8214 }
8215
8216 // If the original exponent was negative, invert the result, producing
8217 // 1/(x*x*x).
8218 if (Exponent < 0)
8219 Res = Builder.buildFDiv(Dst: Ty, Src0: Builder.buildFConstant(Res: Ty, Val: 1.0), Src1: *Res,
8220 Flags: MI.getFlags());
8221
8222 Builder.buildCopy(Res: Dst, Op: *Res);
8223 MI.eraseFromParent();
8224}
8225
8226bool CombinerHelper::matchFoldAPlusC1MinusC2(const MachineInstr &MI,
8227 BuildFnTy &MatchInfo) const {
8228 // fold (A+C1)-C2 -> A+(C1-C2)
8229 const GSub *Sub = cast<GSub>(Val: &MI);
8230 Register A, C1Reg;
8231 if (!mi_match(R: Sub->getLHSReg(), MRI, P: m_GAdd(L: m_Reg(R&: A), R: m_Reg(R&: C1Reg))))
8232 return false;
8233
8234 if (!MRI.hasOneNonDBGUse(RegNo: Sub->getLHSReg()))
8235 return false;
8236
8237 APInt C2 = getIConstantFromReg(VReg: Sub->getRHSReg(), MRI);
8238 APInt C1 = getIConstantFromReg(VReg: C1Reg, MRI);
8239
8240 Register Dst = Sub->getReg(Idx: 0);
8241 LLT DstTy = MRI.getType(Reg: Dst);
8242
8243 MatchInfo = [=](MachineIRBuilder &B) {
8244 auto Const = B.buildConstant(Res: DstTy, Val: C1 - C2);
8245 B.buildAdd(Dst, Src0: A, Src1: Const);
8246 };
8247
8248 return true;
8249}
8250
8251bool CombinerHelper::matchFoldC2MinusAPlusC1(const MachineInstr &MI,
8252 BuildFnTy &MatchInfo) const {
8253 // fold C2-(A+C1) -> (C2-C1)-A
8254 const GSub *Sub = cast<GSub>(Val: &MI);
8255 Register A, C1Reg;
8256 if (!mi_match(R: Sub->getRHSReg(), MRI, P: m_GAdd(L: m_Reg(R&: A), R: m_Reg(R&: C1Reg))))
8257 return false;
8258
8259 if (!MRI.hasOneNonDBGUse(RegNo: Sub->getRHSReg()))
8260 return false;
8261
8262 APInt C2 = getIConstantFromReg(VReg: Sub->getLHSReg(), MRI);
8263 APInt C1 = getIConstantFromReg(VReg: C1Reg, MRI);
8264
8265 Register Dst = Sub->getReg(Idx: 0);
8266 LLT DstTy = MRI.getType(Reg: Dst);
8267
8268 MatchInfo = [=](MachineIRBuilder &B) {
8269 auto Const = B.buildConstant(Res: DstTy, Val: C2 - C1);
8270 B.buildSub(Dst, Src0: Const, Src1: A);
8271 };
8272
8273 return true;
8274}
8275
8276bool CombinerHelper::matchFoldAMinusC1MinusC2(const MachineInstr &MI,
8277 BuildFnTy &MatchInfo) const {
8278 // fold (A-C1)-C2 -> A-(C1+C2)
8279 const GSub *Sub1 = cast<GSub>(Val: &MI);
8280 Register A, C1Reg;
8281 if (!mi_match(R: Sub1->getLHSReg(), MRI, P: m_GSub(L: m_Reg(R&: A), R: m_Reg(R&: C1Reg))))
8282 return false;
8283
8284 if (!MRI.hasOneNonDBGUse(RegNo: Sub1->getLHSReg()))
8285 return false;
8286
8287 APInt C2 = getIConstantFromReg(VReg: Sub1->getRHSReg(), MRI);
8288 APInt C1 = getIConstantFromReg(VReg: C1Reg, MRI);
8289
8290 Register Dst = Sub1->getReg(Idx: 0);
8291 LLT DstTy = MRI.getType(Reg: Dst);
8292
8293 MatchInfo = [=](MachineIRBuilder &B) {
8294 auto Const = B.buildConstant(Res: DstTy, Val: C1 + C2);
8295 B.buildSub(Dst, Src0: A, Src1: Const);
8296 };
8297
8298 return true;
8299}
8300
8301bool CombinerHelper::matchFoldC1Minus2MinusC2(const MachineInstr &MI,
8302 BuildFnTy &MatchInfo) const {
8303 // fold (C1-A)-C2 -> (C1-C2)-A
8304 const GSub *Sub1 = cast<GSub>(Val: &MI);
8305 Register C1Reg, A;
8306 if (!mi_match(R: Sub1->getLHSReg(), MRI, P: m_GSub(L: m_Reg(R&: C1Reg), R: m_Reg(R&: A))))
8307 return false;
8308
8309 if (!MRI.hasOneNonDBGUse(RegNo: Sub1->getLHSReg()))
8310 return false;
8311
8312 APInt C2 = getIConstantFromReg(VReg: Sub1->getRHSReg(), MRI);
8313 APInt C1 = getIConstantFromReg(VReg: C1Reg, MRI);
8314
8315 Register Dst = Sub1->getReg(Idx: 0);
8316 LLT DstTy = MRI.getType(Reg: Dst);
8317
8318 MatchInfo = [=](MachineIRBuilder &B) {
8319 auto Const = B.buildConstant(Res: DstTy, Val: C1 - C2);
8320 B.buildSub(Dst, Src0: Const, Src1: A);
8321 };
8322
8323 return true;
8324}
8325
8326bool CombinerHelper::matchFoldAMinusC1PlusC2(const MachineInstr &MI,
8327 BuildFnTy &MatchInfo) const {
8328 // fold ((A-C1)+C2) -> (A+(C2-C1))
8329 const GAdd *Add = cast<GAdd>(Val: &MI);
8330 Register A, C1Reg;
8331 if (!mi_match(R: Add->getLHSReg(), MRI, P: m_GSub(L: m_Reg(R&: A), R: m_Reg(R&: C1Reg))))
8332 return false;
8333
8334 if (!MRI.hasOneNonDBGUse(RegNo: Add->getLHSReg()))
8335 return false;
8336
8337 APInt C2 = getIConstantFromReg(VReg: Add->getRHSReg(), MRI);
8338 APInt C1 = getIConstantFromReg(VReg: C1Reg, MRI);
8339
8340 Register Dst = Add->getReg(Idx: 0);
8341 LLT DstTy = MRI.getType(Reg: Dst);
8342
8343 MatchInfo = [=](MachineIRBuilder &B) {
8344 auto Const = B.buildConstant(Res: DstTy, Val: C2 - C1);
8345 B.buildAdd(Dst, Src0: A, Src1: Const);
8346 };
8347
8348 return true;
8349}
8350
8351bool CombinerHelper::matchUnmergeValuesAnyExtBuildVector(
8352 const MachineInstr &MI, BuildFnTy &MatchInfo) const {
8353 const GUnmerge *Unmerge = cast<GUnmerge>(Val: &MI);
8354
8355 if (!MRI.hasOneNonDBGUse(RegNo: Unmerge->getSourceReg()))
8356 return false;
8357
8358 LLT DstTy = MRI.getType(Reg: Unmerge->getReg(Idx: 0));
8359
8360 // $bv:_(<8 x s8>) = G_BUILD_VECTOR ....
8361 // $any:_(<8 x s16>) = G_ANYEXT $bv
8362 // $uv:_(<4 x s16>), $uv1:_(<4 x s16>) = G_UNMERGE_VALUES $any
8363 //
8364 // ->
8365 //
8366 // $any:_(s16) = G_ANYEXT $bv[0]
8367 // $any1:_(s16) = G_ANYEXT $bv[1]
8368 // $any2:_(s16) = G_ANYEXT $bv[2]
8369 // $any3:_(s16) = G_ANYEXT $bv[3]
8370 // $any4:_(s16) = G_ANYEXT $bv[4]
8371 // $any5:_(s16) = G_ANYEXT $bv[5]
8372 // $any6:_(s16) = G_ANYEXT $bv[6]
8373 // $any7:_(s16) = G_ANYEXT $bv[7]
8374 // $uv:_(<4 x s16>) = G_BUILD_VECTOR $any, $any1, $any2, $any3
8375 // $uv1:_(<4 x s16>) = G_BUILD_VECTOR $any4, $any5, $any6, $any7
8376
8377 // We want to unmerge into vectors.
8378 if (!DstTy.isFixedVector())
8379 return false;
8380
8381 Register AnySrcReg;
8382 if (!mi_match(R: Unmerge->getSourceReg(), MRI, P: m_GAnyExt(Src: m_Reg(R&: AnySrcReg))))
8383 return false;
8384
8385 GBuildVector *BV;
8386 if (mi_match(R: AnySrcReg, MRI, P: m_GBuildVector(Inst&: BV))) {
8387 // G_UNMERGE_VALUES G_ANYEXT G_BUILD_VECTOR
8388
8389 if (!MRI.hasOneNonDBGUse(RegNo: BV->getReg(Idx: 0)))
8390 return false;
8391
8392 // FIXME: check element types?
8393 if (BV->getNumSources() % Unmerge->getNumDefs() != 0)
8394 return false;
8395
8396 LLT BigBvTy = MRI.getType(Reg: BV->getReg(Idx: 0));
8397 LLT SmallBvTy = DstTy;
8398 LLT SmallBvElemenTy = SmallBvTy.getElementType();
8399
8400 if (!isLegalOrBeforeLegalizer(
8401 Query: {TargetOpcode::G_BUILD_VECTOR, {SmallBvTy, SmallBvElemenTy}}))
8402 return false;
8403
8404 // We check the legality of scalar anyext.
8405 if (!isLegalOrBeforeLegalizer(
8406 Query: {TargetOpcode::G_ANYEXT,
8407 {SmallBvElemenTy, BigBvTy.getElementType()}}))
8408 return false;
8409
8410 MatchInfo = [=](MachineIRBuilder &B) {
8411 // Build into each G_UNMERGE_VALUES def
8412 // a small build vector with anyext from the source build vector.
8413 for (unsigned I = 0; I < Unmerge->getNumDefs(); ++I) {
8414 SmallVector<Register> Ops;
8415 for (unsigned J = 0; J < SmallBvTy.getNumElements(); ++J) {
8416 Register SourceArray =
8417 BV->getSourceReg(I: I * SmallBvTy.getNumElements() + J);
8418 auto AnyExt = B.buildAnyExt(Res: SmallBvElemenTy, Op: SourceArray);
8419 Ops.push_back(Elt: AnyExt.getReg(Idx: 0));
8420 }
8421 B.buildBuildVector(Res: Unmerge->getOperand(i: I).getReg(), Ops);
8422 };
8423 };
8424 return true;
8425 };
8426
8427 return false;
8428}
8429
8430bool CombinerHelper::matchShuffleUndefRHS(MachineInstr &MI,
8431 BuildFnTy &MatchInfo) const {
8432
8433 bool Changed = false;
8434 auto &Shuffle = cast<GShuffleVector>(Val&: MI);
8435 ArrayRef<int> OrigMask = Shuffle.getMask();
8436 SmallVector<int, 16> NewMask;
8437 const LLT SrcTy = MRI.getType(Reg: Shuffle.getSrc1Reg());
8438 const unsigned NumSrcElems = SrcTy.isVector() ? SrcTy.getNumElements() : 1;
8439 const unsigned NumDstElts = OrigMask.size();
8440 for (unsigned i = 0; i != NumDstElts; ++i) {
8441 int Idx = OrigMask[i];
8442 if (Idx >= (int)NumSrcElems) {
8443 Idx = -1;
8444 Changed = true;
8445 }
8446 NewMask.push_back(Elt: Idx);
8447 }
8448
8449 if (!Changed)
8450 return false;
8451
8452 MatchInfo = [&, NewMask = std::move(NewMask)](MachineIRBuilder &B) {
8453 B.buildShuffleVector(Res: MI.getOperand(i: 0), Src1: MI.getOperand(i: 1), Src2: MI.getOperand(i: 2),
8454 Mask: std::move(NewMask));
8455 };
8456
8457 return true;
8458}
8459
8460static void commuteMask(MutableArrayRef<int> Mask, const unsigned NumElems) {
8461 const unsigned MaskSize = Mask.size();
8462 for (unsigned I = 0; I < MaskSize; ++I) {
8463 int Idx = Mask[I];
8464 if (Idx < 0)
8465 continue;
8466
8467 if (Idx < (int)NumElems)
8468 Mask[I] = Idx + NumElems;
8469 else
8470 Mask[I] = Idx - NumElems;
8471 }
8472}
8473
8474bool CombinerHelper::matchShuffleDisjointMask(MachineInstr &MI,
8475 BuildFnTy &MatchInfo) const {
8476
8477 auto &Shuffle = cast<GShuffleVector>(Val&: MI);
8478 // If any of the two inputs is already undef, don't check the mask again to
8479 // prevent infinite loop
8480 if (getOpcodeDef(Opcode: TargetOpcode::G_IMPLICIT_DEF, Reg: Shuffle.getSrc1Reg(), MRI))
8481 return false;
8482
8483 if (getOpcodeDef(Opcode: TargetOpcode::G_IMPLICIT_DEF, Reg: Shuffle.getSrc2Reg(), MRI))
8484 return false;
8485
8486 const LLT DstTy = MRI.getType(Reg: Shuffle.getReg(Idx: 0));
8487 const LLT Src1Ty = MRI.getType(Reg: Shuffle.getSrc1Reg());
8488 if (!isLegalOrBeforeLegalizer(
8489 Query: {TargetOpcode::G_SHUFFLE_VECTOR, {DstTy, Src1Ty}}))
8490 return false;
8491
8492 ArrayRef<int> Mask = Shuffle.getMask();
8493 const unsigned NumSrcElems = Src1Ty.getNumElements();
8494
8495 bool TouchesSrc1 = false;
8496 bool TouchesSrc2 = false;
8497 const unsigned NumElems = Mask.size();
8498 for (unsigned Idx = 0; Idx < NumElems; ++Idx) {
8499 if (Mask[Idx] < 0)
8500 continue;
8501
8502 if (Mask[Idx] < (int)NumSrcElems)
8503 TouchesSrc1 = true;
8504 else
8505 TouchesSrc2 = true;
8506 }
8507
8508 if (TouchesSrc1 == TouchesSrc2)
8509 return false;
8510
8511 Register NewSrc1 = Shuffle.getSrc1Reg();
8512 SmallVector<int, 16> NewMask(Mask);
8513 if (TouchesSrc2) {
8514 NewSrc1 = Shuffle.getSrc2Reg();
8515 commuteMask(Mask: NewMask, NumElems: NumSrcElems);
8516 }
8517
8518 MatchInfo = [=, &Shuffle](MachineIRBuilder &B) {
8519 auto Undef = B.buildUndef(Res: Src1Ty);
8520 B.buildShuffleVector(Res: Shuffle.getReg(Idx: 0), Src1: NewSrc1, Src2: Undef, Mask: NewMask);
8521 };
8522
8523 return true;
8524}
8525
8526bool CombinerHelper::matchSuboCarryOut(const MachineInstr &MI,
8527 BuildFnTy &MatchInfo) const {
8528 const GSubCarryOut *Subo = cast<GSubCarryOut>(Val: &MI);
8529
8530 Register Dst = Subo->getReg(Idx: 0);
8531 Register LHS = Subo->getLHSReg();
8532 Register RHS = Subo->getRHSReg();
8533 Register Carry = Subo->getCarryOutReg();
8534 LLT DstTy = MRI.getType(Reg: Dst);
8535 LLT CarryTy = MRI.getType(Reg: Carry);
8536
8537 // Check legality before known bits.
8538 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_SUB, {DstTy}}) ||
8539 !isConstantLegalOrBeforeLegalizer(Ty: CarryTy))
8540 return false;
8541
8542 ConstantRange KBLHS =
8543 ConstantRange::fromKnownBits(Known: VT->getKnownBits(R: LHS),
8544 /* IsSigned=*/Subo->isSigned());
8545 ConstantRange KBRHS =
8546 ConstantRange::fromKnownBits(Known: VT->getKnownBits(R: RHS),
8547 /* IsSigned=*/Subo->isSigned());
8548
8549 if (Subo->isSigned()) {
8550 // G_SSUBO
8551 switch (KBLHS.signedSubMayOverflow(Other: KBRHS)) {
8552 case ConstantRange::OverflowResult::MayOverflow:
8553 return false;
8554 case ConstantRange::OverflowResult::NeverOverflows: {
8555 MatchInfo = [=](MachineIRBuilder &B) {
8556 B.buildSub(Dst, Src0: LHS, Src1: RHS, Flags: MachineInstr::MIFlag::NoSWrap);
8557 B.buildConstant(Res: Carry, Val: 0);
8558 };
8559 return true;
8560 }
8561 case ConstantRange::OverflowResult::AlwaysOverflowsLow:
8562 case ConstantRange::OverflowResult::AlwaysOverflowsHigh: {
8563 MatchInfo = [=](MachineIRBuilder &B) {
8564 B.buildSub(Dst, Src0: LHS, Src1: RHS);
8565 B.buildConstant(Res: Carry, Val: getICmpTrueVal(TLI: getTargetLowering(),
8566 /*isVector=*/IsVector: CarryTy.isVector(),
8567 /*isFP=*/IsFP: false));
8568 };
8569 return true;
8570 }
8571 }
8572 return false;
8573 }
8574
8575 // G_USUBO
8576 switch (KBLHS.unsignedSubMayOverflow(Other: KBRHS)) {
8577 case ConstantRange::OverflowResult::MayOverflow:
8578 return false;
8579 case ConstantRange::OverflowResult::NeverOverflows: {
8580 MatchInfo = [=](MachineIRBuilder &B) {
8581 B.buildSub(Dst, Src0: LHS, Src1: RHS, Flags: MachineInstr::MIFlag::NoUWrap);
8582 B.buildConstant(Res: Carry, Val: 0);
8583 };
8584 return true;
8585 }
8586 case ConstantRange::OverflowResult::AlwaysOverflowsLow:
8587 case ConstantRange::OverflowResult::AlwaysOverflowsHigh: {
8588 MatchInfo = [=](MachineIRBuilder &B) {
8589 B.buildSub(Dst, Src0: LHS, Src1: RHS);
8590 B.buildConstant(Res: Carry, Val: getICmpTrueVal(TLI: getTargetLowering(),
8591 /*isVector=*/IsVector: CarryTy.isVector(),
8592 /*isFP=*/IsFP: false));
8593 };
8594 return true;
8595 }
8596 }
8597
8598 return false;
8599}
8600
8601// Fold (ctlz (xor x, (sra x, bitwidth-1))) -> (add (ctls x), 1).
8602// Fold (ctlz (or (shl (xor x, (sra x, bitwidth-1)), 1), 1) -> (ctls x)
8603bool CombinerHelper::matchCtls(MachineInstr &CtlzMI,
8604 BuildFnTy &MatchInfo) const {
8605 assert((CtlzMI.getOpcode() == TargetOpcode::G_CTLZ ||
8606 CtlzMI.getOpcode() == TargetOpcode::G_CTLZ_ZERO_POISON) &&
8607 "Expected G_CTLZ variant");
8608
8609 const Register Dst = CtlzMI.getOperand(i: 0).getReg();
8610 Register Src = CtlzMI.getOperand(i: 1).getReg();
8611
8612 LLT Ty = MRI.getType(Reg: Dst);
8613 LLT SrcTy = MRI.getType(Reg: Src);
8614
8615 if (!(Ty.isValid() && Ty.isScalar()))
8616 return false;
8617
8618 if (!LI)
8619 return false;
8620
8621 SmallVector<LLT, 2> QueryTypes = {Ty, SrcTy};
8622 LegalityQuery Query(TargetOpcode::G_CTLS, QueryTypes);
8623
8624 switch (LI->getAction(Query).Action) {
8625 default:
8626 return false;
8627 case LegalizeActions::Legal:
8628 case LegalizeActions::Custom:
8629 case LegalizeActions::WidenScalar:
8630 break;
8631 }
8632
8633 // Src = or(shl(V, 1), 1) -> Src=V; NeedAdd = False
8634 Register V;
8635 bool NeedAdd = true;
8636 if (mi_match(R: Src, MRI,
8637 P: m_OneUse(SP: m_GOr(L: m_OneUse(SP: m_GShl(L: m_Reg(R&: V), R: m_SpecificICst(RequestedValue: 1))),
8638 R: m_SpecificICst(RequestedValue: 1))))) {
8639 NeedAdd = false;
8640 Src = V;
8641 }
8642
8643 unsigned BitWidth = Ty.getScalarSizeInBits();
8644
8645 Register X;
8646 if (!mi_match(R: Src, MRI,
8647 P: m_OneUse(SP: m_GXor(L: m_Reg(R&: X), R: m_OneUse(SP: m_GAShr(
8648 L: m_DeferredReg(R&: X),
8649 R: m_SpecificICst(RequestedValue: BitWidth - 1)))))))
8650 return false;
8651
8652 MatchInfo = [=](MachineIRBuilder &B) {
8653 if (!NeedAdd) {
8654 B.buildCTLS(Dst, Src0: X);
8655 return;
8656 }
8657
8658 auto Ctls = B.buildCTLS(Dst: Ty, Src0: X);
8659 auto One = B.buildConstant(Res: Ty, Val: 1);
8660
8661 B.buildAdd(Dst, Src0: Ctls, Src1: One);
8662 };
8663
8664 return true;
8665}
8666
8667// Fold shr ( add ( ext X, ext Y ), 1 ) -> avgfloor ( x, y )
8668// Fold shr ( add ( ext X, ext Y, 1 ), 1 ) -> avgceil ( x, y )
8669bool CombinerHelper::matchAVG(MachineInstr &MI, MachineRegisterInfo &MRI,
8670 Register X, Register Y,
8671 unsigned TargetOpc) const {
8672 assert((MI.getOpcode() == TargetOpcode::G_LSHR ||
8673 MI.getOpcode() == TargetOpcode::G_ASHR) &&
8674 "Expected G_LSHR/G_ASHR");
8675
8676 LLT XTy = MRI.getType(Reg: X);
8677 return XTy == MRI.getType(Reg: Y) && isLegal(Query: {TargetOpc, {XTy}});
8678}
8679
8680static unsigned getCountZeroPoisonOpcode(const MachineInstr &MI) {
8681 assert((MI.getOpcode() == TargetOpcode::G_CTLZ ||
8682 MI.getOpcode() == TargetOpcode::G_CTTZ) &&
8683 "Expected count-zero opcode");
8684 switch (MI.getOpcode()) {
8685 case TargetOpcode::G_CTLZ:
8686 return TargetOpcode::G_CTLZ_ZERO_POISON;
8687 case TargetOpcode::G_CTTZ:
8688 return TargetOpcode::G_CTTZ_ZERO_POISON;
8689 default:
8690 llvm_unreachable("Unexpected count-zero opcode");
8691 }
8692}
8693
8694bool CombinerHelper::matchCountZeroToZeroPoison(MachineInstr &MI) const {
8695 if (!VT)
8696 return false;
8697
8698 unsigned ZPOpc = getCountZeroPoisonOpcode(MI);
8699 Register Src = MI.getOperand(i: 1).getReg();
8700 if (!VT->isKnownNeverZero(R: Src))
8701 return false;
8702
8703 LLT DstTy = MRI.getType(Reg: MI.getOperand(i: 0).getReg());
8704 LLT SrcTy = MRI.getType(Reg: Src);
8705 return isLegalOrBeforeLegalizer(Query: {ZPOpc, {DstTy, SrcTy}});
8706}
8707
8708void CombinerHelper::applyCountZeroToZeroPoison(MachineInstr &MI) const {
8709 replaceOpcodeWith(FromMI&: MI, ToOpcode: getCountZeroPoisonOpcode(MI));
8710}
8711