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