1//===- CombinerHelperCasts.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//
9// This file implements CombinerHelper for G_ANYEXT, G_SEXT, G_TRUNC, and
10// G_ZEXT
11//
12//===----------------------------------------------------------------------===//
13#include "llvm/CodeGen/GlobalISel/CombinerHelper.h"
14#include "llvm/CodeGen/GlobalISel/LegalizerHelper.h"
15#include "llvm/CodeGen/GlobalISel/LegalizerInfo.h"
16#include "llvm/CodeGen/GlobalISel/MIPatternMatch.h"
17#include "llvm/CodeGen/GlobalISel/MachineIRBuilder.h"
18#include "llvm/CodeGen/GlobalISel/Utils.h"
19#include "llvm/CodeGen/LowLevelTypeUtils.h"
20#include "llvm/CodeGen/MachineOperand.h"
21#include "llvm/CodeGen/MachineRegisterInfo.h"
22#include "llvm/CodeGen/TargetOpcodes.h"
23#include "llvm/Support/Casting.h"
24
25#define DEBUG_TYPE "gi-combiner"
26
27using namespace llvm;
28using namespace MIPatternMatch;
29
30bool CombinerHelper::matchSextOfTrunc(const MachineOperand &MO,
31 BuildFnTy &MatchInfo) const {
32 GSext *Sext = cast<GSext>(Val: getDefIgnoringCopies(Reg: MO.getReg(), MRI));
33 GTrunc *Trunc = cast<GTrunc>(Val: getDefIgnoringCopies(Reg: Sext->getSrcReg(), MRI));
34
35 Register Dst = Sext->getReg(Idx: 0);
36 Register Src = Trunc->getSrcReg();
37
38 LLT DstTy = MRI.getType(Reg: Dst);
39 LLT SrcTy = MRI.getType(Reg: Src);
40
41 // Combines without nsw trunc.
42 if (!Trunc->getFlag(Flag: MachineInstr::NoSWrap)) {
43 // Do this for 8 bit values and up. We don't want to do it for e.g. G_TRUNC
44 // to i1.
45 unsigned TruncWidth = MRI.getType(Reg: Trunc->getReg(Idx: 0)).getScalarSizeInBits();
46 if (TruncWidth < 8)
47 return false;
48
49 if (DstTy != SrcTy ||
50 !isLegalOrBeforeLegalizer(
51 Query: {TargetOpcode::G_SEXT_INREG, {DstTy, SrcTy}, {}, {TruncWidth}}))
52 return false;
53
54 MatchInfo = [=](MachineIRBuilder &B) {
55 B.buildSExtInReg(Res: Dst, Op: Src, ImmOp: TruncWidth);
56 };
57 return true;
58 }
59
60 // Combines for nsw trunc.
61
62 if (DstTy == SrcTy) {
63 MatchInfo = [=](MachineIRBuilder &B) { B.buildCopy(Res: Dst, Op: Src); };
64 return true;
65 }
66
67 if (DstTy.getScalarSizeInBits() < SrcTy.getScalarSizeInBits() &&
68 isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_TRUNC, {DstTy, SrcTy}})) {
69 MatchInfo = [=](MachineIRBuilder &B) {
70 B.buildTrunc(Res: Dst, Op: Src, Flags: MachineInstr::MIFlag::NoSWrap);
71 };
72 return true;
73 }
74
75 if (DstTy.getScalarSizeInBits() > SrcTy.getScalarSizeInBits() &&
76 isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_SEXT, {DstTy, SrcTy}})) {
77 MatchInfo = [=](MachineIRBuilder &B) { B.buildSExt(Res: Dst, Op: Src); };
78 return true;
79 }
80
81 return false;
82}
83
84bool CombinerHelper::matchZextOfTrunc(const MachineOperand &MO,
85 BuildFnTy &MatchInfo) const {
86 GZext *Zext = cast<GZext>(Val: getDefIgnoringCopies(Reg: MO.getReg(), MRI));
87 GTrunc *Trunc = cast<GTrunc>(Val: getDefIgnoringCopies(Reg: Zext->getSrcReg(), MRI));
88
89 Register Dst = Zext->getReg(Idx: 0);
90 Register Src = Trunc->getSrcReg();
91
92 LLT DstTy = MRI.getType(Reg: Dst);
93 LLT SrcTy = MRI.getType(Reg: Src);
94
95 if (DstTy == SrcTy) {
96 MatchInfo = [=](MachineIRBuilder &B) { B.buildCopy(Res: Dst, Op: Src); };
97 return true;
98 }
99
100 if (DstTy.getScalarSizeInBits() < SrcTy.getScalarSizeInBits() &&
101 isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_TRUNC, {DstTy, SrcTy}})) {
102 MatchInfo = [=](MachineIRBuilder &B) {
103 B.buildTrunc(Res: Dst, Op: Src, Flags: MachineInstr::MIFlag::NoUWrap);
104 };
105 return true;
106 }
107
108 if (DstTy.getScalarSizeInBits() > SrcTy.getScalarSizeInBits() &&
109 isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_ZEXT, {DstTy, SrcTy}})) {
110 MatchInfo = [=](MachineIRBuilder &B) {
111 B.buildZExt(Res: Dst, Op: Src, Flags: MachineInstr::MIFlag::NonNeg);
112 };
113 return true;
114 }
115
116 return false;
117}
118
119bool CombinerHelper::matchNonNegZext(const MachineOperand &MO,
120 BuildFnTy &MatchInfo) const {
121 Register Dst = MO.getReg();
122 Register Src;
123 if (!mi_match(R: Dst, MRI, P: m_GZExt(Src: m_Reg(R&: Src))))
124 return false;
125
126 LLT DstTy = MRI.getType(Reg: Dst);
127 LLT SrcTy = MRI.getType(Reg: Src);
128 const auto &TLI = getTargetLowering();
129
130 // Convert zext nneg to sext if sext is the preferred form for the target.
131 if (isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_SEXT, {DstTy, SrcTy}}) &&
132 TLI.isSExtCheaperThanZExt(FromTy: getMVTForLLT(Ty: SrcTy), ToTy: getMVTForLLT(Ty: DstTy))) {
133 MatchInfo = [=](MachineIRBuilder &B) { B.buildSExt(Res: Dst, Op: Src); };
134 return true;
135 }
136
137 return false;
138}
139
140bool CombinerHelper::matchTruncateOfExt(const MachineInstr &Root,
141 const MachineInstr &ExtMI,
142 BuildFnTy &MatchInfo) const {
143 const GTrunc *Trunc = cast<GTrunc>(Val: &Root);
144 const GExtOp *Ext = cast<GExtOp>(Val: &ExtMI);
145
146 if (!MRI.hasOneNonDBGUse(RegNo: Ext->getReg(Idx: 0)))
147 return false;
148
149 Register Dst = Trunc->getReg(Idx: 0);
150 Register Src = Ext->getSrcReg();
151 LLT DstTy = MRI.getType(Reg: Dst);
152 LLT SrcTy = MRI.getType(Reg: Src);
153
154 if (SrcTy == DstTy) {
155 // The source and the destination are equally sized. We need to copy.
156 MatchInfo = [=](MachineIRBuilder &B) { B.buildCopy(Res: Dst, Op: Src); };
157
158 return true;
159 }
160
161 if (SrcTy.getScalarSizeInBits() < DstTy.getScalarSizeInBits()) {
162 // If the source is smaller than the destination, we need to extend.
163
164 if (!isLegalOrBeforeLegalizer(Query: {Ext->getOpcode(), {DstTy, SrcTy}}))
165 return false;
166
167 MatchInfo = [=](MachineIRBuilder &B) {
168 B.buildInstr(Opc: Ext->getOpcode(), DstOps: {Dst}, SrcOps: {Src});
169 };
170
171 return true;
172 }
173
174 if (SrcTy.getScalarSizeInBits() > DstTy.getScalarSizeInBits()) {
175 // If the source is larger than the destination, then we need to truncate.
176
177 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_TRUNC, {DstTy, SrcTy}}))
178 return false;
179
180 MatchInfo = [=](MachineIRBuilder &B) { B.buildTrunc(Res: Dst, Op: Src); };
181
182 return true;
183 }
184
185 return false;
186}
187
188bool CombinerHelper::isCastFree(unsigned Opcode, LLT ToTy, LLT FromTy) const {
189 const TargetLowering &TLI = getTargetLowering();
190 LLVMContext &Ctx = getContext();
191
192 switch (Opcode) {
193 case TargetOpcode::G_ANYEXT:
194 case TargetOpcode::G_ZEXT:
195 return TLI.isZExtFree(FromTy, ToTy, Ctx);
196 case TargetOpcode::G_TRUNC:
197 return TLI.isTruncateFree(FromTy, ToTy, Ctx);
198 default:
199 return false;
200 }
201}
202
203bool CombinerHelper::matchCastOfSelect(const MachineInstr &CastMI,
204 const MachineInstr &SelectMI,
205 BuildFnTy &MatchInfo) const {
206 const GExtOrTruncOp *Cast = cast<GExtOrTruncOp>(Val: &CastMI);
207 const GSelect *Select = cast<GSelect>(Val: &SelectMI);
208
209 if (!MRI.hasOneNonDBGUse(RegNo: Select->getReg(Idx: 0)))
210 return false;
211
212 Register Dst = Cast->getReg(Idx: 0);
213 LLT DstTy = MRI.getType(Reg: Dst);
214 LLT CondTy = MRI.getType(Reg: Select->getCondReg());
215 Register TrueReg = Select->getTrueReg();
216 Register FalseReg = Select->getFalseReg();
217 LLT SrcTy = MRI.getType(Reg: TrueReg);
218 Register Cond = Select->getCondReg();
219
220 if (!isLegalOrBeforeLegalizer(Query: {TargetOpcode::G_SELECT, {DstTy, CondTy}}))
221 return false;
222
223 if (!isCastFree(Opcode: Cast->getOpcode(), ToTy: DstTy, FromTy: SrcTy))
224 return false;
225
226 MatchInfo = [=](MachineIRBuilder &B) {
227 auto True = B.buildInstr(Opc: Cast->getOpcode(), DstOps: {DstTy}, SrcOps: {TrueReg});
228 auto False = B.buildInstr(Opc: Cast->getOpcode(), DstOps: {DstTy}, SrcOps: {FalseReg});
229 B.buildSelect(Res: Dst, Tst: Cond, Op0: True, Op1: False);
230 };
231
232 return true;
233}
234
235bool CombinerHelper::matchExtOfExt(const MachineInstr &FirstMI,
236 const MachineInstr &SecondMI,
237 BuildFnTy &MatchInfo) const {
238 const GExtOp *First = cast<GExtOp>(Val: &FirstMI);
239 const GExtOp *Second = cast<GExtOp>(Val: &SecondMI);
240
241 Register Dst = First->getReg(Idx: 0);
242 Register Src = Second->getSrcReg();
243 LLT DstTy = MRI.getType(Reg: Dst);
244 LLT SrcTy = MRI.getType(Reg: Src);
245
246 if (!MRI.hasOneNonDBGUse(RegNo: Second->getReg(Idx: 0)))
247 return false;
248
249 // ext of ext -> later ext
250 if (First->getOpcode() == Second->getOpcode() &&
251 isLegalOrBeforeLegalizer(Query: {Second->getOpcode(), {DstTy, SrcTy}})) {
252 if (Second->getOpcode() == TargetOpcode::G_ZEXT) {
253 MachineInstr::MIFlag Flag = MachineInstr::MIFlag::NoFlags;
254 if (Second->getFlag(Flag: MachineInstr::MIFlag::NonNeg))
255 Flag = MachineInstr::MIFlag::NonNeg;
256 MatchInfo = [=](MachineIRBuilder &B) { B.buildZExt(Res: Dst, Op: Src, Flags: Flag); };
257 return true;
258 }
259 // not zext -> no flags
260 MatchInfo = [=](MachineIRBuilder &B) {
261 B.buildInstr(Opc: Second->getOpcode(), DstOps: {Dst}, SrcOps: {Src});
262 };
263 return true;
264 }
265
266 // anyext of sext/zext -> sext/zext
267 // -> pick anyext as second ext, then ext of ext
268 if (First->getOpcode() == TargetOpcode::G_ANYEXT &&
269 isLegalOrBeforeLegalizer(Query: {Second->getOpcode(), {DstTy, SrcTy}})) {
270 if (Second->getOpcode() == TargetOpcode::G_ZEXT) {
271 MachineInstr::MIFlag Flag = MachineInstr::MIFlag::NoFlags;
272 if (Second->getFlag(Flag: MachineInstr::MIFlag::NonNeg))
273 Flag = MachineInstr::MIFlag::NonNeg;
274 MatchInfo = [=](MachineIRBuilder &B) { B.buildZExt(Res: Dst, Op: Src, Flags: Flag); };
275 return true;
276 }
277 MatchInfo = [=](MachineIRBuilder &B) { B.buildSExt(Res: Dst, Op: Src); };
278 return true;
279 }
280
281 // sext/zext of anyext -> sext/zext
282 // -> pick anyext as first ext, then ext of ext
283 if (Second->getOpcode() == TargetOpcode::G_ANYEXT &&
284 isLegalOrBeforeLegalizer(Query: {First->getOpcode(), {DstTy, SrcTy}})) {
285 if (First->getOpcode() == TargetOpcode::G_ZEXT) {
286 MachineInstr::MIFlag Flag = MachineInstr::MIFlag::NoFlags;
287 if (First->getFlag(Flag: MachineInstr::MIFlag::NonNeg))
288 Flag = MachineInstr::MIFlag::NonNeg;
289 MatchInfo = [=](MachineIRBuilder &B) { B.buildZExt(Res: Dst, Op: Src, Flags: Flag); };
290 return true;
291 }
292 MatchInfo = [=](MachineIRBuilder &B) { B.buildSExt(Res: Dst, Op: Src); };
293 return true;
294 }
295
296 return false;
297}
298
299bool CombinerHelper::matchCastOfBuildVector(const MachineInstr &CastMI,
300 const MachineInstr &BVMI,
301 BuildFnTy &MatchInfo) const {
302 const GExtOrTruncOp *Cast = cast<GExtOrTruncOp>(Val: &CastMI);
303 const GBuildVector *BV = cast<GBuildVector>(Val: &BVMI);
304
305 if (!MRI.hasOneNonDBGUse(RegNo: BV->getReg(Idx: 0)))
306 return false;
307
308 Register Dst = Cast->getReg(Idx: 0);
309 // The type of the new build vector.
310 LLT DstTy = MRI.getType(Reg: Dst);
311 // The scalar or element type of the new build vector.
312 LLT ElemTy = DstTy.getScalarType();
313 // The scalar or element type of the old build vector.
314 LLT InputElemTy = MRI.getType(Reg: BV->getReg(Idx: 0)).getElementType();
315
316 // Check legality of new build vector, the scalar casts, and profitability of
317 // the many casts.
318 if (!isLegalOrBeforeLegalizer(
319 Query: {TargetOpcode::G_BUILD_VECTOR, {DstTy, ElemTy}}) ||
320 !isLegalOrBeforeLegalizer(Query: {Cast->getOpcode(), {ElemTy, InputElemTy}}) ||
321 !isCastFree(Opcode: Cast->getOpcode(), ToTy: ElemTy, FromTy: InputElemTy))
322 return false;
323
324 MatchInfo = [=](MachineIRBuilder &B) {
325 SmallVector<Register> Casts;
326 unsigned Elements = BV->getNumSources();
327 for (unsigned I = 0; I < Elements; ++I) {
328 auto CastI =
329 B.buildInstr(Opc: Cast->getOpcode(), DstOps: {ElemTy}, SrcOps: {BV->getSourceReg(I)});
330 Casts.push_back(Elt: CastI.getReg(Idx: 0));
331 }
332
333 B.buildBuildVector(Res: Dst, Ops: Casts);
334 };
335
336 return true;
337}
338
339bool CombinerHelper::matchNarrowBinop(const MachineInstr &TruncMI,
340 const MachineInstr &BinopMI,
341 BuildFnTy &MatchInfo) const {
342 const GTrunc *Trunc = cast<GTrunc>(Val: &TruncMI);
343 const GBinOp *BinOp = cast<GBinOp>(Val: &BinopMI);
344
345 if (!MRI.hasOneNonDBGUse(RegNo: BinOp->getReg(Idx: 0)))
346 return false;
347
348 Register Dst = Trunc->getReg(Idx: 0);
349 LLT DstTy = MRI.getType(Reg: Dst);
350
351 // Is narrow binop legal?
352 if (!isLegalOrBeforeLegalizer(Query: {BinOp->getOpcode(), {DstTy}}))
353 return false;
354
355 MatchInfo = [=](MachineIRBuilder &B) {
356 auto LHS = B.buildTrunc(Res: DstTy, Op: BinOp->getLHSReg());
357 auto RHS = B.buildTrunc(Res: DstTy, Op: BinOp->getRHSReg());
358 B.buildInstr(Opc: BinOp->getOpcode(), DstOps: {Dst}, SrcOps: {LHS, RHS});
359 };
360
361 return true;
362}
363
364bool CombinerHelper::matchCastOfInteger(const MachineInstr &CastMI,
365 APInt &MatchInfo) const {
366 const GExtOrTruncOp *Cast = cast<GExtOrTruncOp>(Val: &CastMI);
367
368 APInt Input = getIConstantFromReg(VReg: Cast->getSrcReg(), MRI);
369
370 LLT DstTy = MRI.getType(Reg: Cast->getReg(Idx: 0));
371
372 if (!isConstantLegalOrBeforeLegalizer(Ty: DstTy))
373 return false;
374
375 switch (Cast->getOpcode()) {
376 case TargetOpcode::G_TRUNC: {
377 MatchInfo = Input.trunc(width: DstTy.getScalarSizeInBits());
378 return true;
379 }
380 default:
381 return false;
382 }
383}
384
385bool CombinerHelper::matchRedundantSextInReg(MachineInstr &Root,
386 MachineInstr &Other,
387 BuildFnTy &MatchInfo) const {
388 assert(Root.getOpcode() == TargetOpcode::G_SEXT_INREG &&
389 Other.getOpcode() == TargetOpcode::G_SEXT_INREG);
390
391 unsigned RootWidth = Root.getOperand(i: 2).getImm();
392 unsigned OtherWidth = Other.getOperand(i: 2).getImm();
393
394 Register Dst = Root.getOperand(i: 0).getReg();
395 Register OtherDst = Other.getOperand(i: 0).getReg();
396 Register Src = Other.getOperand(i: 1).getReg();
397
398 if (RootWidth >= OtherWidth) {
399 // The root sext_inreg is entirely redundant because the other one
400 // is narrower.
401 if (!canReplaceReg(DstReg: Dst, SrcReg: OtherDst, MRI))
402 return false;
403
404 MatchInfo = [=](MachineIRBuilder &B) {
405 Observer.changingAllUsesOfReg(MRI, Reg: Dst);
406 MRI.replaceRegWith(FromReg: Dst, ToReg: OtherDst);
407 Observer.finishedChangingAllUsesOfReg();
408 };
409 } else {
410 // RootWidth < OtherWidth, rewrite this G_SEXT_INREG with the source of the
411 // other G_SEXT_INREG.
412 MatchInfo = [=](MachineIRBuilder &B) {
413 B.buildSExtInReg(Res: Dst, Op: Src, ImmOp: RootWidth);
414 };
415 }
416
417 return true;
418}
419