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_FABS:
197 case TargetOpcode::G_FSQRT:
198 case TargetOpcode::COPY:
199 case TargetOpcode::G_STRICT_FMA:
200 case TargetOpcode::G_INTRINSIC_TRUNC:
201 case TargetOpcode::G_INTRINSIC_ROUNDEVEN:
202 ResType = deduceTypeFromResultRegister(Use: &Use, UseRegister: Reg, GR, MIB);
203 break;
204 case TargetOpcode::G_SELECT:
205 if (Reg == Use.getOperand(i: 2).getReg() ||
206 Reg == Use.getOperand(i: 3).getReg())
207 ResType = deduceTypeFromResultRegister(Use: &Use, UseRegister: Reg, GR, MIB);
208 break;
209 case TargetOpcode::G_LOAD:
210 case TargetOpcode::G_STORE:
211 if (Reg == Use.getOperand(i: 1).getReg())
212 ResType = deducePointerTypeFromResultRegister(Use: &Use, UseRegister: Reg, GR, MIB);
213 else
214 ResType = deduceTypeFromPointerOperand(Use: &Use, UseRegister: Reg, GR, MIB);
215 break;
216 case TargetOpcode::G_INTRINSIC_W_SIDE_EFFECTS:
217 case TargetOpcode::G_INTRINSIC: {
218 auto IntrinsicID = cast<GIntrinsic>(Val&: Use).getIntrinsicID();
219 if (IntrinsicID == Intrinsic::spv_insertelt) {
220 if (Reg == Use.getOperand(i: 2).getReg())
221 ResType = deduceTypeFromResultRegister(Use: &Use, UseRegister: Reg, GR, MIB);
222 } else if (IntrinsicID == Intrinsic::spv_extractelt) {
223 if (Reg == Use.getOperand(i: 2).getReg())
224 ResType = deduceTypeFromResultRegister(Use: &Use, UseRegister: Reg, GR, MIB);
225 }
226 break;
227 }
228 }
229 if (ResType) {
230 LLVM_DEBUG(dbgs() << "Deduced type from use " << *ResType);
231 return ResType;
232 }
233 }
234 return nullptr;
235}
236
237static SPIRVTypeInst deduceGEPType(MachineInstr *I, SPIRVGlobalRegistry *GR,
238 MachineIRBuilder &MIB) {
239 LLVM_DEBUG(dbgs() << "Deducing GEP type for: " << *I);
240 Register PtrReg = I->getOperand(i: 3).getReg();
241 SPIRVTypeInst PtrType = GR->getSPIRVTypeForVReg(VReg: PtrReg);
242 if (!PtrType) {
243 LLVM_DEBUG(dbgs() << " Could not get type for pointer operand.\n");
244 return nullptr;
245 }
246
247 SPIRVTypeInst PointeeType = GR->getPointeeType(PtrType);
248 if (!PointeeType) {
249 LLVM_DEBUG(dbgs() << " Could not get pointee type from pointer type.\n");
250 return nullptr;
251 }
252
253 MachineRegisterInfo *MRI = MIB.getMRI();
254
255 // The first index (operand 4) steps over the pointer, so the type doesn't
256 // change.
257 for (unsigned i = 5; i < I->getNumOperands(); ++i) {
258 LLVM_DEBUG(dbgs() << " Traversing index " << i
259 << ", current type: " << *PointeeType);
260 switch (PointeeType->getOpcode()) {
261 case SPIRV::OpTypeArray:
262 case SPIRV::OpTypeRuntimeArray:
263 case SPIRV::OpTypeVector:
264 case SPIRV::OpTypeVectorIdEXT: {
265 Register ElemTypeReg = PointeeType->getOperand(i: 1).getReg();
266 PointeeType = GR->getSPIRVTypeForVReg(VReg: ElemTypeReg);
267 break;
268 }
269 case SPIRV::OpTypeStruct: {
270 MachineOperand &IdxOp = I->getOperand(i);
271 if (!IdxOp.isReg()) {
272 LLVM_DEBUG(dbgs() << " Index is not a register.\n");
273 return nullptr;
274 }
275 MachineInstr *Def = MRI->getVRegDef(Reg: IdxOp.getReg());
276 if (!Def) {
277 LLVM_DEBUG(
278 dbgs() << " Could not find definition for index register.\n");
279 return nullptr;
280 }
281
282 uint64_t IndexVal = foldImm(MO: IdxOp, MRI);
283 if (IndexVal >= PointeeType->getNumOperands() - 1) {
284 LLVM_DEBUG(dbgs() << " Struct index out of bounds.\n");
285 return nullptr;
286 }
287
288 Register MemberTypeReg = PointeeType->getOperand(i: IndexVal + 1).getReg();
289 PointeeType = GR->getSPIRVTypeForVReg(VReg: MemberTypeReg);
290 break;
291 }
292 default:
293 LLVM_DEBUG(dbgs() << " Unknown type opcode for GEP traversal.\n");
294 return nullptr;
295 }
296
297 if (!PointeeType) {
298 LLVM_DEBUG(dbgs() << " Could not resolve next pointee type.\n");
299 return nullptr;
300 }
301 }
302 LLVM_DEBUG(dbgs() << " Final pointee type: " << *PointeeType);
303
304 SPIRV::StorageClass::StorageClass SC = GR->getPointerStorageClass(Type: PtrType);
305 SPIRVTypeInst Res = GR->getOrCreateSPIRVPointerType(BaseType: PointeeType, MIRBuilder&: MIB, SC);
306 LLVM_DEBUG(dbgs() << " Deduced GEP type: " << *Res);
307 return Res;
308}
309
310static SPIRVTypeInst deduceResultTypeFromOperands(MachineInstr *I,
311 SPIRVGlobalRegistry *GR,
312 MachineIRBuilder &MIB) {
313 Register ResVReg = I->getOperand(i: 0).getReg();
314 switch (I->getOpcode()) {
315 case TargetOpcode::G_CONSTANT:
316 case TargetOpcode::G_ANYEXT:
317 case TargetOpcode::G_SEXT:
318 case TargetOpcode::G_ZEXT:
319 case TargetOpcode::G_TRUNC:
320 return deduceIntTypeFromResult(ResVReg, MIB, GR);
321 case TargetOpcode::G_BUILD_VECTOR:
322 return deduceTypeFromOperandRange(I, MIB, GR, StartOp: 1, EndOp: I->getNumOperands());
323 case TargetOpcode::G_SHUFFLE_VECTOR:
324 return deduceTypeFromOperandRange(I, MIB, GR, StartOp: 1, EndOp: 3);
325 case TargetOpcode::G_SELECT:
326 return deduceTypeFromOperandRange(I, MIB, GR, StartOp: 2, EndOp: 4);
327 case TargetOpcode::G_INTRINSIC_W_SIDE_EFFECTS:
328 case TargetOpcode::G_INTRINSIC: {
329 auto IntrinsicID = cast<GIntrinsic>(Val: I)->getIntrinsicID();
330 if (IntrinsicID == Intrinsic::spv_gep)
331 return deduceGEPType(I, GR, MIB);
332 break;
333 }
334 case TargetOpcode::G_LOAD: {
335 SPIRVTypeInst PtrType = deduceTypeFromSingleOperand(I, MIB, GR, OpIdx: 1);
336 return PtrType ? GR->getPointeeType(PtrType) : nullptr;
337 }
338 case TargetOpcode::G_PHI: {
339 for (unsigned Idx = 1; Idx < I->getNumOperands(); Idx += 2) {
340 Register OpReg = I->getOperand(i: Idx).getReg();
341 if (SPIRVTypeInst OpType = GR->getSPIRVTypeForVReg(VReg: OpReg))
342 return OpType;
343 }
344 return nullptr;
345 }
346 default:
347 if (I->getNumDefs() == 1 && I->getNumOperands() > 1 &&
348 I->getOperand(i: 1).isReg())
349 return deduceTypeFromSingleOperand(I, MIB, GR, OpIdx: 1);
350 }
351 return nullptr;
352}
353
354static bool deduceAndAssignTypeForGUnmerge(MachineInstr *I, MachineFunction &MF,
355 SPIRVGlobalRegistry *GR,
356 MachineIRBuilder &MIB) {
357 MachineRegisterInfo &MRI = MF.getRegInfo();
358 Register SrcReg = I->getOperand(i: I->getNumOperands() - 1).getReg();
359 SPIRVTypeInst ScalarType = nullptr;
360 if (SPIRVTypeInst DefType = GR->getSPIRVTypeForVReg(VReg: SrcReg)) {
361 assert(isVectorType(DefType));
362 ScalarType = GR->getScalarOrVectorComponentType(Type: DefType);
363 }
364
365 if (!ScalarType) {
366 // If we could not deduce the type from the source, try to deduce it from
367 // the uses of the results.
368 for (unsigned i = 0; i < I->getNumDefs(); ++i) {
369 Register DefReg = I->getOperand(i).getReg();
370 ScalarType = deduceTypeFromUses(Reg: DefReg, MF, GR, MIB);
371 if (ScalarType) {
372 ScalarType = GR->getScalarOrVectorComponentType(Type: ScalarType);
373 break;
374 }
375 }
376 }
377
378 if (!ScalarType)
379 return false;
380
381 for (unsigned i = 0; i < I->getNumOperands(); ++i) {
382 Register DefReg = I->getOperand(i).getReg();
383 if (GR->getSPIRVTypeForVReg(VReg: DefReg))
384 continue;
385
386 LLT DefLLT = MRI.getType(Reg: DefReg);
387 SPIRVTypeInst ResType =
388 DefLLT.isVector()
389 ? GR->getOrCreateSPIRVVectorType(
390 BaseType: ScalarType, NumElements: DefLLT.getNumElements(), I&: *I,
391 TII: *MF.getSubtarget<SPIRVSubtarget>().getInstrInfo())
392 : ScalarType;
393 setRegClassType(Reg: DefReg, SpvType: ResType, GR, MRI: &MRI, MF);
394 }
395 return true;
396}
397
398static bool deduceAndAssignSpirvType(MachineInstr *I, MachineFunction &MF,
399 SPIRVGlobalRegistry *GR,
400 MachineIRBuilder &MIB) {
401 LLVM_DEBUG(dbgs() << "\nProcessing instruction: " << *I);
402 MachineRegisterInfo &MRI = MF.getRegInfo();
403 Register ResVReg = I->getOperand(i: 0).getReg();
404
405 // G_UNMERGE_VALUES is handled separately because it has multiple definitions,
406 // unlike the other instructions which have a single result register. The main
407 // deduction logic is designed for the single-definition case.
408 if (I->getOpcode() == TargetOpcode::G_UNMERGE_VALUES)
409 return deduceAndAssignTypeForGUnmerge(I, MF, GR, MIB);
410
411 LLVM_DEBUG(dbgs() << "Inferring type from operands\n");
412 SPIRVTypeInst ResType = deduceResultTypeFromOperands(I, GR, MIB);
413 if (!ResType) {
414 LLVM_DEBUG(dbgs() << "Inferring type from uses\n");
415 ResType = deduceTypeFromUses(Reg: ResVReg, MF, GR, MIB);
416 }
417
418 if (!ResType)
419 return false;
420
421 LLVM_DEBUG(dbgs() << "Assigned type to " << *I << ": " << *ResType);
422 setRegClassType(Reg: ResVReg, SpvType: ResType, GR, MRI: &MRI, MF);
423 return true;
424}
425
426static bool requiresSpirvType(MachineInstr &I, SPIRVGlobalRegistry *GR,
427 MachineRegisterInfo &MRI) {
428 LLVM_DEBUG(dbgs() << "Checking if instruction requires a SPIR-V type: "
429 << I;);
430 if (I.getNumDefs() == 0) {
431 LLVM_DEBUG(dbgs() << "Instruction does not have a definition.\n");
432 return false;
433 }
434
435 if (!I.isPreISelOpcode()) {
436 LLVM_DEBUG(dbgs() << "Instruction is not a generic instruction.\n");
437 return false;
438 }
439
440 Register ResultRegister = I.defs().begin()->getReg();
441 if (GR->getSPIRVTypeForVReg(VReg: ResultRegister)) {
442 LLVM_DEBUG(dbgs() << "Instruction already has a SPIR-V type.\n");
443 if (!MRI.getRegClassOrNull(Reg: ResultRegister)) {
444 LLVM_DEBUG(dbgs() << "Updating the register class.\n");
445 setRegClassType(Reg: ResultRegister, SpvType: GR->getSPIRVTypeForVReg(VReg: ResultRegister),
446 GR, MRI: &MRI, MF: *GR->CurMF, Force: true);
447 }
448 return false;
449 }
450
451 return true;
452}
453
454static void registerSpirvTypeForNewInstructions(MachineFunction &MF,
455 SPIRVGlobalRegistry *GR) {
456 MachineRegisterInfo &MRI = MF.getRegInfo();
457 SmallVector<MachineInstr *, 8> Worklist;
458 for (MachineBasicBlock &MBB : MF) {
459 for (MachineInstr &I : MBB) {
460 if (requiresSpirvType(I, GR, MRI)) {
461 Worklist.push_back(Elt: &I);
462 }
463 }
464 }
465
466 if (Worklist.empty()) {
467 LLVM_DEBUG(dbgs() << "Initial worklist is empty.\n");
468 return;
469 }
470
471 LLVM_DEBUG(dbgs() << "Initial worklist:\n";
472 for (auto *I : Worklist) { I->dump(); });
473
474 bool Changed;
475 do {
476 Changed = false;
477 SmallVector<MachineInstr *, 8> NextWorklist;
478
479 for (MachineInstr *I : Worklist) {
480 MachineIRBuilder MIB(*I);
481 if (deduceAndAssignSpirvType(I, MF, GR, MIB)) {
482 Changed = true;
483 } else {
484 NextWorklist.push_back(Elt: I);
485 }
486 }
487 Worklist = std::move(NextWorklist);
488 LLVM_DEBUG(dbgs() << "Worklist size: " << Worklist.size() << "\n");
489 } while (Changed);
490
491 if (Worklist.empty())
492 return;
493
494 for (auto *I : Worklist) {
495 MachineIRBuilder MIB(*I);
496 LLVM_DEBUG(dbgs() << "Assigning default type to results in " << *I);
497 for (unsigned Idx = 0; Idx < I->getNumDefs(); ++Idx) {
498 Register ResVReg = I->getOperand(i: Idx).getReg();
499 if (GR->getSPIRVTypeForVReg(VReg: ResVReg))
500 continue;
501 const LLT &ResLLT = MRI.getType(Reg: ResVReg);
502 SPIRVTypeInst ResType = nullptr;
503 if (ResLLT.isVector()) {
504 SPIRVTypeInst CompType = GR->getOrCreateSPIRVIntegerType(
505 BitWidth: ResLLT.getElementType().getSizeInBits(), MIRBuilder&: MIB);
506 ResType = GR->getOrCreateSPIRVVectorType(
507 BaseType: CompType, NumElements: ResLLT.getNumElements(), MIRBuilder&: MIB, EmitIR: false);
508 } else {
509 ResType = GR->getOrCreateSPIRVIntegerType(BitWidth: ResLLT.getSizeInBits(), MIRBuilder&: MIB);
510 }
511 setRegClassType(Reg: ResVReg, SpvType: ResType, GR, MRI: &MRI, MF, Force: true);
512 }
513 }
514}
515
516static bool hasAssignType(Register Reg, MachineRegisterInfo &MRI) {
517 for (MachineInstr &UseInstr : MRI.use_nodbg_instructions(Reg)) {
518 if (UseInstr.getOpcode() == SPIRV::ASSIGN_TYPE) {
519 return true;
520 }
521 }
522 return false;
523}
524
525static void generateAssignType(MachineInstr &MI, Register ResultRegister,
526 SPIRVTypeInst ResultType,
527 SPIRVGlobalRegistry *GR,
528 MachineRegisterInfo &MRI) {
529 LLVM_DEBUG(dbgs() << " Adding ASSIGN_TYPE for ResultRegister: "
530 << printReg(ResultRegister, MRI.getTargetRegisterInfo())
531 << " with type: " << *ResultType);
532 MachineIRBuilder MIB(MI);
533 updateRegType(Reg: ResultRegister, Ty: nullptr, SpirvTy: ResultType, GR, MIB, MRI);
534 MIB.setInsertPt(MBB&: *MI.getParent(), II: std::next(x: MI.getIterator()));
535
536 // Tablegen definition assumes SPIRV::ASSIGN_TYPE pseudo-instruction is
537 // present after each auto-folded instruction to take a type reference
538 // from.
539 Register NewReg =
540 MRI.createGenericVirtualRegister(Ty: MRI.getType(Reg: ResultRegister));
541 const auto *RegClass = GR->getRegClass(SpvType: ResultType);
542 MRI.setRegClass(Reg: NewReg, RC: RegClass);
543 MRI.setRegClass(Reg: ResultRegister, RC: RegClass);
544
545 GR->assignSPIRVTypeToVReg(Type: ResultType, VReg: ResultRegister, MF: MIB.getMF());
546 // This is to make it convenient for Legalizer to get the SPIRVType
547 // when processing the actual MI (i.e. not pseudo one).
548 GR->assignSPIRVTypeToVReg(Type: ResultType, VReg: NewReg, MF: MIB.getMF());
549 // Copy MIFlags from Def to ASSIGN_TYPE instruction. It's required to
550 // keep the flags after instruction selection.
551 const uint32_t Flags = MI.getFlags();
552 MIB.buildInstr(Opcode: SPIRV::ASSIGN_TYPE)
553 .addDef(RegNo: ResultRegister)
554 .addUse(RegNo: NewReg)
555 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: ResultType))
556 .setMIFlags(Flags);
557 for (unsigned I = 0, E = MI.getNumDefs(); I != E; ++I) {
558 MachineOperand &MO = MI.getOperand(i: I);
559 if (MO.getReg() == ResultRegister) {
560 MO.setReg(NewReg);
561 break;
562 }
563 }
564}
565
566static void ensureAssignTypeForTypeFolding(MachineFunction &MF,
567 SPIRVGlobalRegistry *GR) {
568 LLVM_DEBUG(dbgs() << "Entering ensureAssignTypeForTypeFolding for function "
569 << MF.getName() << "\n");
570 MachineRegisterInfo &MRI = MF.getRegInfo();
571 for (MachineBasicBlock &MBB : MF) {
572 for (MachineInstr &MI : MBB) {
573 if (!isTypeFoldingSupported(Opcode: MI.getOpcode()))
574 continue;
575
576 LLVM_DEBUG(dbgs() << "Processing instruction: " << MI);
577
578 Register ResultRegister = MI.defs().begin()->getReg();
579 if (hasAssignType(Reg: ResultRegister, MRI)) {
580 LLVM_DEBUG(dbgs() << " Instruction already has ASSIGN_TYPE\n");
581 continue;
582 }
583
584 SPIRVTypeInst ResultType = GR->getSPIRVTypeForVReg(VReg: ResultRegister);
585 generateAssignType(MI, ResultRegister, ResultType, GR, MRI);
586 }
587 }
588}
589
590static bool runPostLegalizer(MachineFunction &MF) {
591 // Initialize the type registry.
592 const SPIRVSubtarget &ST = MF.getSubtarget<SPIRVSubtarget>();
593 SPIRVGlobalRegistry *GR = ST.getSPIRVGlobalRegistry();
594 GR->setCurrentFunc(MF);
595 registerSpirvTypeForNewInstructions(MF, GR);
596 ensureAssignTypeForTypeFolding(MF, GR);
597 return true;
598}
599
600INITIALIZE_PASS(SPIRVPostLegalizerLegacy, DEBUG_TYPE, "SPIRV post legalizer",
601 false, false)
602
603char SPIRVPostLegalizerLegacy::ID = 0;
604
605FunctionPass *llvm::createSPIRVPostLegalizerLegacyPass() {
606 return new SPIRVPostLegalizerLegacy();
607}
608
609bool SPIRVPostLegalizerLegacy::runOnMachineFunction(MachineFunction &MF) {
610 return runPostLegalizer(MF);
611}
612
613PreservedAnalyses
614SPIRVPostLegalizerPass::run(MachineFunction &MF,
615 MachineFunctionAnalysisManager &MFAM) {
616 return runPostLegalizer(MF) ? getMachineFunctionPassPreservedAnalyses()
617 : PreservedAnalyses::all();
618}
619