1//===-- SPIRVPostLegalizer.cpp - amend info after legalization -*- C++ -*-===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// The pass partially applies pre-legalization logic to new instructions
10// inserted as a result of legalization:
11// - assigns SPIR-V types to registers for new instructions.
12// - inserts ASSIGN_TYPE pseudo-instructions required for type folding.
13//
14//===----------------------------------------------------------------------===//
15
16#include "SPIRV.h"
17#include "SPIRVSubtarget.h"
18#include "SPIRVUtils.h"
19#include "llvm/CodeGen/GlobalISel/GenericMachineInstrs.h"
20#include "llvm/CodeGen/MachineFrameInfo.h"
21#include "llvm/CodeGen/MachineFunction.h"
22#include "llvm/CodeGen/MachineFunctionAnalysisManager.h"
23#include "llvm/CodeGen/MachinePassManager.h"
24#include "llvm/IR/Analysis.h"
25#include "llvm/IR/IntrinsicsSPIRV.h"
26#include "llvm/Support/Debug.h"
27#include <stack>
28
29#define DEBUG_TYPE "spirv-postlegalizer"
30
31using namespace llvm;
32
33namespace {
34class SPIRVPostLegalizerLegacy : public MachineFunctionPass {
35public:
36 static char ID;
37 SPIRVPostLegalizerLegacy() : MachineFunctionPass(ID) {}
38 bool runOnMachineFunction(MachineFunction &MF) override;
39};
40} // namespace
41
42namespace llvm {
43// Defined in SPIRVPreLegalizer.cpp.
44extern void updateRegType(Register Reg, Type *Ty, SPIRVTypeInst SpirvTy,
45 SPIRVGlobalRegistry *GR, MachineIRBuilder &MIB,
46 MachineRegisterInfo &MRI);
47extern void processInstr(MachineInstr &MI, MachineIRBuilder &MIB,
48 MachineRegisterInfo &MRI, SPIRVGlobalRegistry *GR,
49 SPIRVTypeInst KnownResType);
50} // namespace llvm
51
52static SPIRVTypeInst deduceIntTypeFromResult(Register ResVReg,
53 MachineIRBuilder &MIB,
54 SPIRVGlobalRegistry *GR) {
55 const LLT &Ty = MIB.getMRI()->getType(Reg: ResVReg);
56 SPIRVTypeInst ScalarType =
57 GR->getOrCreateSPIRVIntegerType(BitWidth: Ty.getScalarSizeInBits(), MIRBuilder&: MIB);
58 if (Ty.isVector())
59 return GR->getOrCreateSPIRVVectorType(BaseType: ScalarType, NumElements: Ty.getNumElements(), MIRBuilder&: MIB,
60 EmitIR: false);
61 return ScalarType;
62}
63
64static SPIRVTypeInst deduceTypeFromSingleOperand(MachineInstr *I,
65 MachineIRBuilder &MIB,
66 SPIRVGlobalRegistry *GR,
67 unsigned OpIdx) {
68 Register OpReg = I->getOperand(i: OpIdx).getReg();
69 if (SPIRVTypeInst OpType = GR->getSPIRVTypeForVReg(VReg: OpReg)) {
70 if (SPIRVTypeInst CompType = GR->getScalarOrVectorComponentType(Type: OpType)) {
71 Register ResVReg = I->getOperand(i: 0).getReg();
72 const LLT &ResLLT = MIB.getMRI()->getType(Reg: ResVReg);
73 if (ResLLT.isVector())
74 return GR->getOrCreateSPIRVVectorType(BaseType: CompType, NumElements: ResLLT.getNumElements(),
75 MIRBuilder&: MIB, EmitIR: false);
76 return CompType;
77 }
78 }
79 return nullptr;
80}
81
82static SPIRVTypeInst deduceTypeFromOperandRange(MachineInstr *I,
83 MachineIRBuilder &MIB,
84 SPIRVGlobalRegistry *GR,
85 unsigned StartOp,
86 unsigned EndOp) {
87 SPIRVTypeInst ResType = nullptr;
88 for (unsigned i = StartOp; i < EndOp; ++i) {
89 if (SPIRVTypeInst Type = deduceTypeFromSingleOperand(I, MIB, GR, OpIdx: i)) {
90#ifdef EXPENSIVE_CHECKS
91 assert(!ResType || Type == ResType && "Conflicting type from operands.");
92 ResType = Type;
93#else
94 return Type;
95#endif
96 }
97 }
98 return ResType;
99}
100
101static SPIRVTypeInst deduceTypeFromResultRegister(MachineInstr *Use,
102 Register UseRegister,
103 SPIRVGlobalRegistry *GR,
104 MachineIRBuilder &MIB) {
105 for (const MachineOperand &MO : Use->defs()) {
106 if (!MO.isReg())
107 continue;
108 if (SPIRVTypeInst OpType = GR->getSPIRVTypeForVReg(VReg: MO.getReg())) {
109 if (SPIRVTypeInst CompType = GR->getScalarOrVectorComponentType(Type: OpType)) {
110 const LLT &ResLLT = MIB.getMRI()->getType(Reg: UseRegister);
111 if (ResLLT.isVector())
112 return GR->getOrCreateSPIRVVectorType(
113 BaseType: CompType, NumElements: ResLLT.getNumElements(), MIRBuilder&: MIB, EmitIR: false);
114 return CompType;
115 }
116 }
117 }
118 return nullptr;
119}
120
121static SPIRVTypeInst
122deducePointerTypeFromResultRegister(MachineInstr *Use, Register UseRegister,
123 SPIRVGlobalRegistry *GR,
124 MachineIRBuilder &MIB) {
125 assert(Use->getOpcode() == TargetOpcode::G_LOAD ||
126 Use->getOpcode() == TargetOpcode::G_STORE);
127
128 Register ValueReg = Use->getOperand(i: 0).getReg();
129 SPIRVTypeInst ValueType = GR->getSPIRVTypeForVReg(VReg: ValueReg);
130 if (!ValueType)
131 return nullptr;
132
133 return GR->getOrCreateSPIRVPointerType(BaseType: ValueType, MIRBuilder&: MIB,
134 SC: SPIRV::StorageClass::Function);
135}
136
137static SPIRVTypeInst deduceTypeFromPointerOperand(MachineInstr *Use,
138 Register UseRegister,
139 SPIRVGlobalRegistry *GR,
140 MachineIRBuilder &MIB) {
141 assert(Use->getOpcode() == TargetOpcode::G_LOAD ||
142 Use->getOpcode() == TargetOpcode::G_STORE);
143
144 Register PtrReg = Use->getOperand(i: 1).getReg();
145 SPIRVTypeInst PtrType = GR->getSPIRVTypeForVReg(VReg: PtrReg);
146 if (!PtrType)
147 return nullptr;
148
149 return GR->getPointeeType(PtrType);
150}
151
152static SPIRVTypeInst deduceTypeFromUses(Register Reg, MachineFunction &MF,
153 SPIRVGlobalRegistry *GR,
154 MachineIRBuilder &MIB) {
155 MachineRegisterInfo &MRI = MF.getRegInfo();
156 for (MachineInstr &Use : MRI.use_nodbg_instructions(Reg)) {
157 SPIRVTypeInst ResType = nullptr;
158 LLVM_DEBUG(dbgs() << "Looking at use " << Use);
159 switch (Use.getOpcode()) {
160 case TargetOpcode::G_BUILD_VECTOR:
161 case TargetOpcode::G_SHUFFLE_VECTOR:
162 case TargetOpcode::G_EXTRACT_VECTOR_ELT:
163 case TargetOpcode::G_UNMERGE_VALUES:
164 case TargetOpcode::G_ADD:
165 case TargetOpcode::G_SUB:
166 case TargetOpcode::G_MUL:
167 case TargetOpcode::G_SDIV:
168 case TargetOpcode::G_UDIV:
169 case TargetOpcode::G_SREM:
170 case TargetOpcode::G_UREM:
171 case TargetOpcode::G_FADD:
172 case TargetOpcode::G_FSUB:
173 case TargetOpcode::G_FMUL:
174 case TargetOpcode::G_FDIV:
175 case TargetOpcode::G_FREM:
176 case TargetOpcode::G_FMA:
177 case TargetOpcode::G_FATAN2:
178 case TargetOpcode::G_FPOW:
179 case TargetOpcode::COPY:
180 case TargetOpcode::G_STRICT_FMA:
181 ResType = deduceTypeFromResultRegister(Use: &Use, UseRegister: Reg, GR, MIB);
182 break;
183 case TargetOpcode::G_LOAD:
184 case TargetOpcode::G_STORE:
185 if (Reg == Use.getOperand(i: 1).getReg())
186 ResType = deducePointerTypeFromResultRegister(Use: &Use, UseRegister: Reg, GR, MIB);
187 else
188 ResType = deduceTypeFromPointerOperand(Use: &Use, UseRegister: Reg, GR, MIB);
189 break;
190 case TargetOpcode::G_INTRINSIC_W_SIDE_EFFECTS:
191 case TargetOpcode::G_INTRINSIC: {
192 auto IntrinsicID = cast<GIntrinsic>(Val&: Use).getIntrinsicID();
193 if (IntrinsicID == Intrinsic::spv_insertelt) {
194 if (Reg == Use.getOperand(i: 2).getReg())
195 ResType = deduceTypeFromResultRegister(Use: &Use, UseRegister: Reg, GR, MIB);
196 } else if (IntrinsicID == Intrinsic::spv_extractelt) {
197 if (Reg == Use.getOperand(i: 2).getReg())
198 ResType = deduceTypeFromResultRegister(Use: &Use, UseRegister: Reg, GR, MIB);
199 }
200 break;
201 }
202 }
203 if (ResType) {
204 LLVM_DEBUG(dbgs() << "Deduced type from use " << *ResType);
205 return ResType;
206 }
207 }
208 return nullptr;
209}
210
211static SPIRVTypeInst deduceGEPType(MachineInstr *I, SPIRVGlobalRegistry *GR,
212 MachineIRBuilder &MIB) {
213 LLVM_DEBUG(dbgs() << "Deducing GEP type for: " << *I);
214 Register PtrReg = I->getOperand(i: 3).getReg();
215 SPIRVTypeInst PtrType = GR->getSPIRVTypeForVReg(VReg: PtrReg);
216 if (!PtrType) {
217 LLVM_DEBUG(dbgs() << " Could not get type for pointer operand.\n");
218 return nullptr;
219 }
220
221 SPIRVTypeInst PointeeType = GR->getPointeeType(PtrType);
222 if (!PointeeType) {
223 LLVM_DEBUG(dbgs() << " Could not get pointee type from pointer type.\n");
224 return nullptr;
225 }
226
227 MachineRegisterInfo *MRI = MIB.getMRI();
228
229 // The first index (operand 4) steps over the pointer, so the type doesn't
230 // change.
231 for (unsigned i = 5; i < I->getNumOperands(); ++i) {
232 LLVM_DEBUG(dbgs() << " Traversing index " << i
233 << ", current type: " << *PointeeType);
234 switch (PointeeType->getOpcode()) {
235 case SPIRV::OpTypeArray:
236 case SPIRV::OpTypeRuntimeArray:
237 case SPIRV::OpTypeVector:
238 case SPIRV::OpTypeVectorIdEXT: {
239 Register ElemTypeReg = PointeeType->getOperand(i: 1).getReg();
240 PointeeType = GR->getSPIRVTypeForVReg(VReg: ElemTypeReg);
241 break;
242 }
243 case SPIRV::OpTypeStruct: {
244 MachineOperand &IdxOp = I->getOperand(i);
245 if (!IdxOp.isReg()) {
246 LLVM_DEBUG(dbgs() << " Index is not a register.\n");
247 return nullptr;
248 }
249 MachineInstr *Def = MRI->getVRegDef(Reg: IdxOp.getReg());
250 if (!Def) {
251 LLVM_DEBUG(
252 dbgs() << " Could not find definition for index register.\n");
253 return nullptr;
254 }
255
256 uint64_t IndexVal = foldImm(MO: IdxOp, MRI);
257 if (IndexVal >= PointeeType->getNumOperands() - 1) {
258 LLVM_DEBUG(dbgs() << " Struct index out of bounds.\n");
259 return nullptr;
260 }
261
262 Register MemberTypeReg = PointeeType->getOperand(i: IndexVal + 1).getReg();
263 PointeeType = GR->getSPIRVTypeForVReg(VReg: MemberTypeReg);
264 break;
265 }
266 default:
267 LLVM_DEBUG(dbgs() << " Unknown type opcode for GEP traversal.\n");
268 return nullptr;
269 }
270
271 if (!PointeeType) {
272 LLVM_DEBUG(dbgs() << " Could not resolve next pointee type.\n");
273 return nullptr;
274 }
275 }
276 LLVM_DEBUG(dbgs() << " Final pointee type: " << *PointeeType);
277
278 SPIRV::StorageClass::StorageClass SC = GR->getPointerStorageClass(Type: PtrType);
279 SPIRVTypeInst Res = GR->getOrCreateSPIRVPointerType(BaseType: PointeeType, MIRBuilder&: MIB, SC);
280 LLVM_DEBUG(dbgs() << " Deduced GEP type: " << *Res);
281 return Res;
282}
283
284static SPIRVTypeInst deduceResultTypeFromOperands(MachineInstr *I,
285 SPIRVGlobalRegistry *GR,
286 MachineIRBuilder &MIB) {
287 Register ResVReg = I->getOperand(i: 0).getReg();
288 switch (I->getOpcode()) {
289 case TargetOpcode::G_CONSTANT:
290 case TargetOpcode::G_ANYEXT:
291 case TargetOpcode::G_SEXT:
292 case TargetOpcode::G_ZEXT:
293 case TargetOpcode::G_TRUNC:
294 return deduceIntTypeFromResult(ResVReg, MIB, GR);
295 case TargetOpcode::G_BUILD_VECTOR:
296 return deduceTypeFromOperandRange(I, MIB, GR, StartOp: 1, EndOp: I->getNumOperands());
297 case TargetOpcode::G_SHUFFLE_VECTOR:
298 return deduceTypeFromOperandRange(I, MIB, GR, StartOp: 1, EndOp: 3);
299 case TargetOpcode::G_INTRINSIC_W_SIDE_EFFECTS:
300 case TargetOpcode::G_INTRINSIC: {
301 auto IntrinsicID = cast<GIntrinsic>(Val: I)->getIntrinsicID();
302 if (IntrinsicID == Intrinsic::spv_gep)
303 return deduceGEPType(I, GR, MIB);
304 break;
305 }
306 case TargetOpcode::G_LOAD: {
307 SPIRVTypeInst PtrType = deduceTypeFromSingleOperand(I, MIB, GR, OpIdx: 1);
308 return PtrType ? GR->getPointeeType(PtrType) : nullptr;
309 }
310 case TargetOpcode::G_PHI: {
311 for (unsigned Idx = 1; Idx < I->getNumOperands(); Idx += 2) {
312 Register OpReg = I->getOperand(i: Idx).getReg();
313 if (SPIRVTypeInst OpType = GR->getSPIRVTypeForVReg(VReg: OpReg))
314 return OpType;
315 }
316 return nullptr;
317 }
318 default:
319 if (I->getNumDefs() == 1 && I->getNumOperands() > 1 &&
320 I->getOperand(i: 1).isReg())
321 return deduceTypeFromSingleOperand(I, MIB, GR, OpIdx: 1);
322 }
323 return nullptr;
324}
325
326static bool deduceAndAssignTypeForGUnmerge(MachineInstr *I, MachineFunction &MF,
327 SPIRVGlobalRegistry *GR,
328 MachineIRBuilder &MIB) {
329 MachineRegisterInfo &MRI = MF.getRegInfo();
330 Register SrcReg = I->getOperand(i: I->getNumOperands() - 1).getReg();
331 SPIRVTypeInst ScalarType = nullptr;
332 if (SPIRVTypeInst DefType = GR->getSPIRVTypeForVReg(VReg: SrcReg)) {
333 assert(isVectorType(DefType));
334 ScalarType = GR->getScalarOrVectorComponentType(Type: DefType);
335 }
336
337 if (!ScalarType) {
338 // If we could not deduce the type from the source, try to deduce it from
339 // the uses of the results.
340 for (unsigned i = 0; i < I->getNumDefs(); ++i) {
341 Register DefReg = I->getOperand(i).getReg();
342 ScalarType = deduceTypeFromUses(Reg: DefReg, MF, GR, MIB);
343 if (ScalarType) {
344 ScalarType = GR->getScalarOrVectorComponentType(Type: ScalarType);
345 break;
346 }
347 }
348 }
349
350 if (!ScalarType)
351 return false;
352
353 for (unsigned i = 0; i < I->getNumOperands(); ++i) {
354 Register DefReg = I->getOperand(i).getReg();
355 if (GR->getSPIRVTypeForVReg(VReg: DefReg))
356 continue;
357
358 LLT DefLLT = MRI.getType(Reg: DefReg);
359 SPIRVTypeInst ResType =
360 DefLLT.isVector()
361 ? GR->getOrCreateSPIRVVectorType(
362 BaseType: ScalarType, NumElements: DefLLT.getNumElements(), I&: *I,
363 TII: *MF.getSubtarget<SPIRVSubtarget>().getInstrInfo())
364 : ScalarType;
365 setRegClassType(Reg: DefReg, SpvType: ResType, GR, MRI: &MRI, MF);
366 }
367 return true;
368}
369
370static bool deduceAndAssignSpirvType(MachineInstr *I, MachineFunction &MF,
371 SPIRVGlobalRegistry *GR,
372 MachineIRBuilder &MIB) {
373 LLVM_DEBUG(dbgs() << "\nProcessing instruction: " << *I);
374 MachineRegisterInfo &MRI = MF.getRegInfo();
375 Register ResVReg = I->getOperand(i: 0).getReg();
376
377 // G_UNMERGE_VALUES is handled separately because it has multiple definitions,
378 // unlike the other instructions which have a single result register. The main
379 // deduction logic is designed for the single-definition case.
380 if (I->getOpcode() == TargetOpcode::G_UNMERGE_VALUES)
381 return deduceAndAssignTypeForGUnmerge(I, MF, GR, MIB);
382
383 LLVM_DEBUG(dbgs() << "Inferring type from operands\n");
384 SPIRVTypeInst ResType = deduceResultTypeFromOperands(I, GR, MIB);
385 if (!ResType) {
386 LLVM_DEBUG(dbgs() << "Inferring type from uses\n");
387 ResType = deduceTypeFromUses(Reg: ResVReg, MF, GR, MIB);
388 }
389
390 if (!ResType)
391 return false;
392
393 LLVM_DEBUG(dbgs() << "Assigned type to " << *I << ": " << *ResType);
394 setRegClassType(Reg: ResVReg, SpvType: ResType, GR, MRI: &MRI, MF);
395 return true;
396}
397
398static bool requiresSpirvType(MachineInstr &I, SPIRVGlobalRegistry *GR,
399 MachineRegisterInfo &MRI) {
400 LLVM_DEBUG(dbgs() << "Checking if instruction requires a SPIR-V type: "
401 << I;);
402 if (I.getNumDefs() == 0) {
403 LLVM_DEBUG(dbgs() << "Instruction does not have a definition.\n");
404 return false;
405 }
406
407 if (!I.isPreISelOpcode()) {
408 LLVM_DEBUG(dbgs() << "Instruction is not a generic instruction.\n");
409 return false;
410 }
411
412 Register ResultRegister = I.defs().begin()->getReg();
413 if (GR->getSPIRVTypeForVReg(VReg: ResultRegister)) {
414 LLVM_DEBUG(dbgs() << "Instruction already has a SPIR-V type.\n");
415 if (!MRI.getRegClassOrNull(Reg: ResultRegister)) {
416 LLVM_DEBUG(dbgs() << "Updating the register class.\n");
417 setRegClassType(Reg: ResultRegister, SpvType: GR->getSPIRVTypeForVReg(VReg: ResultRegister),
418 GR, MRI: &MRI, MF: *GR->CurMF, Force: true);
419 }
420 return false;
421 }
422
423 return true;
424}
425
426static void registerSpirvTypeForNewInstructions(MachineFunction &MF,
427 SPIRVGlobalRegistry *GR) {
428 MachineRegisterInfo &MRI = MF.getRegInfo();
429 SmallVector<MachineInstr *, 8> Worklist;
430 for (MachineBasicBlock &MBB : MF) {
431 for (MachineInstr &I : MBB) {
432 if (requiresSpirvType(I, GR, MRI)) {
433 Worklist.push_back(Elt: &I);
434 }
435 }
436 }
437
438 if (Worklist.empty()) {
439 LLVM_DEBUG(dbgs() << "Initial worklist is empty.\n");
440 return;
441 }
442
443 LLVM_DEBUG(dbgs() << "Initial worklist:\n";
444 for (auto *I : Worklist) { I->dump(); });
445
446 bool Changed;
447 do {
448 Changed = false;
449 SmallVector<MachineInstr *, 8> NextWorklist;
450
451 for (MachineInstr *I : Worklist) {
452 MachineIRBuilder MIB(*I);
453 if (deduceAndAssignSpirvType(I, MF, GR, MIB)) {
454 Changed = true;
455 } else {
456 NextWorklist.push_back(Elt: I);
457 }
458 }
459 Worklist = std::move(NextWorklist);
460 LLVM_DEBUG(dbgs() << "Worklist size: " << Worklist.size() << "\n");
461 } while (Changed);
462
463 if (Worklist.empty())
464 return;
465
466 for (auto *I : Worklist) {
467 MachineIRBuilder MIB(*I);
468 LLVM_DEBUG(dbgs() << "Assigning default type to results in " << *I);
469 for (unsigned Idx = 0; Idx < I->getNumDefs(); ++Idx) {
470 Register ResVReg = I->getOperand(i: Idx).getReg();
471 if (GR->getSPIRVTypeForVReg(VReg: ResVReg))
472 continue;
473 const LLT &ResLLT = MRI.getType(Reg: ResVReg);
474 SPIRVTypeInst ResType = nullptr;
475 if (ResLLT.isVector()) {
476 SPIRVTypeInst CompType = GR->getOrCreateSPIRVIntegerType(
477 BitWidth: ResLLT.getElementType().getSizeInBits(), MIRBuilder&: MIB);
478 ResType = GR->getOrCreateSPIRVVectorType(
479 BaseType: CompType, NumElements: ResLLT.getNumElements(), MIRBuilder&: MIB, EmitIR: false);
480 } else {
481 ResType = GR->getOrCreateSPIRVIntegerType(BitWidth: ResLLT.getSizeInBits(), MIRBuilder&: MIB);
482 }
483 setRegClassType(Reg: ResVReg, SpvType: ResType, GR, MRI: &MRI, MF, Force: true);
484 }
485 }
486}
487
488static bool hasAssignType(Register Reg, MachineRegisterInfo &MRI) {
489 for (MachineInstr &UseInstr : MRI.use_nodbg_instructions(Reg)) {
490 if (UseInstr.getOpcode() == SPIRV::ASSIGN_TYPE) {
491 return true;
492 }
493 }
494 return false;
495}
496
497static void generateAssignType(MachineInstr &MI, Register ResultRegister,
498 SPIRVTypeInst ResultType,
499 SPIRVGlobalRegistry *GR,
500 MachineRegisterInfo &MRI) {
501 LLVM_DEBUG(dbgs() << " Adding ASSIGN_TYPE for ResultRegister: "
502 << printReg(ResultRegister, MRI.getTargetRegisterInfo())
503 << " with type: " << *ResultType);
504 MachineIRBuilder MIB(MI);
505 updateRegType(Reg: ResultRegister, Ty: nullptr, SpirvTy: ResultType, GR, MIB, MRI);
506
507 // Tablegen definition assumes SPIRV::ASSIGN_TYPE pseudo-instruction is
508 // present after each auto-folded instruction to take a type reference
509 // from.
510 Register NewReg =
511 MRI.createGenericVirtualRegister(Ty: MRI.getType(Reg: ResultRegister));
512 const auto *RegClass = GR->getRegClass(SpvType: ResultType);
513 MRI.setRegClass(Reg: NewReg, RC: RegClass);
514 MRI.setRegClass(Reg: ResultRegister, RC: RegClass);
515
516 GR->assignSPIRVTypeToVReg(Type: ResultType, VReg: ResultRegister, MF: MIB.getMF());
517 // This is to make it convenient for Legalizer to get the SPIRVType
518 // when processing the actual MI (i.e. not pseudo one).
519 GR->assignSPIRVTypeToVReg(Type: ResultType, VReg: NewReg, MF: MIB.getMF());
520 // Copy MIFlags from Def to ASSIGN_TYPE instruction. It's required to
521 // keep the flags after instruction selection.
522 const uint32_t Flags = MI.getFlags();
523 MIB.buildInstr(Opcode: SPIRV::ASSIGN_TYPE)
524 .addDef(RegNo: ResultRegister)
525 .addUse(RegNo: NewReg)
526 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: ResultType))
527 .setMIFlags(Flags);
528 for (unsigned I = 0, E = MI.getNumDefs(); I != E; ++I) {
529 MachineOperand &MO = MI.getOperand(i: I);
530 if (MO.getReg() == ResultRegister) {
531 MO.setReg(NewReg);
532 break;
533 }
534 }
535}
536
537static void ensureAssignTypeForTypeFolding(MachineFunction &MF,
538 SPIRVGlobalRegistry *GR) {
539 LLVM_DEBUG(dbgs() << "Entering ensureAssignTypeForTypeFolding for function "
540 << MF.getName() << "\n");
541 MachineRegisterInfo &MRI = MF.getRegInfo();
542 for (MachineBasicBlock &MBB : MF) {
543 for (MachineInstr &MI : MBB) {
544 if (!isTypeFoldingSupported(Opcode: MI.getOpcode()))
545 continue;
546
547 LLVM_DEBUG(dbgs() << "Processing instruction: " << MI);
548
549 Register ResultRegister = MI.defs().begin()->getReg();
550 if (hasAssignType(Reg: ResultRegister, MRI)) {
551 LLVM_DEBUG(dbgs() << " Instruction already has ASSIGN_TYPE\n");
552 continue;
553 }
554
555 SPIRVTypeInst ResultType = GR->getSPIRVTypeForVReg(VReg: ResultRegister);
556 generateAssignType(MI, ResultRegister, ResultType, GR, MRI);
557 }
558 }
559}
560
561static bool runPostLegalizer(MachineFunction &MF) {
562 // Initialize the type registry.
563 const SPIRVSubtarget &ST = MF.getSubtarget<SPIRVSubtarget>();
564 SPIRVGlobalRegistry *GR = ST.getSPIRVGlobalRegistry();
565 GR->setCurrentFunc(MF);
566 registerSpirvTypeForNewInstructions(MF, GR);
567 ensureAssignTypeForTypeFolding(MF, GR);
568 return true;
569}
570
571INITIALIZE_PASS(SPIRVPostLegalizerLegacy, DEBUG_TYPE, "SPIRV post legalizer",
572 false, false)
573
574char SPIRVPostLegalizerLegacy::ID = 0;
575
576FunctionPass *llvm::createSPIRVPostLegalizerLegacyPass() {
577 return new SPIRVPostLegalizerLegacy();
578}
579
580bool SPIRVPostLegalizerLegacy::runOnMachineFunction(MachineFunction &MF) {
581 return runPostLegalizer(MF);
582}
583
584PreservedAnalyses
585SPIRVPostLegalizerPass::run(MachineFunction &MF,
586 MachineFunctionAnalysisManager &MFAM) {
587 return runPostLegalizer(MF) ? getMachineFunctionPassPreservedAnalyses()
588 : PreservedAnalyses::all();
589}
590