1//===-- SPIRVGlobalRegistry.cpp - SPIR-V Global Registry --------*- 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// This file contains the implementation of the SPIRVGlobalRegistry class,
10// which is used to maintain rich type information required for SPIR-V even
11// after lowering from LLVM IR to GMIR. It can convert an llvm::Type into
12// an OpTypeXXX instruction, and map it to a virtual register. Also it builds
13// and supports consistency of constants and global variables.
14//
15//===----------------------------------------------------------------------===//
16
17#include "SPIRVGlobalRegistry.h"
18#include "SPIRV.h"
19#include "SPIRVBuiltins.h"
20#include "SPIRVSubtarget.h"
21#include "SPIRVUtils.h"
22#include "llvm/ADT/APInt.h"
23#include "llvm/IR/Constants.h"
24#include "llvm/IR/DiagnosticInfo.h"
25#include "llvm/IR/Function.h"
26#include "llvm/IR/IntrinsicInst.h"
27#include "llvm/IR/Intrinsics.h"
28#include "llvm/IR/IntrinsicsSPIRV.h"
29#include "llvm/IR/Type.h"
30#include "llvm/Support/Casting.h"
31#include "llvm/Support/MathExtras.h"
32#include <cassert>
33#include <functional>
34
35using namespace llvm;
36
37static bool allowEmitFakeUse(const Value *Arg) {
38 if (isSpvIntrinsic(Arg))
39 return false;
40 if (isa<AtomicCmpXchgInst, InsertValueInst, UndefValue>(Val: Arg))
41 return false;
42 if (const auto *LI = dyn_cast<LoadInst>(Val: Arg))
43 if (LI->getType()->isAggregateType())
44 return false;
45 return true;
46}
47
48static unsigned typeToAddressSpace(const Type *Ty) {
49 if (auto PType = dyn_cast<TypedPointerType>(Val: Ty))
50 return PType->getAddressSpace();
51 if (auto PType = dyn_cast<PointerType>(Val: Ty))
52 return PType->getAddressSpace();
53 if (auto *ExtTy = dyn_cast<TargetExtType>(Val: Ty);
54 ExtTy && isTypedPointerWrapper(ExtTy))
55 return ExtTy->getIntParameter(i: 0);
56 reportFatalInternalError(reason: "Unable to convert LLVM type to SPIRVType");
57}
58
59static bool
60storageClassRequiresExplictLayout(SPIRV::StorageClass::StorageClass SC) {
61 switch (SC) {
62 case SPIRV::StorageClass::Uniform:
63 case SPIRV::StorageClass::PushConstant:
64 case SPIRV::StorageClass::StorageBuffer:
65 case SPIRV::StorageClass::PhysicalStorageBufferEXT:
66 return true;
67 case SPIRV::StorageClass::UniformConstant:
68 case SPIRV::StorageClass::Input:
69 case SPIRV::StorageClass::Output:
70 case SPIRV::StorageClass::Workgroup:
71 case SPIRV::StorageClass::CrossWorkgroup:
72 case SPIRV::StorageClass::Private:
73 case SPIRV::StorageClass::Function:
74 case SPIRV::StorageClass::Generic:
75 case SPIRV::StorageClass::AtomicCounter:
76 case SPIRV::StorageClass::Image:
77 case SPIRV::StorageClass::CallableDataNV:
78 case SPIRV::StorageClass::IncomingCallableDataNV:
79 case SPIRV::StorageClass::RayPayloadNV:
80 case SPIRV::StorageClass::HitAttributeNV:
81 case SPIRV::StorageClass::IncomingRayPayloadNV:
82 case SPIRV::StorageClass::ShaderRecordBufferNV:
83 case SPIRV::StorageClass::CodeSectionINTEL:
84 case SPIRV::StorageClass::DeviceOnlyINTEL:
85 case SPIRV::StorageClass::HostOnlyINTEL:
86 return false;
87 }
88 llvm_unreachable("Unknown SPIRV::StorageClass enum");
89}
90
91SPIRVGlobalRegistry::SPIRVGlobalRegistry(DataLayout DL)
92 : DL(DL), Bound(0), CurMF(nullptr) {}
93
94SPIRVTypeInst
95SPIRVGlobalRegistry::assignIntTypeToVReg(unsigned BitWidth, Register VReg,
96 MachineInstr &I,
97 const SPIRVInstrInfo &TII) {
98 SPIRVTypeInst SpirvType = getOrCreateSPIRVIntegerType(BitWidth, I, TII);
99 assignSPIRVTypeToVReg(Type: SpirvType, VReg, MF: *CurMF);
100 return SpirvType;
101}
102
103SPIRVTypeInst
104SPIRVGlobalRegistry::assignFloatTypeToVReg(unsigned BitWidth, Register VReg,
105 MachineInstr &I,
106 const SPIRVInstrInfo &TII) {
107 SPIRVTypeInst SpirvType = getOrCreateSPIRVFloatType(BitWidth, I, TII);
108 assignSPIRVTypeToVReg(Type: SpirvType, VReg, MF: *CurMF);
109 return SpirvType;
110}
111
112SPIRVTypeInst SPIRVGlobalRegistry::assignVectTypeToVReg(
113 SPIRVTypeInst BaseType, unsigned NumElements, Register VReg,
114 MachineInstr &I, const SPIRVInstrInfo &TII) {
115 SPIRVTypeInst SpirvType =
116 getOrCreateSPIRVVectorType(BaseType, NumElements, I, TII);
117 assignSPIRVTypeToVReg(Type: SpirvType, VReg, MF: *CurMF);
118 return SpirvType;
119}
120
121SPIRVTypeInst SPIRVGlobalRegistry::assignTypeToVReg(
122 const Type *Type, Register VReg, MachineIRBuilder &MIRBuilder,
123 SPIRV::AccessQualifier::AccessQualifier AccessQual, bool EmitIR) {
124 SPIRVTypeInst SpirvType =
125 getOrCreateSPIRVType(Type, MIRBuilder, AQ: AccessQual, EmitIR);
126 assignSPIRVTypeToVReg(Type: SpirvType, VReg, MF: MIRBuilder.getMF());
127 return SpirvType;
128}
129
130void SPIRVGlobalRegistry::assignSPIRVTypeToVReg(SPIRVTypeInst SpirvType,
131 Register VReg,
132 const MachineFunction &MF) {
133 VRegToTypeMap[&MF][VReg] = SpirvType;
134}
135
136static Register createTypeVReg(MachineRegisterInfo &MRI) {
137 auto Res = MRI.createGenericVirtualRegister(Ty: LLT::scalar(SizeInBits: 64));
138 MRI.setRegClass(Reg: Res, RC: &SPIRV::TYPERegClass);
139 return Res;
140}
141
142inline Register createTypeVReg(MachineIRBuilder &MIRBuilder) {
143 return createTypeVReg(MRI&: MIRBuilder.getMF().getRegInfo());
144}
145
146SPIRVTypeInst SPIRVGlobalRegistry::getOpTypeBool(MachineIRBuilder &MIRBuilder) {
147 return createConstOrTypeAtFunctionEntry(
148 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
149 return MIRBuilder.buildInstr(Opcode: SPIRV::OpTypeBool)
150 .addDef(RegNo: createTypeVReg(MIRBuilder));
151 });
152}
153
154unsigned SPIRVGlobalRegistry::adjustOpTypeIntWidth(unsigned Width) const {
155 const SPIRVSubtarget &ST = cast<SPIRVSubtarget>(Val: CurMF->getSubtarget());
156 if (ST.canUseExtension(
157 E: SPIRV::Extension::SPV_ALTERA_arbitrary_precision_integers) ||
158 (Width == 4 && ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_int4)))
159 return Width;
160 if (Width <= 8)
161 return 8;
162 else if (Width <= 16)
163 return 16;
164 else if (Width <= 32)
165 return 32;
166 else if (Width <= 64)
167 return 64;
168 else if (Width <= 128)
169 return 128;
170 reportFatalUsageError(reason: "Unsupported Integer width!");
171}
172
173SPIRVTypeInst SPIRVGlobalRegistry::getOpTypeInt(unsigned Width,
174 MachineIRBuilder &MIRBuilder,
175 bool IsSigned) {
176 Width = adjustOpTypeIntWidth(Width);
177 const SPIRVSubtarget &ST =
178 cast<SPIRVSubtarget>(Val: MIRBuilder.getMF().getSubtarget());
179 return createConstOrTypeAtFunctionEntry(MIRBuilder, Op: [&](MachineIRBuilder
180 &MIRBuilder) {
181 if (Width == 4 && ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_int4)) {
182 MIRBuilder.buildInstr(Opcode: SPIRV::OpExtension)
183 .addImm(Val: SPIRV::Extension::SPV_INTEL_int4);
184 MIRBuilder.buildInstr(Opcode: SPIRV::OpCapability)
185 .addImm(Val: SPIRV::Capability::Int4TypeINTEL);
186 } else if ((!isPowerOf2_32(Value: Width) || Width < 8) &&
187 ST.canUseExtension(
188 E: SPIRV::Extension::SPV_ALTERA_arbitrary_precision_integers)) {
189 MIRBuilder.buildInstr(Opcode: SPIRV::OpExtension)
190 .addImm(Val: SPIRV::Extension::SPV_ALTERA_arbitrary_precision_integers);
191 MIRBuilder.buildInstr(Opcode: SPIRV::OpCapability)
192 .addImm(Val: SPIRV::Capability::ArbitraryPrecisionIntegersALTERA);
193 }
194 return MIRBuilder.buildInstr(Opcode: SPIRV::OpTypeInt)
195 .addDef(RegNo: createTypeVReg(MIRBuilder))
196 .addImm(Val: Width)
197 .addImm(Val: IsSigned ? 1 : 0);
198 });
199}
200
201SPIRVTypeInst
202SPIRVGlobalRegistry::getOpTypeFloat(uint32_t Width,
203 MachineIRBuilder &MIRBuilder) {
204 return createConstOrTypeAtFunctionEntry(MIRBuilder, Op: [&](MachineIRBuilder
205 &MIRBuilder) {
206 return MIRBuilder.buildInstr(Opcode: SPIRV::OpTypeFloat)
207 .addDef(RegNo: createTypeVReg(MIRBuilder))
208 .addImm(Val: Width);
209 });
210}
211
212SPIRVTypeInst
213SPIRVGlobalRegistry::getOpTypeFloat(uint32_t Width,
214 MachineIRBuilder &MIRBuilder,
215 SPIRV::FPEncoding::FPEncoding FPEncode) {
216 return createConstOrTypeAtFunctionEntry(MIRBuilder, Op: [&](MachineIRBuilder
217 &MIRBuilder) {
218 return MIRBuilder.buildInstr(Opcode: SPIRV::OpTypeFloat)
219 .addDef(RegNo: createTypeVReg(MIRBuilder))
220 .addImm(Val: Width)
221 .addImm(Val: FPEncode);
222 });
223}
224
225SPIRVTypeInst SPIRVGlobalRegistry::getOpTypeVoid(MachineIRBuilder &MIRBuilder) {
226 return createConstOrTypeAtFunctionEntry(
227 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
228 return MIRBuilder.buildInstr(Opcode: SPIRV::OpTypeVoid)
229 .addDef(RegNo: createTypeVReg(MIRBuilder));
230 });
231}
232
233void SPIRVGlobalRegistry::invalidateMachineInstr(MachineInstr *MI) {
234 // Other maps that may hold MachineInstr*:
235 // - VRegToTypeMap: We cannot remove the definitions of `MI` from
236 // VRegToTypeMap because some calls to invalidateMachineInstr are replacing MI
237 // with another instruction defining the same register. We expect that if MI
238 // is a type instruction, and it is still referenced in VRegToTypeMap, then
239 // those registers are dead or the VRegToTypeMap is out-of-date. We do not
240 // expect passes to ask for the SPIR-V type of a dead register. If the
241 // VRegToTypeMap is out-of-date already, then there was an error before. We
242 // cannot add an assert to verify this because the VRegToTypeMap can be
243 // out-of-date.
244 // - FunctionToInstr & FunctionToInstrRev: At this point, we should not be
245 // deleting functions. No need to update.
246 // - AliasInstMDMap: Would require a linear search, and the Intel Alias
247 // instruction are not instructions instruction selection will be able to
248 // remove.
249
250 const SPIRVSubtarget &ST = MI->getMF()->getSubtarget<SPIRVSubtarget>();
251 [[maybe_unused]] const SPIRVInstrInfo *TII = ST.getInstrInfo();
252 assert(!TII->isAliasingInstr(*MI) &&
253 "Cannot invalidate aliasing instructions.");
254 assert(MI->getOpcode() != SPIRV::OpFunction &&
255 "Cannot invalidate OpFunction.");
256
257 if (MI->getOpcode() == SPIRV::OpFunctionCall) {
258 if (const auto *F = dyn_cast<Function>(Val: MI->getOperand(i: 2).getGlobal())) {
259 auto It = ForwardCalls.find(Val: F);
260 if (It != ForwardCalls.end()) {
261 It->second.erase(Ptr: MI);
262 if (It->second.empty())
263 ForwardCalls.erase(I: It);
264 }
265 }
266 }
267
268 const MachineFunction *MF = MI->getMF();
269 auto It = LastInsertedTypeMap.find(Val: MF);
270 if (It != LastInsertedTypeMap.end() && It->second == MI)
271 LastInsertedTypeMap.erase(Val: MF);
272 // remove from the duplicate tracker to avoid incorrect reuse
273 erase(MI);
274}
275
276const MachineInstr *SPIRVGlobalRegistry::createConstOrTypeAtFunctionEntry(
277 MachineIRBuilder &MIRBuilder,
278 std::function<MachineInstr *(MachineIRBuilder &)> Op) {
279 auto oldInsertPoint = MIRBuilder.getInsertPt();
280 MachineBasicBlock *OldMBB = &MIRBuilder.getMBB();
281 MachineBasicBlock *NewMBB = &*MIRBuilder.getMF().begin();
282
283 auto LastInsertedType = LastInsertedTypeMap.find(Val: CurMF);
284 if (LastInsertedType != LastInsertedTypeMap.end()) {
285 auto It = LastInsertedType->second->getIterator();
286 // It might happen that this instruction was removed from the first MBB,
287 // hence the Parent's check.
288 MachineBasicBlock::iterator InsertAt;
289 if (It->getParent() != NewMBB)
290 InsertAt = oldInsertPoint->getParent() == NewMBB
291 ? oldInsertPoint
292 : getInsertPtValidEnd(MBB: NewMBB);
293 else if (It->getNextNode())
294 InsertAt = It->getNextNode()->getIterator();
295 else
296 InsertAt = getInsertPtValidEnd(MBB: NewMBB);
297 MIRBuilder.setInsertPt(MBB&: *NewMBB, II: InsertAt);
298 } else {
299 MIRBuilder.setInsertPt(MBB&: *NewMBB, II: NewMBB->begin());
300 auto Result = LastInsertedTypeMap.try_emplace(Key: CurMF, Args: nullptr);
301 assert(Result.second);
302 LastInsertedType = Result.first;
303 }
304
305 MachineInstr *ConstOrType = Op(MIRBuilder);
306 // We expect all users of this function to insert definitions at the insertion
307 // point set above that is always the first MBB.
308 assert(ConstOrType->getParent() == NewMBB);
309 LastInsertedType->second = ConstOrType;
310 // Advance past any continued instructions so that the next type/constant
311 // is inserted after the full group, preserving required adjacency.
312 while (auto *Next = LastInsertedType->second->getNextNode()) {
313 unsigned Opc = Next->getOpcode();
314 if (Opc == SPIRV::OpTypeStructContinuedINTEL ||
315 Opc == SPIRV::OpConstantCompositeContinuedINTEL ||
316 Opc == SPIRV::OpSpecConstantCompositeContinuedINTEL ||
317 Opc == SPIRV::OpCompositeConstructContinuedINTEL)
318 LastInsertedType->second = Next;
319 else
320 break;
321 }
322
323 MIRBuilder.setInsertPt(MBB&: *OldMBB, II: oldInsertPoint);
324 return ConstOrType;
325}
326
327SPIRVTypeInst
328SPIRVGlobalRegistry::getOpTypeVector(uint32_t NumElems, SPIRVTypeInst ElemType,
329 MachineIRBuilder &MIRBuilder) {
330 auto EleOpc = ElemType->getOpcode();
331 assert(NumElems >= 2 && "SPIR-V OpTypeVector requires at least 2 components");
332
333 if (EleOpc == SPIRV::OpTypePointer) {
334 if (!cast<SPIRVSubtarget>(Val: MIRBuilder.getMF().getSubtarget())
335 .canUseExtension(
336 E: SPIRV::Extension::SPV_INTEL_masked_gather_scatter)) {
337 const Function &F = MIRBuilder.getMF().getFunction();
338 F.getContext().diagnose(DI: DiagnosticInfoUnsupported(
339 F,
340 "Vector of pointers requires SPV_INTEL_masked_gather_scatter "
341 "extension",
342 DebugLoc(), DS_Error));
343 }
344 } else {
345 assert((EleOpc == SPIRV::OpTypeInt || EleOpc == SPIRV::OpTypeFloat ||
346 EleOpc == SPIRV::OpTypeBool) &&
347 "Invalid vector element type");
348 }
349
350 return createConstOrTypeAtFunctionEntry(
351 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
352 return MIRBuilder.buildInstr(Opcode: SPIRV::OpTypeVector)
353 .addDef(RegNo: createTypeVReg(MIRBuilder))
354 .addUse(RegNo: getSPIRVTypeID(SpirvType: ElemType))
355 .addImm(Val: NumElems);
356 });
357}
358
359Register SPIRVGlobalRegistry::getOrCreateConstFP(APFloat Val, MachineInstr &I,
360 SPIRVTypeInst SpvType,
361 const SPIRVInstrInfo &TII,
362 bool ZeroAsNull) {
363 LLVMContext &Ctx = CurMF->getFunction().getContext();
364 auto *const CF = ConstantFP::get(Context&: Ctx, V: Val);
365 const MachineInstr *MI = findMI(Obj: CF, MF: CurMF);
366 if (MI && (MI->getOpcode() == SPIRV::OpConstantNull ||
367 MI->getOpcode() == SPIRV::OpConstantF))
368 return MI->getOperand(i: 0).getReg();
369 return createConstFP(CF, I, SpvType, TII, ZeroAsNull);
370}
371
372Register SPIRVGlobalRegistry::createConstFP(const ConstantFP *CF,
373 MachineInstr &I,
374 SPIRVTypeInst SpvType,
375 const SPIRVInstrInfo &TII,
376 bool ZeroAsNull) {
377 unsigned BitWidth = getScalarOrVectorBitWidth(Type: SpvType);
378 LLT LLTy = LLT::scalar(SizeInBits: BitWidth);
379 Register Res = CurMF->getRegInfo().createGenericVirtualRegister(Ty: LLTy);
380 CurMF->getRegInfo().setRegClass(Reg: Res, RC: &SPIRV::fIDRegClass);
381 assignSPIRVTypeToVReg(SpirvType: SpvType, VReg: Res, MF: *CurMF);
382
383 MachineInstr *DepMI =
384 const_cast<MachineInstr *>(static_cast<const MachineInstr *>(SpvType));
385 MachineIRBuilder MIRBuilder(*DepMI->getParent(), DepMI->getIterator());
386 const MachineInstr *Const = createConstOrTypeAtFunctionEntry(
387 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
388 MachineInstrBuilder MIB;
389 // In OpenCL OpConstantNull - Scalar floating point: +0.0 (all bits 0)
390 if (CF->getValue().isPosZero() && ZeroAsNull) {
391 MIB = MIRBuilder.buildInstr(Opcode: SPIRV::OpConstantNull)
392 .addDef(RegNo: Res)
393 .addUse(RegNo: getSPIRVTypeID(SpirvType: SpvType));
394 } else {
395 MIB = MIRBuilder.buildInstr(Opcode: SPIRV::OpConstantF)
396 .addDef(RegNo: Res)
397 .addUse(RegNo: getSPIRVTypeID(SpirvType: SpvType));
398 addNumImm(Imm: APInt(BitWidth,
399 CF->getValueAPF().bitcastToAPInt().getZExtValue()),
400 MIB);
401 }
402 const auto &ST = CurMF->getSubtarget();
403 constrainSelectedInstRegOperands(I&: *MIB, TII: *ST.getInstrInfo(),
404 TRI: *ST.getRegisterInfo(),
405 RBI: *ST.getRegBankInfo());
406 return MIB;
407 });
408 add(V: CF, MI: Const);
409 return Res;
410}
411
412Register SPIRVGlobalRegistry::getOrCreateConstInt(uint64_t Val, MachineInstr &I,
413 SPIRVTypeInst SpvType,
414 const SPIRVInstrInfo &TII,
415 bool ZeroAsNull) {
416 return getOrCreateConstInt(Val: APInt(getScalarOrVectorBitWidth(Type: SpvType), Val), I,
417 SpvType, TII, ZeroAsNull);
418}
419
420Register SPIRVGlobalRegistry::getOrCreateConstInt(const APInt &Val,
421 MachineInstr &I,
422 SPIRVTypeInst SpvType,
423 const SPIRVInstrInfo &TII,
424 bool ZeroAsNull) {
425 auto *const CI = ConstantInt::get(
426 Context&: cast<IntegerType>(Val: getTypeForSPIRVType(Ty: SpvType))->getContext(), V: Val);
427 const MachineInstr *MI = findMI(Obj: CI, MF: CurMF);
428 if (MI && (MI->getOpcode() == SPIRV::OpConstantNull ||
429 MI->getOpcode() == SPIRV::OpConstantI))
430 return MI->getOperand(i: 0).getReg();
431 return createConstInt(CI, I, SpvType, TII, ZeroAsNull);
432}
433
434Register SPIRVGlobalRegistry::createConstInt(const ConstantInt *CI,
435 MachineInstr &I,
436 SPIRVTypeInst SpvType,
437 const SPIRVInstrInfo &TII,
438 bool ZeroAsNull) {
439 unsigned BitWidth = getScalarOrVectorBitWidth(Type: SpvType);
440 LLT LLTy = LLT::scalar(SizeInBits: BitWidth);
441 Register Res = CurMF->getRegInfo().createGenericVirtualRegister(Ty: LLTy);
442 CurMF->getRegInfo().setRegClass(Reg: Res, RC: &SPIRV::iIDRegClass);
443 assignIntTypeToVReg(BitWidth, VReg: Res, I, TII);
444
445 MachineInstr *DepMI =
446 const_cast<MachineInstr *>(static_cast<const MachineInstr *>(SpvType));
447 MachineIRBuilder MIRBuilder(*DepMI->getParent(), DepMI->getIterator());
448 const MachineInstr *Const = createConstOrTypeAtFunctionEntry(
449 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
450 MachineInstrBuilder MIB;
451 if (BitWidth == 1) {
452 MIB = MIRBuilder
453 .buildInstr(Opcode: CI->isZero() ? SPIRV::OpConstantFalse
454 : SPIRV::OpConstantTrue)
455 .addDef(RegNo: Res)
456 .addUse(RegNo: getSPIRVTypeID(SpirvType: SpvType));
457 } else if (!CI->isZero() || !ZeroAsNull) {
458 MIB = MIRBuilder.buildInstr(Opcode: SPIRV::OpConstantI)
459 .addDef(RegNo: Res)
460 .addUse(RegNo: getSPIRVTypeID(SpirvType: SpvType));
461 addNumImm(Imm: CI->getValue(), MIB);
462 } else {
463 MIB = MIRBuilder.buildInstr(Opcode: SPIRV::OpConstantNull)
464 .addDef(RegNo: Res)
465 .addUse(RegNo: getSPIRVTypeID(SpirvType: SpvType));
466 }
467 const auto &ST = CurMF->getSubtarget();
468 constrainSelectedInstRegOperands(I&: *MIB, TII: *ST.getInstrInfo(),
469 TRI: *ST.getRegisterInfo(),
470 RBI: *ST.getRegBankInfo());
471 return MIB;
472 });
473 add(V: CI, MI: Const);
474 return Res;
475}
476
477Register SPIRVGlobalRegistry::buildConstantInt(uint64_t Val,
478 MachineIRBuilder &MIRBuilder,
479 SPIRVTypeInst SpvType,
480 bool EmitIR, bool ZeroAsNull) {
481 assert(SpvType);
482 auto &MF = MIRBuilder.getMF();
483 const IntegerType *Ty = cast<IntegerType>(Val: getTypeForSPIRVType(Ty: SpvType));
484 // TODO: Avoid implicit trunc?
485 // See https://github.com/llvm/llvm-project/issues/112510.
486 auto *const CI = ConstantInt::get(Ty: const_cast<IntegerType *>(Ty), V: Val,
487 /*IsSigned=*/false, /*ImplicitTrunc=*/true);
488 Register Res = find(V: CI, MF: &MF);
489 if (Res.isValid())
490 return Res;
491
492 unsigned BitWidth = getScalarOrVectorBitWidth(Type: SpvType);
493 LLT LLTy = LLT::scalar(SizeInBits: BitWidth);
494 MachineRegisterInfo &MRI = MF.getRegInfo();
495 Res = MRI.createGenericVirtualRegister(Ty: LLTy);
496 MRI.setRegClass(Reg: Res, RC: &SPIRV::iIDRegClass);
497 assignTypeToVReg(Type: Ty, VReg: Res, MIRBuilder, AccessQual: SPIRV::AccessQualifier::ReadWrite,
498 EmitIR);
499
500 const MachineInstr *Const = createConstOrTypeAtFunctionEntry(
501 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
502 if (EmitIR)
503 return MIRBuilder.buildConstant(Res, Val: *CI);
504 Register SpvTypeReg = getSPIRVTypeID(SpirvType: SpvType);
505 MachineInstrBuilder MIB;
506 if (Val || !ZeroAsNull) {
507 MIB = MIRBuilder.buildInstr(Opcode: SPIRV::OpConstantI)
508 .addDef(RegNo: Res)
509 .addUse(RegNo: SpvTypeReg);
510 addNumImm(Imm: APInt(BitWidth, Val), MIB);
511 } else {
512 MIB = MIRBuilder.buildInstr(Opcode: SPIRV::OpConstantNull)
513 .addDef(RegNo: Res)
514 .addUse(RegNo: SpvTypeReg);
515 }
516 const auto &Subtarget = CurMF->getSubtarget();
517 constrainSelectedInstRegOperands(I&: *MIB, TII: *Subtarget.getInstrInfo(),
518 TRI: *Subtarget.getRegisterInfo(),
519 RBI: *Subtarget.getRegBankInfo());
520 return MIB;
521 });
522 add(V: CI, MI: Const);
523 return Res;
524}
525
526Register SPIRVGlobalRegistry::buildConstantFP(APFloat Val,
527 MachineIRBuilder &MIRBuilder,
528 SPIRVTypeInst SpvType) {
529 auto &MF = MIRBuilder.getMF();
530 LLVMContext &Ctx = MF.getFunction().getContext();
531 if (!SpvType)
532 SpvType = getOrCreateSPIRVType(Type: Type::getFloatTy(C&: Ctx), MIRBuilder,
533 AQ: SPIRV::AccessQualifier::ReadWrite, EmitIR: true);
534 auto *const CF = ConstantFP::get(Context&: Ctx, V: Val);
535 Register Res = find(V: CF, MF: &MF);
536 if (Res.isValid())
537 return Res;
538
539 LLT LLTy = LLT::scalar(SizeInBits: getScalarOrVectorBitWidth(Type: SpvType));
540 Res = MF.getRegInfo().createGenericVirtualRegister(Ty: LLTy);
541 MF.getRegInfo().setRegClass(Reg: Res, RC: &SPIRV::fIDRegClass);
542 assignSPIRVTypeToVReg(SpirvType: SpvType, VReg: Res, MF);
543
544 const MachineInstr *Const = createConstOrTypeAtFunctionEntry(
545 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
546 MachineInstrBuilder MIB;
547 MIB = MIRBuilder.buildInstr(Opcode: SPIRV::OpConstantF)
548 .addDef(RegNo: Res)
549 .addUse(RegNo: getSPIRVTypeID(SpirvType: SpvType));
550 addNumImm(Imm: CF->getValueAPF().bitcastToAPInt(), MIB);
551 return MIB;
552 });
553 add(V: CF, MI: Const);
554 return Res;
555}
556
557Register SPIRVGlobalRegistry::getOrCreateBaseRegister(
558 Constant *Val, MachineInstr &I, SPIRVTypeInst SpvType,
559 const SPIRVInstrInfo &TII, unsigned BitWidth, bool ZeroAsNull) {
560 SPIRVTypeInst Type = SpvType;
561 if (SpvType->getOpcode() == SPIRV::OpTypeVector ||
562 SpvType->getOpcode() == SPIRV::OpTypeArray) {
563 auto EleTypeReg = SpvType->getOperand(i: 1).getReg();
564 Type = getSPIRVTypeForVReg(VReg: EleTypeReg);
565 }
566 if (Type->getOpcode() == SPIRV::OpTypeFloat) {
567 SPIRVTypeInst SpvBaseType = getOrCreateSPIRVFloatType(BitWidth, I, TII);
568 return getOrCreateConstFP(Val: cast<ConstantFP>(Val)->getValue(), I, SpvType: SpvBaseType,
569 TII, ZeroAsNull);
570 }
571 assert(Type->getOpcode() == SPIRV::OpTypeInt);
572 SPIRVTypeInst SpvBaseType = getOrCreateSPIRVIntegerType(BitWidth, I, TII);
573 return getOrCreateConstInt(Val: Val->getUniqueInteger(), I, SpvType: SpvBaseType, TII,
574 ZeroAsNull);
575}
576
577Register SPIRVGlobalRegistry::getOrCreateCompositeOrNull(
578 Constant *Val, MachineInstr &I, SPIRVTypeInst SpvType,
579 const SPIRVInstrInfo &TII, Constant *CA, unsigned BitWidth,
580 unsigned ElemCnt, bool ZeroAsNull) {
581 if (Register R = find(V: CA, MF: CurMF); R.isValid())
582 return R;
583
584 bool IsNull = Val->isNullValue() && ZeroAsNull;
585 Register ElemReg;
586 if (!IsNull)
587 ElemReg =
588 getOrCreateBaseRegister(Val, I, SpvType, TII, BitWidth, ZeroAsNull);
589
590 LLT LLTy = LLT::scalar(SizeInBits: 64);
591 Register Res = CurMF->getRegInfo().createGenericVirtualRegister(Ty: LLTy);
592 CurMF->getRegInfo().setRegClass(Reg: Res, RC: getRegClass(SpvType));
593 assignSPIRVTypeToVReg(SpirvType: SpvType, VReg: Res, MF: *CurMF);
594
595 MachineInstr *DepMI =
596 const_cast<MachineInstr *>(static_cast<const MachineInstr *>(SpvType));
597 MachineIRBuilder MIRBuilder(*DepMI->getParent(), DepMI->getIterator());
598 const MachineInstr *NewMI = createConstOrTypeAtFunctionEntry(
599 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
600 MachineInstrBuilder MIB;
601 if (!IsNull) {
602 MIB = MIRBuilder.buildInstr(Opcode: SPIRV::OpConstantComposite)
603 .addDef(RegNo: Res)
604 .addUse(RegNo: getSPIRVTypeID(SpirvType: SpvType));
605 for (unsigned i = 0; i < ElemCnt; ++i)
606 MIB.addUse(RegNo: ElemReg);
607 } else {
608 MIB = MIRBuilder.buildInstr(Opcode: SPIRV::OpConstantNull)
609 .addDef(RegNo: Res)
610 .addUse(RegNo: getSPIRVTypeID(SpirvType: SpvType));
611 }
612 const auto &Subtarget = CurMF->getSubtarget();
613 constrainSelectedInstRegOperands(I&: *MIB, TII: *Subtarget.getInstrInfo(),
614 TRI: *Subtarget.getRegisterInfo(),
615 RBI: *Subtarget.getRegBankInfo());
616 return MIB;
617 });
618 add(V: CA, MI: NewMI);
619 return Res;
620}
621
622Register SPIRVGlobalRegistry::getOrCreateConstVector(uint64_t Val,
623 MachineInstr &I,
624 SPIRVTypeInst SpvType,
625 const SPIRVInstrInfo &TII,
626 bool ZeroAsNull) {
627 return getOrCreateConstVector(Val: APInt(getScalarOrVectorBitWidth(Type: SpvType), Val),
628 I, SpvType, TII, ZeroAsNull);
629}
630
631Register SPIRVGlobalRegistry::getOrCreateConstVector(const APInt &Val,
632 MachineInstr &I,
633 SPIRVTypeInst SpvType,
634 const SPIRVInstrInfo &TII,
635 bool ZeroAsNull) {
636 const Type *LLVMTy = getTypeForSPIRVType(Ty: SpvType);
637 assert(LLVMTy->isVectorTy() &&
638 "Expected vector type for constant vector creation");
639 const FixedVectorType *LLVMVecTy = cast<FixedVectorType>(Val: LLVMTy);
640 Type *LLVMBaseTy = LLVMVecTy->getElementType();
641 assert(LLVMBaseTy->isIntegerTy() &&
642 "Expected integer element type for APInt constant vector");
643 auto *ConstVal = cast<ConstantInt>(Val: ConstantInt::get(Ty: LLVMBaseTy, V: Val));
644 auto *ConstVec =
645 ConstantVector::getSplat(EC: LLVMVecTy->getElementCount(), Elt: ConstVal);
646 unsigned BW = getScalarOrVectorBitWidth(Type: SpvType);
647 return getOrCreateCompositeOrNull(Val: ConstVal, I, SpvType, TII, CA: ConstVec, BitWidth: BW,
648 ElemCnt: getScalarOrVectorComponentCount(Type: SpvType),
649 ZeroAsNull);
650}
651
652Register SPIRVGlobalRegistry::getOrCreateConstVector(APFloat Val,
653 MachineInstr &I,
654 SPIRVTypeInst SpvType,
655 const SPIRVInstrInfo &TII,
656 bool ZeroAsNull) {
657 const Type *LLVMTy = getTypeForSPIRVType(Ty: SpvType);
658 assert(LLVMTy->isVectorTy());
659 const FixedVectorType *LLVMVecTy = cast<FixedVectorType>(Val: LLVMTy);
660 Type *LLVMBaseTy = LLVMVecTy->getElementType();
661 assert(LLVMBaseTy->isFloatingPointTy());
662 auto *ConstVal = ConstantFP::get(Ty: LLVMBaseTy, V: Val);
663 auto *ConstVec =
664 ConstantVector::getSplat(EC: LLVMVecTy->getElementCount(), Elt: ConstVal);
665 unsigned BW = getScalarOrVectorBitWidth(Type: SpvType);
666 return getOrCreateCompositeOrNull(Val: ConstVal, I, SpvType, TII, CA: ConstVec, BitWidth: BW,
667 ElemCnt: getScalarOrVectorComponentCount(Type: SpvType),
668 ZeroAsNull);
669}
670
671Register SPIRVGlobalRegistry::getOrCreateConstIntArray(
672 uint64_t Val, size_t Num, MachineInstr &I, SPIRVTypeInst SpvType,
673 const SPIRVInstrInfo &TII) {
674 const Type *LLVMTy = getTypeForSPIRVType(Ty: SpvType);
675 assert(LLVMTy->isArrayTy());
676 const ArrayType *LLVMArrTy = cast<ArrayType>(Val: LLVMTy);
677 Type *LLVMBaseTy = LLVMArrTy->getElementType();
678 Constant *CI = ConstantInt::get(Ty: LLVMBaseTy, V: Val);
679 SPIRVTypeInst SpvBaseTy =
680 getSPIRVTypeForVReg(VReg: SpvType->getOperand(i: 1).getReg());
681 unsigned BW = getScalarOrVectorBitWidth(Type: SpvBaseTy);
682 // The following is reasonably unique key that is better that [Val]. The naive
683 // alternative would be something along the lines of:
684 // SmallVector<Constant *> NumCI(Num, CI);
685 // Constant *UniqueKey =
686 // ConstantArray::get(const_cast<ArrayType*>(LLVMArrTy), NumCI);
687 // that would be a truly unique but dangerous key, because it could lead to
688 // the creation of constants of arbitrary length (that is, the parameter of
689 // memset) which were missing in the original module.
690 Type *I64Ty = Type::getInt64Ty(C&: LLVMBaseTy->getContext());
691 Constant *UniqueKey = ConstantStruct::getAnon(
692 V: {PoisonValue::get(T: const_cast<ArrayType *>(LLVMArrTy)),
693 ConstantInt::get(Ty: LLVMBaseTy, V: Val), ConstantInt::get(Ty: I64Ty, V: Num)});
694 return getOrCreateCompositeOrNull(Val: CI, I, SpvType, TII, CA: UniqueKey, BitWidth: BW,
695 ElemCnt: LLVMArrTy->getNumElements());
696}
697
698Register SPIRVGlobalRegistry::getOrCreateIntCompositeOrNull(
699 uint64_t Val, MachineIRBuilder &MIRBuilder, SPIRVTypeInst SpvType,
700 bool EmitIR, Constant *CA, unsigned BitWidth, unsigned ElemCnt) {
701 if (Register R = find(V: CA, MF: CurMF); R.isValid())
702 return R;
703
704 Register ElemReg;
705 if (Val || EmitIR) {
706 SPIRVTypeInst SpvBaseType =
707 getOrCreateSPIRVIntegerType(BitWidth, MIRBuilder);
708 ElemReg = buildConstantInt(Val, MIRBuilder, SpvType: SpvBaseType, EmitIR);
709 }
710 LLT LLTy = EmitIR ? LLT::fixed_vector(NumElements: ElemCnt, ScalarSizeInBits: BitWidth) : LLT::scalar(SizeInBits: 64);
711 Register Res = CurMF->getRegInfo().createGenericVirtualRegister(Ty: LLTy);
712 CurMF->getRegInfo().setRegClass(Reg: Res, RC: &SPIRV::iIDRegClass);
713 assignSPIRVTypeToVReg(SpirvType: SpvType, VReg: Res, MF: *CurMF);
714
715 const MachineInstr *NewMI = createConstOrTypeAtFunctionEntry(
716 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
717 if (EmitIR)
718 return MIRBuilder.buildSplatBuildVector(Res, Src: ElemReg);
719
720 if (Val) {
721 auto MIB = MIRBuilder.buildInstr(Opcode: SPIRV::OpConstantComposite)
722 .addDef(RegNo: Res)
723 .addUse(RegNo: getSPIRVTypeID(SpirvType: SpvType));
724 for (unsigned i = 0; i < ElemCnt; ++i)
725 MIB.addUse(RegNo: ElemReg);
726 return MIB;
727 }
728
729 return MIRBuilder.buildInstr(Opcode: SPIRV::OpConstantNull)
730 .addDef(RegNo: Res)
731 .addUse(RegNo: getSPIRVTypeID(SpirvType: SpvType));
732 });
733 add(V: CA, MI: NewMI);
734 return Res;
735}
736
737Register SPIRVGlobalRegistry::getOrCreateConsIntVector(
738 uint64_t Val, MachineIRBuilder &MIRBuilder, SPIRVTypeInst SpvType,
739 bool EmitIR) {
740 const Type *LLVMTy = getTypeForSPIRVType(Ty: SpvType);
741 assert(LLVMTy->isVectorTy());
742 const FixedVectorType *LLVMVecTy = cast<FixedVectorType>(Val: LLVMTy);
743 Type *LLVMBaseTy = LLVMVecTy->getElementType();
744 const auto ConstInt = ConstantInt::get(Ty: LLVMBaseTy, V: Val);
745 auto ConstVec =
746 ConstantVector::getSplat(EC: LLVMVecTy->getElementCount(), Elt: ConstInt);
747 unsigned BW = getScalarOrVectorBitWidth(Type: SpvType);
748 return getOrCreateIntCompositeOrNull(
749 Val, MIRBuilder, SpvType, EmitIR, CA: ConstVec, BitWidth: BW,
750 ElemCnt: getScalarOrVectorComponentCount(Type: SpvType));
751}
752
753Register
754SPIRVGlobalRegistry::getOrCreateConstNullPtr(MachineIRBuilder &MIRBuilder,
755 SPIRVTypeInst SpvType) {
756 const Type *Ty = getTypeForSPIRVType(Ty: SpvType);
757 unsigned AddressSpace = typeToAddressSpace(Ty);
758 Type *ElemTy = ::getPointeeType(Ty);
759 assert(ElemTy);
760 const Constant *CP = ConstantTargetNone::get(
761 T: dyn_cast<TargetExtType>(Val: getTypedPointerWrapper(ElemTy, AS: AddressSpace)));
762 Register Res = find(V: CP, MF: CurMF);
763 if (Res.isValid())
764 return Res;
765
766 LLT LLTy = LLT::pointer(AddressSpace, SizeInBits: getPointerSize());
767 Res = CurMF->getRegInfo().createGenericVirtualRegister(Ty: LLTy);
768 CurMF->getRegInfo().setRegClass(Reg: Res, RC: &SPIRV::pIDRegClass);
769 assignSPIRVTypeToVReg(SpirvType: SpvType, VReg: Res, MF: *CurMF);
770
771 const MachineInstr *NewMI = createConstOrTypeAtFunctionEntry(
772 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
773 return MIRBuilder.buildInstr(Opcode: SPIRV::OpConstantNull)
774 .addDef(RegNo: Res)
775 .addUse(RegNo: getSPIRVTypeID(SpirvType: SpvType));
776 });
777 add(V: CP, MI: NewMI);
778 return Res;
779}
780
781Register
782SPIRVGlobalRegistry::buildConstantSampler(Register ResReg, unsigned AddrMode,
783 unsigned Param, unsigned FilerMode,
784 MachineIRBuilder &MIRBuilder) {
785 auto Sampler =
786 ResReg.isValid()
787 ? ResReg
788 : MIRBuilder.getMRI()->createVirtualRegister(RegClass: &SPIRV::iIDRegClass);
789 SPIRVTypeInst TypeSampler = getOrCreateOpTypeSampler(MIRBuilder);
790 Register TypeSamplerReg = getSPIRVTypeID(SpirvType: TypeSampler);
791 // We cannot use createOpType() logic here, because of the
792 // GlobalISel/IRTranslator.cpp check for a tail call that expects that
793 // MIRBuilder.getInsertPt() has a previous instruction. If this constant is
794 // inserted as a result of "__translate_sampler_initializer()" this would
795 // break this IRTranslator assumption.
796 MIRBuilder.buildInstr(Opcode: SPIRV::OpConstantSampler)
797 .addDef(RegNo: Sampler)
798 .addUse(RegNo: TypeSamplerReg)
799 .addImm(Val: AddrMode)
800 .addImm(Val: Param)
801 .addImm(Val: FilerMode);
802 return Sampler;
803}
804
805Register SPIRVGlobalRegistry::buildGlobalVariable(
806 Register ResVReg, SPIRVTypeInst BaseType, StringRef Name,
807 const GlobalValue *GV, SPIRV::StorageClass::StorageClass Storage,
808 const MachineInstr *Init, bool IsConst,
809 const std::optional<SPIRV::LinkageType::LinkageType> &LinkageType,
810 MachineIRBuilder &MIRBuilder, bool IsInstSelector) {
811 const GlobalVariable *GVar = nullptr;
812 if (GV) {
813 GVar = cast<const GlobalVariable>(Val: GV);
814 } else {
815 // If GV is not passed explicitly, use the name to find or construct
816 // the global variable.
817 Module *M = MIRBuilder.getMF().getFunction().getParent();
818 GVar = M->getGlobalVariable(Name);
819 if (GVar == nullptr) {
820 const Type *Ty = getTypeForSPIRVType(Ty: BaseType); // TODO: check type.
821 if (auto *TPTy = dyn_cast<TypedPointerType>(Val: Ty))
822 Ty = PointerType::get(C&: M->getContext(), AddressSpace: TPTy->getAddressSpace());
823 // Module takes ownership of the global var.
824 GVar = new GlobalVariable(*M, const_cast<Type *>(Ty), false,
825 GlobalValue::ExternalLinkage, nullptr,
826 Twine(Name));
827 }
828 GV = GVar;
829 }
830
831 const MachineFunction *MF = &MIRBuilder.getMF();
832 Register Reg = find(V: GVar, MF);
833 if (Reg.isValid()) {
834 if (Reg != ResVReg)
835 MIRBuilder.buildCopy(Res: ResVReg, Op: Reg);
836 return ResVReg;
837 }
838
839 // Emit the OpVariable into the entry block to ensure the def dominates
840 // all uses across all MBBs.
841 MachineBasicBlock &EntryBB = MIRBuilder.getMF().front();
842 MachineIRBuilder GVBuilder(MIRBuilder.getState());
843 if (&GVBuilder.getMBB() != &EntryBB)
844 GVBuilder.setInsertPt(MBB&: EntryBB, II: EntryBB.getFirstTerminator());
845
846 auto MIB = GVBuilder.buildInstr(Opcode: SPIRV::OpVariable)
847 .addDef(RegNo: ResVReg)
848 .addUse(RegNo: getSPIRVTypeID(SpirvType: BaseType))
849 .addImm(Val: static_cast<uint32_t>(Storage));
850 if (Init)
851 MIB.addUse(RegNo: Init->getOperand(i: 0).getReg());
852 // ISel may introduce a new register on this step, so we need to add it to
853 // DT and correct its type avoiding fails on the next stage.
854 if (IsInstSelector) {
855 const auto &Subtarget = CurMF->getSubtarget();
856 constrainSelectedInstRegOperands(I&: *MIB, TII: *Subtarget.getInstrInfo(),
857 TRI: *Subtarget.getRegisterInfo(),
858 RBI: *Subtarget.getRegBankInfo());
859 }
860 add(V: GVar, MI: MIB);
861
862 Reg = MIB->getOperand(i: 0).getReg();
863 addGlobalObject(V: GVar, MF, R: Reg);
864
865 // Set to Reg the same type as ResVReg has.
866 auto MRI = MIRBuilder.getMRI();
867 if (Reg != ResVReg) {
868 LLT RegLLTy =
869 LLT::pointer(AddressSpace: MRI->getType(Reg: ResVReg).getAddressSpace(), SizeInBits: getPointerSize());
870 MRI->setType(VReg: Reg, Ty: RegLLTy);
871 assignSPIRVTypeToVReg(SpirvType: BaseType, VReg: Reg, MF: MIRBuilder.getMF());
872 } else {
873 // Our knowledge about the type may be updated.
874 // If that's the case, we need to update a type
875 // associated with the register.
876 SPIRVTypeInst DefType = getSPIRVTypeForVReg(VReg: ResVReg);
877 if (!DefType || DefType != SPIRVTypeInst(BaseType))
878 assignSPIRVTypeToVReg(SpirvType: BaseType, VReg: Reg, MF: MIRBuilder.getMF());
879 }
880
881 // If it's a global variable with name, output OpName for it.
882 if (GVar && GVar->hasName())
883 buildOpName(Target: Reg, Name: GVar->getName(), MIRBuilder);
884
885 // Output decorations for the GV.
886 // TODO: maybe move to GenerateDecorations pass.
887 const SPIRVSubtarget &ST =
888 cast<SPIRVSubtarget>(Val: MIRBuilder.getMF().getSubtarget());
889 if (IsConst && !ST.isShader())
890 buildOpDecorate(Reg, MIRBuilder, Dec: SPIRV::Decoration::Constant, DecArgs: {});
891
892 if (GVar && GVar->getAlign().valueOrOne().value() != 1 && !ST.isShader()) {
893 unsigned Alignment = (unsigned)GVar->getAlign().valueOrOne().value();
894 buildOpDecorate(Reg, MIRBuilder, Dec: SPIRV::Decoration::Alignment, DecArgs: {Alignment});
895 }
896
897 if (LinkageType)
898 buildOpDecorate(Reg, MIRBuilder, Dec: SPIRV::Decoration::LinkageAttributes,
899 DecArgs: {static_cast<uint32_t>(*LinkageType)}, StrImm: Name);
900
901 SPIRV::BuiltIn::BuiltIn BuiltInId;
902 if (getSpirvBuiltInIdByName(Name, BI&: BuiltInId))
903 buildOpDecorate(Reg, MIRBuilder, Dec: SPIRV::Decoration::BuiltIn,
904 DecArgs: {static_cast<uint32_t>(BuiltInId)});
905
906 // If it's a global variable with "spirv.Decorations" metadata node
907 // recognize it as a SPIR-V friendly LLVM IR and parse "spirv.Decorations"
908 // arguments.
909 MDNode *GVarMD = nullptr;
910 if (GVar && (GVarMD = GVar->getMetadata(Kind: "spirv.Decorations")) != nullptr)
911 buildOpSpirvDecorations(Reg, MIRBuilder, GVarMD, ST);
912
913 return Reg;
914}
915
916// Returns a name based on the Type. Notes that this does not look at
917// decorations, and will return the same string for two types that are the same
918// except for decorations.
919Register SPIRVGlobalRegistry::getOrCreateGlobalVariableWithBinding(
920 SPIRVTypeInst VarType, uint32_t Set, uint32_t Binding, StringRef Name,
921 MachineIRBuilder &MIRBuilder) {
922 Register VarReg =
923 MIRBuilder.getMRI()->createVirtualRegister(RegClass: &SPIRV::iIDRegClass);
924
925 buildGlobalVariable(ResVReg: VarReg, BaseType: VarType, Name, GV: nullptr,
926 Storage: getPointerStorageClass(Type: VarType), Init: nullptr, IsConst: false,
927 LinkageType: std::nullopt, MIRBuilder, IsInstSelector: false);
928
929 buildOpDecorate(Reg: VarReg, MIRBuilder, Dec: SPIRV::Decoration::DescriptorSet, DecArgs: {Set});
930 buildOpDecorate(Reg: VarReg, MIRBuilder, Dec: SPIRV::Decoration::Binding, DecArgs: {Binding});
931 return VarReg;
932}
933
934// TODO: Double check the calls to getOpTypeArray to make sure that `ElemType`
935// is explicitly laid out when required.
936SPIRVTypeInst SPIRVGlobalRegistry::getOpTypeArray(uint32_t NumElems,
937 SPIRVTypeInst ElemType,
938 MachineIRBuilder &MIRBuilder,
939 bool ExplicitLayoutRequired,
940 bool EmitIR) {
941 assert((ElemType->getOpcode() != SPIRV::OpTypeVoid) &&
942 "Invalid array element type");
943 SPIRVTypeInst SpvTypeInt32 = getOrCreateSPIRVIntegerType(BitWidth: 32, MIRBuilder);
944 SPIRVTypeInst ArrayType = nullptr;
945 const SPIRVSubtarget &ST =
946 cast<SPIRVSubtarget>(Val: MIRBuilder.getMF().getSubtarget());
947 if (NumElems != 0) {
948 Register NumElementsVReg =
949 buildConstantInt(Val: NumElems, MIRBuilder, SpvType: SpvTypeInt32, EmitIR);
950 ArrayType = createConstOrTypeAtFunctionEntry(
951 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
952 return MIRBuilder.buildInstr(Opcode: SPIRV::OpTypeArray)
953 .addDef(RegNo: createTypeVReg(MIRBuilder))
954 .addUse(RegNo: getSPIRVTypeID(SpirvType: ElemType))
955 .addUse(RegNo: NumElementsVReg);
956 });
957 } else if (ST.getTargetTriple().getVendor() == Triple::VendorType::AMD) {
958 // We set the array size to the token UINT64_MAX value, which is generally
959 // illegal (the maximum legal size is 61-bits) for the foreseeable future.
960 SPIRVTypeInst SpvTypeInt64 = getOrCreateSPIRVIntegerType(BitWidth: 64, MIRBuilder);
961 Register NumElementsVReg =
962 buildConstantInt(UINT64_MAX, MIRBuilder, SpvType: SpvTypeInt64, EmitIR);
963 ArrayType = createConstOrTypeAtFunctionEntry(
964 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
965 return MIRBuilder.buildInstr(Opcode: SPIRV::OpTypeArray)
966 .addDef(RegNo: createTypeVReg(MIRBuilder))
967 .addUse(RegNo: getSPIRVTypeID(SpirvType: ElemType))
968 .addUse(RegNo: NumElementsVReg);
969 });
970 } else {
971 if (!ST.isShader()) {
972 llvm::reportFatalUsageError(
973 reason: "Runtime arrays are not allowed in non-shader "
974 "SPIR-V modules");
975 return nullptr;
976 }
977 ArrayType = createConstOrTypeAtFunctionEntry(
978 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
979 return MIRBuilder.buildInstr(Opcode: SPIRV::OpTypeRuntimeArray)
980 .addDef(RegNo: createTypeVReg(MIRBuilder))
981 .addUse(RegNo: getSPIRVTypeID(SpirvType: ElemType));
982 });
983 }
984
985 if (ExplicitLayoutRequired && !isResourceType(Type: ElemType)) {
986 Type *ET = const_cast<Type *>(getTypeForSPIRVType(Ty: ElemType));
987 addArrayStrideDecorations(Reg: ArrayType->defs().begin()->getReg(), ElementType: ET,
988 MIRBuilder);
989 }
990
991 return ArrayType;
992}
993
994SPIRVTypeInst
995SPIRVGlobalRegistry::getOpTypeOpaque(const StructType *Ty,
996 MachineIRBuilder &MIRBuilder) {
997 assert(Ty->hasName());
998 StringRef Name = Ty->hasName() ? Ty->getName() : "";
999 Register ResVReg = createTypeVReg(MIRBuilder);
1000 return createConstOrTypeAtFunctionEntry(
1001 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
1002 auto MIB = MIRBuilder.buildInstr(Opcode: SPIRV::OpTypeOpaque).addDef(RegNo: ResVReg);
1003 addStringImm(Str: Name, MIB);
1004 buildOpName(Target: ResVReg, Name, MIRBuilder);
1005 return MIB;
1006 });
1007}
1008
1009SPIRVTypeInst SPIRVGlobalRegistry::getOpTypeStruct(
1010 const StructType *Ty, MachineIRBuilder &MIRBuilder,
1011 SPIRV::AccessQualifier::AccessQualifier AccQual,
1012 StructOffsetDecorator Decorator, bool EmitIR) {
1013 Type *OriginalElementType = nullptr;
1014 uint64_t TotalSize = 0;
1015 if (matchPeeledArrayPattern(Ty, OriginalElementType, TotalSize)) {
1016 SPIRVTypeInst ElementSPIRVType = findSPIRVType(
1017 Ty: OriginalElementType, MIRBuilder, accessQual: AccQual,
1018 /* ExplicitLayoutRequired= */ Decorator != nullptr, EmitIR);
1019 return getOpTypeArray(NumElems: TotalSize, ElemType: ElementSPIRVType, MIRBuilder,
1020 /*ExplicitLayoutRequired=*/Decorator != nullptr,
1021 EmitIR);
1022 }
1023
1024 const SPIRVSubtarget &ST =
1025 cast<SPIRVSubtarget>(Val: MIRBuilder.getMF().getSubtarget());
1026 SmallVector<Register, 4> FieldTypes;
1027 constexpr unsigned MaxWordCount = UINT16_MAX;
1028 const size_t NumElements = Ty->getNumElements();
1029
1030 size_t MaxNumElements = MaxWordCount - 2;
1031 size_t SPIRVStructNumElements = NumElements;
1032 if (NumElements > MaxNumElements) {
1033 // Do adjustments for continued instructions.
1034 SPIRVStructNumElements = MaxNumElements;
1035 MaxNumElements = MaxWordCount - 1;
1036 }
1037
1038 for (const auto &Elem : Ty->elements()) {
1039 SPIRVTypeInst ElemTy = findSPIRVType(
1040 Ty: toTypedPointer(Ty: Elem), MIRBuilder, accessQual: AccQual,
1041 /* ExplicitLayoutRequired= */ Decorator != nullptr, EmitIR);
1042 assert(ElemTy && ElemTy->getOpcode() != SPIRV::OpTypeVoid &&
1043 "Invalid struct element type");
1044 FieldTypes.push_back(Elt: getSPIRVTypeID(SpirvType: ElemTy));
1045 }
1046 Register ResVReg = createTypeVReg(MIRBuilder);
1047 if (Ty->hasName())
1048 buildOpName(Target: ResVReg, Name: Ty->getName(), MIRBuilder);
1049 if (Ty->isPacked() && !ST.isShader())
1050 buildOpDecorate(Reg: ResVReg, MIRBuilder, Dec: SPIRV::Decoration::CPacked, DecArgs: {});
1051
1052 SPIRVTypeInst SPVType = createConstOrTypeAtFunctionEntry(
1053 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
1054 auto MIBStruct =
1055 MIRBuilder.buildInstr(Opcode: SPIRV::OpTypeStruct).addDef(RegNo: ResVReg);
1056 for (size_t I = 0; I < SPIRVStructNumElements; ++I)
1057 MIBStruct.addUse(RegNo: FieldTypes[I]);
1058 for (size_t I = SPIRVStructNumElements; I < NumElements;
1059 I += MaxNumElements) {
1060 auto MIBCont =
1061 MIRBuilder.buildInstr(Opcode: SPIRV::OpTypeStructContinuedINTEL);
1062 for (size_t J = I; J < std::min(a: I + MaxNumElements, b: NumElements); ++J)
1063 MIBCont.addUse(RegNo: FieldTypes[J]);
1064 }
1065 return MIBStruct;
1066 });
1067
1068 if (Decorator)
1069 Decorator(SPVType->defs().begin()->getReg());
1070
1071 return SPVType;
1072}
1073
1074SPIRVTypeInst SPIRVGlobalRegistry::getOrCreateSpecialType(
1075 const Type *Ty, MachineIRBuilder &MIRBuilder,
1076 SPIRV::AccessQualifier::AccessQualifier AccQual) {
1077 assert(isSpecialOpaqueType(Ty) && "Not a special opaque builtin type");
1078 return SPIRV::lowerBuiltinType(Type: Ty, AccessQual: AccQual, MIRBuilder, GR: this);
1079}
1080
1081SPIRVTypeInst SPIRVGlobalRegistry::getOpTypePointer(
1082 SPIRV::StorageClass::StorageClass SC, SPIRVTypeInst ElemType,
1083 MachineIRBuilder &MIRBuilder, Register Reg) {
1084 if (!Reg.isValid())
1085 Reg = createTypeVReg(MIRBuilder);
1086
1087 return createConstOrTypeAtFunctionEntry(MIRBuilder, Op: [&](MachineIRBuilder
1088 &MIRBuilder) {
1089 return MIRBuilder.buildInstr(Opcode: SPIRV::OpTypePointer)
1090 .addDef(RegNo: Reg)
1091 .addImm(Val: static_cast<uint32_t>(SC))
1092 .addUse(RegNo: getSPIRVTypeID(SpirvType: ElemType));
1093 });
1094}
1095
1096SPIRVTypeInst SPIRVGlobalRegistry::getOpTypeForwardPointer(
1097 SPIRV::StorageClass::StorageClass SC, MachineIRBuilder &MIRBuilder) {
1098 return createConstOrTypeAtFunctionEntry(MIRBuilder, Op: [&](MachineIRBuilder
1099 &MIRBuilder) {
1100 return MIRBuilder.buildInstr(Opcode: SPIRV::OpTypeForwardPointer)
1101 .addUse(RegNo: createTypeVReg(MIRBuilder))
1102 .addImm(Val: static_cast<uint32_t>(SC));
1103 });
1104}
1105
1106SPIRVTypeInst SPIRVGlobalRegistry::getOpTypeFunction(
1107 const FunctionType *Ty, SPIRVTypeInst RetType,
1108 const SmallVectorImpl<SPIRVTypeInst> &ArgTypes,
1109 MachineIRBuilder &MIRBuilder) {
1110 const SPIRVSubtarget *ST =
1111 static_cast<const SPIRVSubtarget *>(&MIRBuilder.getMF().getSubtarget());
1112 if (Ty->isVarArg() && ST->isShader()) {
1113 Function &Fn = MIRBuilder.getMF().getFunction();
1114 Ty->getContext().diagnose(DI: DiagnosticInfoUnsupported(
1115 Fn, "SPIR-V shaders do not support variadic functions",
1116 MIRBuilder.getDebugLoc()));
1117 }
1118 return createConstOrTypeAtFunctionEntry(MIRBuilder, Op: [&](MachineIRBuilder
1119 &MIRBuilder) {
1120 auto MIB = MIRBuilder.buildInstr(Opcode: SPIRV::OpTypeFunction)
1121 .addDef(RegNo: createTypeVReg(MIRBuilder))
1122 .addUse(RegNo: getSPIRVTypeID(SpirvType: RetType));
1123 for (auto &ArgType : ArgTypes)
1124 MIB.addUse(RegNo: getSPIRVTypeID(SpirvType: ArgType));
1125 return MIB;
1126 });
1127}
1128
1129SPIRVTypeInst SPIRVGlobalRegistry::getOrCreateOpTypeFunctionWithArgs(
1130 const Type *Ty, SPIRVTypeInst RetType,
1131 const SmallVectorImpl<SPIRVTypeInst> &ArgTypes,
1132 MachineIRBuilder &MIRBuilder) {
1133 if (const MachineInstr *MI = findMI(T: Ty, RequiresExplicitLayout: false, MF: &MIRBuilder.getMF()))
1134 return MI;
1135 const MachineInstr *NewMI =
1136 getOpTypeFunction(Ty: cast<FunctionType>(Val: Ty), RetType, ArgTypes, MIRBuilder);
1137 add(T: Ty, RequiresExplicitLayout: false, MI: NewMI);
1138 return finishCreatingSPIRVType(LLVMTy: Ty, SpirvType: NewMI);
1139}
1140
1141SPIRVTypeInst SPIRVGlobalRegistry::findSPIRVType(
1142 const Type *Ty, MachineIRBuilder &MIRBuilder,
1143 SPIRV::AccessQualifier::AccessQualifier AccQual,
1144 bool ExplicitLayoutRequired, bool EmitIR) {
1145 // Treat <1 x T> as T.
1146 if (auto *FVT = dyn_cast<FixedVectorType>(Val: Ty);
1147 FVT && FVT->getNumElements() == 1)
1148 return findSPIRVType(Ty: FVT->getElementType(), MIRBuilder, AccQual,
1149 ExplicitLayoutRequired, EmitIR);
1150 Ty = adjustIntTypeByWidth(Ty);
1151 // TODO: findMI needs to know if a layout is required.
1152 if (const MachineInstr *MI =
1153 findMI(T: Ty, RequiresExplicitLayout: ExplicitLayoutRequired, MF: &MIRBuilder.getMF()))
1154 return MI;
1155 if (auto It = ForwardPointerTypes.find(Val: Ty); It != ForwardPointerTypes.end())
1156 return It->second;
1157 return restOfCreateSPIRVType(Type: Ty, MIRBuilder, AccessQual: AccQual, ExplicitLayoutRequired,
1158 EmitIR);
1159}
1160
1161Register SPIRVGlobalRegistry::getSPIRVTypeID(SPIRVTypeInst SpirvType) const {
1162 assert(SpirvType && "Attempting to get type id for nullptr type.");
1163 if (SpirvType->getOpcode() == SPIRV::OpTypeForwardPointer ||
1164 SpirvType->getOpcode() == SPIRV::OpTypeStructContinuedINTEL)
1165 return SpirvType->uses().begin()->getReg();
1166 return SpirvType->defs().begin()->getReg();
1167}
1168
1169// We need to use a new LLVM integer type if there is a mismatch between
1170// number of bits in LLVM and SPIRV integer types to let DuplicateTracker
1171// ensure uniqueness of a SPIRV type by the corresponding LLVM type. Without
1172// such an adjustment SPIRVGlobalRegistry::getOpTypeInt() could create the
1173// same "OpTypeInt 8" type for a series of LLVM integer types with number of
1174// bits less than 8. This would lead to duplicate type definitions
1175// eventually due to the method that DuplicateTracker utilizes to reason
1176// about uniqueness of type records.
1177const Type *SPIRVGlobalRegistry::adjustIntTypeByWidth(const Type *Ty) const {
1178 if (auto IType = dyn_cast<IntegerType>(Val: Ty)) {
1179 unsigned SrcBitWidth = IType->getBitWidth();
1180 if (SrcBitWidth > 1) {
1181 unsigned BitWidth = adjustOpTypeIntWidth(Width: SrcBitWidth);
1182 // Maybe change source LLVM type to keep DuplicateTracker consistent.
1183 if (SrcBitWidth != BitWidth)
1184 Ty = IntegerType::get(C&: Ty->getContext(), NumBits: BitWidth);
1185 }
1186 }
1187 return Ty;
1188}
1189
1190SPIRVTypeInst SPIRVGlobalRegistry::createSPIRVType(
1191 const Type *Ty, MachineIRBuilder &MIRBuilder,
1192 SPIRV::AccessQualifier::AccessQualifier AccQual,
1193 bool ExplicitLayoutRequired, bool EmitIR) {
1194 if (isSpecialOpaqueType(Ty))
1195 return getOrCreateSpecialType(Ty, MIRBuilder, AccQual);
1196
1197 if (const MachineInstr *MI =
1198 findMI(T: Ty, RequiresExplicitLayout: ExplicitLayoutRequired, MF: &MIRBuilder.getMF()))
1199 return MI;
1200
1201 if (auto IType = dyn_cast<IntegerType>(Val: Ty)) {
1202 const unsigned Width = IType->getBitWidth();
1203 return Width == 1 ? getOpTypeBool(MIRBuilder)
1204 : getOpTypeInt(Width, MIRBuilder, IsSigned: false);
1205 }
1206 if (Ty->isFloatingPointTy()) {
1207 if (Ty->isFP128Ty() || Ty->isPPC_FP128Ty())
1208 llvm::reportFatalUsageError(reason: "fp128 is not supported in SPIR-V");
1209 if (Ty->isBFloatTy()) {
1210 return getOpTypeFloat(Width: Ty->getPrimitiveSizeInBits(), MIRBuilder,
1211 FPEncode: SPIRV::FPEncoding::BFloat16KHR);
1212 } else {
1213 return getOpTypeFloat(Width: Ty->getPrimitiveSizeInBits(), MIRBuilder);
1214 }
1215 }
1216 if (Ty->isVoidTy())
1217 return getOpTypeVoid(MIRBuilder);
1218 if (Ty->isVectorTy()) {
1219 SPIRVTypeInst El =
1220 findSPIRVType(Ty: cast<FixedVectorType>(Val: Ty)->getElementType(), MIRBuilder,
1221 AccQual, ExplicitLayoutRequired, EmitIR);
1222 return getOpTypeVector(NumElems: cast<FixedVectorType>(Val: Ty)->getNumElements(), ElemType: El,
1223 MIRBuilder);
1224 }
1225 if (Ty->isArrayTy()) {
1226 SPIRVTypeInst El = findSPIRVType(Ty: Ty->getArrayElementType(), MIRBuilder,
1227 AccQual, ExplicitLayoutRequired, EmitIR);
1228 return getOpTypeArray(NumElems: Ty->getArrayNumElements(), ElemType: El, MIRBuilder,
1229 ExplicitLayoutRequired, EmitIR);
1230 }
1231 if (auto SType = dyn_cast<StructType>(Val: Ty)) {
1232 if (SType->isOpaque())
1233 return getOpTypeOpaque(Ty: SType, MIRBuilder);
1234
1235 StructOffsetDecorator Decorator = nullptr;
1236 if (ExplicitLayoutRequired) {
1237 Decorator = [&MIRBuilder, SType, this](Register Reg) {
1238 addStructOffsetDecorations(Reg, Ty: const_cast<StructType *>(SType),
1239 MIRBuilder);
1240 };
1241 }
1242 return getOpTypeStruct(Ty: SType, MIRBuilder, AccQual, Decorator: std::move(Decorator),
1243 EmitIR);
1244 }
1245 if (auto FType = dyn_cast<FunctionType>(Val: Ty)) {
1246 SPIRVTypeInst RetTy =
1247 findSPIRVType(Ty: FType->getReturnType(), MIRBuilder, AccQual,
1248 ExplicitLayoutRequired, EmitIR);
1249 SmallVector<SPIRVTypeInst, 4> ParamTypes;
1250 for (const auto &ParamTy : FType->params())
1251 ParamTypes.push_back(Elt: findSPIRVType(Ty: ParamTy, MIRBuilder, AccQual,
1252 ExplicitLayoutRequired, EmitIR));
1253 return getOpTypeFunction(Ty: FType, RetType: RetTy, ArgTypes: ParamTypes, MIRBuilder);
1254 }
1255
1256 unsigned AddrSpace = typeToAddressSpace(Ty);
1257
1258 // Get access to information about available extensions
1259 const SPIRVSubtarget *ST =
1260 static_cast<const SPIRVSubtarget *>(&MIRBuilder.getMF().getSubtarget());
1261 auto SC = addressSpaceToStorageClass(AddrSpace, STI: *ST);
1262
1263 SPIRVTypeInst SpvElementType = nullptr;
1264 Type *ElemTy = ::getPointeeType(Ty);
1265 if (ElemTy && isa<FunctionType>(Val: ElemTy) &&
1266 !ST->canUseExtension(E: SPIRV::Extension::SPV_INTEL_function_pointers))
1267 ElemTy = nullptr;
1268 if (ElemTy)
1269 SpvElementType = getOrCreateSPIRVType(Type: ElemTy, MIRBuilder, AQ: AccQual, EmitIR);
1270 else
1271 SpvElementType = getOrCreateSPIRVIntegerType(BitWidth: 8, MIRBuilder);
1272
1273 if (!ElemTy) {
1274 ElemTy = Type::getInt8Ty(C&: MIRBuilder.getContext());
1275 }
1276
1277 // If we have forward pointer associated with this type, use its register
1278 // operand to create OpTypePointer.
1279 if (auto It = ForwardPointerTypes.find(Val: Ty); It != ForwardPointerTypes.end()) {
1280 Register Reg = getSPIRVTypeID(SpirvType: It->second);
1281 // TODO: what does getOpTypePointer do?
1282 return getOpTypePointer(SC, ElemType: SpvElementType, MIRBuilder, Reg);
1283 }
1284
1285 return getOrCreateSPIRVPointerType(BaseType: ElemTy, MIRBuilder, SC);
1286}
1287
1288SPIRVTypeInst SPIRVGlobalRegistry::restOfCreateSPIRVType(
1289 const Type *Ty, MachineIRBuilder &MIRBuilder,
1290 SPIRV::AccessQualifier::AccessQualifier AccessQual,
1291 bool ExplicitLayoutRequired, bool EmitIR) {
1292 // TODO: Could this create a problem if one requires an explicit layout, and
1293 // the next time it does not?
1294 if (TypesInProcessing.count(Ptr: Ty) && !isPointerTyOrWrapper(Ty))
1295 return nullptr;
1296 TypesInProcessing.insert(Ptr: Ty);
1297 SPIRVTypeInst SpirvType = createSPIRVType(Ty, MIRBuilder, AccQual: AccessQual,
1298 ExplicitLayoutRequired, EmitIR);
1299 TypesInProcessing.erase(Ptr: Ty);
1300 VRegToTypeMap[&MIRBuilder.getMF()][getSPIRVTypeID(SpirvType)] = SpirvType;
1301
1302 // TODO: We could end up with two SPIR-V types pointing to the same llvm type.
1303 // Is that a problem?
1304 SPIRVToLLVMType[SpirvType] = unifyPtrType(Ty);
1305
1306 if (SpirvType->getOpcode() == SPIRV::OpTypeForwardPointer ||
1307 findMI(T: Ty, RequiresExplicitLayout: false, MF: &MIRBuilder.getMF()) || isSpecialOpaqueType(Ty))
1308 return SpirvType;
1309
1310 if (auto *ExtTy = dyn_cast<TargetExtType>(Val: Ty);
1311 ExtTy && isTypedPointerWrapper(ExtTy))
1312 add(PointeeTy: ExtTy->getTypeParameter(i: 0), AddressSpace: ExtTy->getIntParameter(i: 0), MI: SpirvType);
1313 else if (!isPointerTy(T: Ty))
1314 add(T: Ty, RequiresExplicitLayout: ExplicitLayoutRequired, MI: SpirvType);
1315 else if (isTypedPointerTy(T: Ty))
1316 add(PointeeTy: cast<TypedPointerType>(Val: Ty)->getElementType(),
1317 AddressSpace: getPointerAddressSpace(T: Ty), MI: SpirvType);
1318 else
1319 add(PointeeTy: Type::getInt8Ty(C&: MIRBuilder.getMF().getFunction().getContext()),
1320 AddressSpace: getPointerAddressSpace(T: Ty), MI: SpirvType);
1321 return SpirvType;
1322}
1323
1324SPIRVTypeInst
1325SPIRVGlobalRegistry::getSPIRVTypeForVReg(Register VReg,
1326 const MachineFunction *MF) const {
1327 auto t = VRegToTypeMap.find(Val: MF ? MF : CurMF);
1328 if (t != VRegToTypeMap.end()) {
1329 auto tt = t->second.find(Val: VReg);
1330 if (tt != t->second.end())
1331 return tt->second;
1332 }
1333 return nullptr;
1334}
1335
1336SPIRVTypeInst SPIRVGlobalRegistry::getResultType(Register VReg,
1337 MachineFunction *MF) {
1338 if (!MF)
1339 MF = CurMF;
1340 MachineInstr *Instr = getVRegDef(MRI&: MF->getRegInfo(), Reg: VReg);
1341 return getSPIRVTypeForVReg(VReg: Instr->getOperand(i: 1).getReg(), MF);
1342}
1343
1344SPIRVTypeInst SPIRVGlobalRegistry::getOrCreateSPIRVType(
1345 const Type *Ty, MachineIRBuilder &MIRBuilder,
1346 SPIRV::AccessQualifier::AccessQualifier AccessQual,
1347 bool ExplicitLayoutRequired, bool EmitIR) {
1348 // SPIR-V doesn't support single-element vectors. Treat <1 x T> as T.
1349 if (auto *FVT = dyn_cast<FixedVectorType>(Val: Ty);
1350 FVT && FVT->getNumElements() == 1)
1351 return getOrCreateSPIRVType(Ty: FVT->getElementType(), MIRBuilder, AccessQual,
1352 ExplicitLayoutRequired, EmitIR);
1353 const MachineFunction *MF = &MIRBuilder.getMF();
1354 Register Reg;
1355 if (auto *ExtTy = dyn_cast<TargetExtType>(Val: Ty);
1356 ExtTy && isTypedPointerWrapper(ExtTy))
1357 Reg = find(PointeeTy: ExtTy->getTypeParameter(i: 0), AddressSpace: ExtTy->getIntParameter(i: 0), MF);
1358 else if (!isPointerTy(T: Ty))
1359 Reg = find(T: Ty = adjustIntTypeByWidth(Ty), RequiresExplicitLayout: ExplicitLayoutRequired, MF);
1360 else if (isTypedPointerTy(T: Ty))
1361 Reg = find(PointeeTy: cast<TypedPointerType>(Val: Ty)->getElementType(),
1362 AddressSpace: getPointerAddressSpace(T: Ty), MF);
1363 else
1364 Reg = find(PointeeTy: Type::getInt8Ty(C&: MIRBuilder.getMF().getFunction().getContext()),
1365 AddressSpace: getPointerAddressSpace(T: Ty), MF);
1366 if (Reg.isValid() && !isSpecialOpaqueType(Ty))
1367 return getSPIRVTypeForVReg(VReg: Reg);
1368
1369 TypesInProcessing.clear();
1370 SPIRVTypeInst STy = restOfCreateSPIRVType(Ty, MIRBuilder, AccessQual,
1371 ExplicitLayoutRequired, EmitIR);
1372 // Create normal pointer types for the corresponding OpTypeForwardPointers.
1373 for (auto &CU : ForwardPointerTypes) {
1374 // Pointer type themselves do not require an explicit layout. The types
1375 // they pointer to might, but that is taken care of when creating the type.
1376 bool PtrNeedsLayout = false;
1377 const Type *Ty2 = CU.first;
1378 SPIRVTypeInst STy2 = CU.second;
1379 if ((Reg = find(T: Ty2, RequiresExplicitLayout: PtrNeedsLayout, MF)).isValid())
1380 STy2 = getSPIRVTypeForVReg(VReg: Reg);
1381 else
1382 STy2 = restOfCreateSPIRVType(Ty: Ty2, MIRBuilder, AccessQual, ExplicitLayoutRequired: PtrNeedsLayout,
1383 EmitIR);
1384 if (Ty == Ty2)
1385 STy = STy2;
1386 }
1387 ForwardPointerTypes.clear();
1388 return STy;
1389}
1390
1391bool SPIRVGlobalRegistry::isScalarOfType(Register VReg,
1392 unsigned TypeOpcode) const {
1393 SPIRVTypeInst Type = getSPIRVTypeForVReg(VReg);
1394 assert(Type && "isScalarOfType VReg has no type assigned");
1395 return Type->getOpcode() == TypeOpcode;
1396}
1397
1398bool SPIRVGlobalRegistry::isScalarOrVectorOfType(Register VReg,
1399 unsigned TypeOpcode) const {
1400 SPIRVTypeInst Type = getSPIRVTypeForVReg(VReg);
1401 assert(Type && "isScalarOrVectorOfType VReg has no type assigned");
1402 if (Type->getOpcode() == TypeOpcode)
1403 return true;
1404 if (Type->getOpcode() == SPIRV::OpTypeVector) {
1405 Register ScalarTypeVReg = Type->getOperand(i: 1).getReg();
1406 SPIRVTypeInst ScalarType = getSPIRVTypeForVReg(VReg: ScalarTypeVReg);
1407 return ScalarType->getOpcode() == TypeOpcode;
1408 }
1409 return false;
1410}
1411
1412bool SPIRVGlobalRegistry::isResourceType(SPIRVTypeInst Type) const {
1413 switch (Type->getOpcode()) {
1414 case SPIRV::OpTypeImage:
1415 case SPIRV::OpTypeSampler:
1416 case SPIRV::OpTypeSampledImage:
1417 return true;
1418 case SPIRV::OpTypeStruct:
1419 return hasBlockDecoration(Type);
1420 default:
1421 return false;
1422 }
1423 return false;
1424}
1425unsigned
1426SPIRVGlobalRegistry::getScalarOrVectorComponentCount(Register VReg) const {
1427 return getScalarOrVectorComponentCount(Type: getSPIRVTypeForVReg(VReg));
1428}
1429
1430unsigned
1431SPIRVGlobalRegistry::getScalarOrVectorComponentCount(SPIRVTypeInst Type) const {
1432 if (!Type)
1433 return 0;
1434 return Type->getOpcode() == SPIRV::OpTypeVector
1435 ? static_cast<unsigned>(Type->getOperand(i: 2).getImm())
1436 : 1;
1437}
1438
1439SPIRVTypeInst
1440SPIRVGlobalRegistry::getScalarOrVectorComponentType(SPIRVTypeInst Type) const {
1441 if (!Type)
1442 return nullptr;
1443 Register ScalarReg = Type->getOpcode() == SPIRV::OpTypeVector
1444 ? Type->getOperand(i: 1).getReg()
1445 : Type->getOperand(i: 0).getReg();
1446 SPIRVTypeInst ScalarType = getSPIRVTypeForVReg(VReg: ScalarReg);
1447 assert(isScalarOrVectorOfType(Type->getOperand(0).getReg(),
1448 ScalarType->getOpcode()));
1449 return ScalarType;
1450}
1451
1452unsigned
1453SPIRVGlobalRegistry::getScalarOrVectorBitWidth(SPIRVTypeInst Type) const {
1454 assert(Type && "Invalid Type pointer");
1455 SPIRVTypeInst ScalarType = getScalarOrVectorComponentType(Type);
1456 if (ScalarType->getOpcode() == SPIRV::OpTypeInt ||
1457 ScalarType->getOpcode() == SPIRV::OpTypeFloat)
1458 return ScalarType->getOperand(i: 1).getImm();
1459 if (ScalarType->getOpcode() == SPIRV::OpTypeBool)
1460 return 1;
1461 llvm_unreachable("Attempting to get bit width of non-integer/float type.");
1462}
1463
1464unsigned SPIRVGlobalRegistry::getNumScalarOrVectorTotalBitWidth(
1465 SPIRVTypeInst Type) const {
1466 assert(Type && "Invalid Type pointer");
1467 unsigned NumElements = getScalarOrVectorComponentCount(Type);
1468 SPIRVTypeInst ScalarType = getScalarOrVectorComponentType(Type);
1469 return ScalarType->getOpcode() == SPIRV::OpTypeInt ||
1470 ScalarType->getOpcode() == SPIRV::OpTypeFloat
1471 ? NumElements * ScalarType->getOperand(i: 1).getImm()
1472 : 0;
1473}
1474
1475SPIRVTypeInst
1476SPIRVGlobalRegistry::retrieveScalarOrVectorIntType(SPIRVTypeInst Type) const {
1477 SPIRVTypeInst ScalarType = getScalarOrVectorComponentType(Type);
1478 return ScalarType && ScalarType->getOpcode() == SPIRV::OpTypeInt ? ScalarType
1479 : nullptr;
1480}
1481
1482bool SPIRVGlobalRegistry::isScalarOrVectorSigned(SPIRVTypeInst Type) const {
1483 SPIRVTypeInst IntType = retrieveScalarOrVectorIntType(Type);
1484 return IntType && IntType->getOperand(i: 2).getImm() != 0;
1485}
1486
1487SPIRVTypeInst SPIRVGlobalRegistry::getPointeeType(SPIRVTypeInst PtrType) {
1488 return PtrType && PtrType->getOpcode() == SPIRV::OpTypePointer
1489 ? getSPIRVTypeForVReg(VReg: PtrType->getOperand(i: 2).getReg())
1490 : nullptr;
1491}
1492
1493unsigned SPIRVGlobalRegistry::getPointeeTypeOp(Register PtrReg) {
1494 SPIRVTypeInst ElemType = getPointeeType(PtrType: getSPIRVTypeForVReg(VReg: PtrReg));
1495 return ElemType ? ElemType->getOpcode() : 0;
1496}
1497
1498bool SPIRVGlobalRegistry::isBitcastCompatible(SPIRVTypeInst Type1,
1499 SPIRVTypeInst Type2) const {
1500 if (!Type1 || !Type2)
1501 return false;
1502 auto Op1 = Type1->getOpcode(), Op2 = Type2->getOpcode();
1503 // Ignore difference between <1.5 and >=1.5 protocol versions:
1504 // it's valid if either Result Type or Operand is a pointer, and the other
1505 // is a pointer, an integer scalar, or an integer vector.
1506 if (Op1 == SPIRV::OpTypePointer &&
1507 (Op2 == SPIRV::OpTypePointer || retrieveScalarOrVectorIntType(Type: Type2)))
1508 return true;
1509 if (Op2 == SPIRV::OpTypePointer &&
1510 (Op1 == SPIRV::OpTypePointer || retrieveScalarOrVectorIntType(Type: Type1)))
1511 return true;
1512 unsigned Bits1 = getNumScalarOrVectorTotalBitWidth(Type: Type1),
1513 Bits2 = getNumScalarOrVectorTotalBitWidth(Type: Type2);
1514 return Bits1 > 0 && Bits1 == Bits2;
1515}
1516
1517SPIRV::StorageClass::StorageClass
1518SPIRVGlobalRegistry::getPointerStorageClass(Register VReg) const {
1519 SPIRVTypeInst Type = getSPIRVTypeForVReg(VReg);
1520 assert(Type && Type->getOpcode() == SPIRV::OpTypePointer &&
1521 Type->getOperand(1).isImm() && "Pointer type is expected");
1522 return getPointerStorageClass(Type);
1523}
1524
1525SPIRV::StorageClass::StorageClass
1526SPIRVGlobalRegistry::getPointerStorageClass(SPIRVTypeInst Type) const {
1527 return static_cast<SPIRV::StorageClass::StorageClass>(
1528 Type->getOperand(i: 1).getImm());
1529}
1530
1531SPIRVTypeInst SPIRVGlobalRegistry::getOrCreateVulkanBufferType(
1532 MachineIRBuilder &MIRBuilder, Type *ElemType,
1533 SPIRV::StorageClass::StorageClass SC, bool IsWritable, bool EmitIr) {
1534 auto Key = SPIRV::irhandle_vkbuffer(ElementType: ElemType, SC, IsWriteable: IsWritable);
1535 if (const MachineInstr *MI = findMI(Handle: Key, MF: &MIRBuilder.getMF()))
1536 return MI;
1537
1538 bool ExplicitLayoutRequired = storageClassRequiresExplictLayout(SC);
1539 // We need to get the SPIR-V type for the element here, so we can add the
1540 // decoration to it.
1541 auto *T = StructType::create(Elements: ElemType);
1542 SPIRVTypeInst BlockType =
1543 getOrCreateSPIRVType(Ty: T, MIRBuilder, AccessQual: SPIRV::AccessQualifier::None,
1544 ExplicitLayoutRequired, EmitIR: EmitIr);
1545
1546 buildOpDecorate(Reg: BlockType->defs().begin()->getReg(), MIRBuilder,
1547 Dec: SPIRV::Decoration::Block, DecArgs: {});
1548
1549 if (!IsWritable) {
1550 buildOpMemberDecorate(Reg: BlockType->defs().begin()->getReg(), MIRBuilder,
1551 Dec: SPIRV::Decoration::NonWritable, Member: 0, DecArgs: {});
1552 }
1553
1554 SPIRVTypeInst R =
1555 getOrCreateSPIRVPointerTypeInternal(BaseType: BlockType, MIRBuilder, SC);
1556 add(Handle: Key, MI: R);
1557 return R;
1558}
1559
1560SPIRVTypeInst
1561SPIRVGlobalRegistry::getOrCreatePaddingType(MachineIRBuilder &MIRBuilder) {
1562 auto Key = SPIRV::irhandle_padding();
1563 if (const MachineInstr *MI = findMI(Handle: Key, MF: &MIRBuilder.getMF()))
1564 return MI;
1565 auto *T = Type::getInt8Ty(C&: MIRBuilder.getContext());
1566 SPIRVTypeInst R = getOrCreateSPIRVIntegerType(BitWidth: 8, MIRBuilder);
1567 finishCreatingSPIRVType(LLVMTy: T, SpirvType: R);
1568 add(Handle: Key, MI: R);
1569 return R;
1570}
1571
1572SPIRVTypeInst SPIRVGlobalRegistry::getOrCreateVulkanPushConstantType(
1573 MachineIRBuilder &MIRBuilder, Type *T) {
1574 const auto SC = SPIRV::StorageClass::PushConstant;
1575
1576 auto Key = SPIRV::irhandle_vkbuffer(ElementType: T, SC, /* IsWritable= */ IsWriteable: false);
1577 if (const MachineInstr *MI = findMI(Handle: Key, MF: &MIRBuilder.getMF()))
1578 return MI;
1579
1580 // We need to get the SPIR-V type for the element here, so we can add the
1581 // decoration to it.
1582 SPIRVTypeInst BlockType = getOrCreateSPIRVType(
1583 Ty: T, MIRBuilder, AccessQual: SPIRV::AccessQualifier::None,
1584 /* ExplicitLayoutRequired= */ true, /* EmitIr= */ EmitIR: false);
1585
1586 buildOpDecorate(Reg: BlockType->defs().begin()->getReg(), MIRBuilder,
1587 Dec: SPIRV::Decoration::Block, DecArgs: {});
1588 SPIRVTypeInst R = BlockType;
1589 add(Handle: Key, MI: R);
1590 return R;
1591}
1592
1593SPIRVTypeInst SPIRVGlobalRegistry::getOrCreateLayoutType(
1594 MachineIRBuilder &MIRBuilder, const TargetExtType *T, bool EmitIr) {
1595 auto Key = SPIRV::handle(Ty: T);
1596 if (const MachineInstr *MI = findMI(Handle: Key, MF: &MIRBuilder.getMF()))
1597 return MI;
1598
1599 StructType *ST = cast<StructType>(Val: T->getTypeParameter(i: 0));
1600 ArrayRef<uint32_t> Offsets = T->int_params().slice(N: 1);
1601 assert(ST->getNumElements() == Offsets.size());
1602
1603 StructOffsetDecorator Decorator = [&MIRBuilder, &Offsets](Register Reg) {
1604 for (uint32_t I = 0; I < Offsets.size(); ++I) {
1605 buildOpMemberDecorate(Reg, MIRBuilder, Dec: SPIRV::Decoration::Offset, Member: I,
1606 DecArgs: {Offsets[I]});
1607 }
1608 };
1609
1610 // We need a new OpTypeStruct instruction because decorations will be
1611 // different from a struct with an explicit layout created from a different
1612 // entry point.
1613 SPIRVTypeInst SPIRVStructType =
1614 getOpTypeStruct(Ty: ST, MIRBuilder, AccQual: SPIRV::AccessQualifier::None,
1615 Decorator: std::move(Decorator), EmitIR: EmitIr);
1616 add(Handle: Key, MI: SPIRVStructType);
1617 return SPIRVStructType;
1618}
1619
1620SPIRVTypeInst SPIRVGlobalRegistry::getImageType(
1621 const TargetExtType *ExtensionType,
1622 const SPIRV::AccessQualifier::AccessQualifier Qualifier,
1623 MachineIRBuilder &MIRBuilder) {
1624 assert(ExtensionType->getNumTypeParameters() == 1 &&
1625 "SPIR-V image builtin type must have sampled type parameter!");
1626 const SPIRVTypeInst SampledType =
1627 getOrCreateSPIRVType(Type: ExtensionType->getTypeParameter(i: 0), MIRBuilder,
1628 AQ: SPIRV::AccessQualifier::ReadWrite, EmitIR: true);
1629 assert((ExtensionType->getNumIntParameters() == 7 ||
1630 ExtensionType->getNumIntParameters() == 6) &&
1631 "Invalid number of parameters for SPIR-V image builtin!");
1632
1633 SPIRV::AccessQualifier::AccessQualifier accessQualifier =
1634 SPIRV::AccessQualifier::None;
1635 if (ExtensionType->getNumIntParameters() == 7) {
1636 accessQualifier = Qualifier == SPIRV::AccessQualifier::WriteOnly
1637 ? SPIRV::AccessQualifier::WriteOnly
1638 : SPIRV::AccessQualifier::AccessQualifier(
1639 ExtensionType->getIntParameter(i: 6));
1640 }
1641
1642 // Create or get an existing type from GlobalRegistry.
1643 SPIRVTypeInst R = getOrCreateOpTypeImage(
1644 MIRBuilder, SampledType,
1645 Dim: SPIRV::Dim::Dim(ExtensionType->getIntParameter(i: 0)),
1646 Depth: ExtensionType->getIntParameter(i: 1), Arrayed: ExtensionType->getIntParameter(i: 2),
1647 Multisampled: ExtensionType->getIntParameter(i: 3), Sampled: ExtensionType->getIntParameter(i: 4),
1648 ImageFormat: SPIRV::ImageFormat::ImageFormat(ExtensionType->getIntParameter(i: 5)),
1649 AccQual: accessQualifier);
1650 SPIRVToLLVMType[R] = ExtensionType;
1651 return R;
1652}
1653
1654SPIRVTypeInst SPIRVGlobalRegistry::getOrCreateOpTypeImage(
1655 MachineIRBuilder &MIRBuilder, SPIRVTypeInst SampledType,
1656 SPIRV::Dim::Dim Dim, uint32_t Depth, uint32_t Arrayed,
1657 uint32_t Multisampled, uint32_t Sampled,
1658 SPIRV::ImageFormat::ImageFormat ImageFormat,
1659 SPIRV::AccessQualifier::AccessQualifier AccessQual) {
1660 auto Key = SPIRV::irhandle_image(SampledTy: SPIRVToLLVMType.lookup(Val: SampledType), Dim,
1661 Depth, Arrayed, MS: Multisampled, Sampled,
1662 ImageFormat, AQ: AccessQual);
1663 if (const MachineInstr *MI = findMI(Handle: Key, MF: &MIRBuilder.getMF()))
1664 return MI;
1665 const MachineInstr *NewMI = createConstOrTypeAtFunctionEntry(
1666 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
1667 auto MIB =
1668 MIRBuilder.buildInstr(Opcode: SPIRV::OpTypeImage)
1669 .addDef(RegNo: createTypeVReg(MIRBuilder))
1670 .addUse(RegNo: getSPIRVTypeID(SpirvType: SampledType))
1671 .addImm(Val: Dim)
1672 .addImm(Val: Depth) // Depth (whether or not it is a Depth image).
1673 .addImm(Val: Arrayed) // Arrayed.
1674 .addImm(Val: Multisampled) // Multisampled (0 = only single-sample).
1675 .addImm(Val: Sampled) // Sampled (0 = usage known at runtime).
1676 .addImm(Val: ImageFormat);
1677 if (AccessQual != SPIRV::AccessQualifier::None)
1678 MIB.addImm(Val: AccessQual);
1679 return MIB;
1680 });
1681 add(Handle: Key, MI: NewMI);
1682 return NewMI;
1683}
1684
1685SPIRVTypeInst
1686SPIRVGlobalRegistry::getOrCreateOpTypeSampler(MachineIRBuilder &MIRBuilder) {
1687 auto Key = SPIRV::irhandle_sampler();
1688 const MachineFunction *MF = &MIRBuilder.getMF();
1689 if (const MachineInstr *MI = findMI(Handle: Key, MF))
1690 return MI;
1691 const MachineInstr *NewMI = createConstOrTypeAtFunctionEntry(
1692 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
1693 return MIRBuilder.buildInstr(Opcode: SPIRV::OpTypeSampler)
1694 .addDef(RegNo: createTypeVReg(MIRBuilder));
1695 });
1696 add(Handle: Key, MI: NewMI);
1697 return NewMI;
1698}
1699
1700SPIRVTypeInst SPIRVGlobalRegistry::getOrCreateOpTypePipe(
1701 MachineIRBuilder &MIRBuilder,
1702 SPIRV::AccessQualifier::AccessQualifier AccessQual) {
1703 auto Key = SPIRV::irhandle_pipe(AQ: AccessQual);
1704 if (const MachineInstr *MI = findMI(Handle: Key, MF: &MIRBuilder.getMF()))
1705 return MI;
1706 const MachineInstr *NewMI = createConstOrTypeAtFunctionEntry(
1707 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
1708 return MIRBuilder.buildInstr(Opcode: SPIRV::OpTypePipe)
1709 .addDef(RegNo: createTypeVReg(MIRBuilder))
1710 .addImm(Val: AccessQual);
1711 });
1712 add(Handle: Key, MI: NewMI);
1713 return NewMI;
1714}
1715
1716SPIRVTypeInst SPIRVGlobalRegistry::getOrCreateOpTypeDeviceEvent(
1717 MachineIRBuilder &MIRBuilder) {
1718 auto Key = SPIRV::irhandle_event();
1719 if (const MachineInstr *MI = findMI(Handle: Key, MF: &MIRBuilder.getMF()))
1720 return MI;
1721 const MachineInstr *NewMI = createConstOrTypeAtFunctionEntry(
1722 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
1723 return MIRBuilder.buildInstr(Opcode: SPIRV::OpTypeDeviceEvent)
1724 .addDef(RegNo: createTypeVReg(MIRBuilder));
1725 });
1726 add(Handle: Key, MI: NewMI);
1727 return NewMI;
1728}
1729
1730SPIRVTypeInst SPIRVGlobalRegistry::getOrCreateOpTypeSampledImage(
1731 SPIRVTypeInst ImageType, MachineIRBuilder &MIRBuilder) {
1732 auto Key = SPIRV::irhandle_sampled_image(
1733 SampledTy: SPIRVToLLVMType.lookup(Val: MIRBuilder.getMF().getRegInfo().getVRegDef(
1734 Reg: ImageType->getOperand(i: 1).getReg())),
1735 ImageTy: ImageType);
1736 if (const MachineInstr *MI = findMI(Handle: Key, MF: &MIRBuilder.getMF()))
1737 return MI;
1738 const MachineInstr *NewMI = createConstOrTypeAtFunctionEntry(
1739 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
1740 return MIRBuilder.buildInstr(Opcode: SPIRV::OpTypeSampledImage)
1741 .addDef(RegNo: createTypeVReg(MIRBuilder))
1742 .addUse(RegNo: getSPIRVTypeID(SpirvType: ImageType));
1743 });
1744 add(Handle: Key, MI: NewMI);
1745 return NewMI;
1746}
1747
1748SPIRVTypeInst SPIRVGlobalRegistry::getOrCreateOpTypeCoopMatr(
1749 MachineIRBuilder &MIRBuilder, const TargetExtType *ExtensionType,
1750 SPIRVTypeInst ElemType, uint32_t Scope, uint32_t Rows, uint32_t Columns,
1751 uint32_t Use, bool EmitIR) {
1752 if (const MachineInstr *MI =
1753 findMI(T: ExtensionType, RequiresExplicitLayout: false, MF: &MIRBuilder.getMF()))
1754 return MI;
1755 const MachineInstr *NewMI = createConstOrTypeAtFunctionEntry(
1756 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
1757 SPIRVTypeInst SpvTypeInt32 =
1758 getOrCreateSPIRVIntegerType(BitWidth: 32, MIRBuilder);
1759 const Type *ET = getTypeForSPIRVType(Ty: ElemType);
1760 if (ET->isIntegerTy() && ET->getIntegerBitWidth() == 4 &&
1761 cast<SPIRVSubtarget>(Val: MIRBuilder.getMF().getSubtarget())
1762 .canUseExtension(E: SPIRV::Extension::SPV_INTEL_int4)) {
1763 MIRBuilder.buildInstr(Opcode: SPIRV::OpCapability)
1764 .addImm(Val: SPIRV::Capability::Int4CooperativeMatrixINTEL);
1765 }
1766 return MIRBuilder.buildInstr(Opcode: SPIRV::OpTypeCooperativeMatrixKHR)
1767 .addDef(RegNo: createTypeVReg(MIRBuilder))
1768 .addUse(RegNo: getSPIRVTypeID(SpirvType: ElemType))
1769 .addUse(RegNo: buildConstantInt(Val: Scope, MIRBuilder, SpvType: SpvTypeInt32, EmitIR))
1770 .addUse(RegNo: buildConstantInt(Val: Rows, MIRBuilder, SpvType: SpvTypeInt32, EmitIR))
1771 .addUse(RegNo: buildConstantInt(Val: Columns, MIRBuilder, SpvType: SpvTypeInt32, EmitIR))
1772 .addUse(RegNo: buildConstantInt(Val: Use, MIRBuilder, SpvType: SpvTypeInt32, EmitIR));
1773 });
1774 add(T: ExtensionType, RequiresExplicitLayout: false, MI: NewMI);
1775 return NewMI;
1776}
1777
1778SPIRVTypeInst SPIRVGlobalRegistry::getOrCreateOpTypeByOpcode(
1779 const Type *Ty, MachineIRBuilder &MIRBuilder, unsigned Opcode) {
1780 if (const MachineInstr *MI = findMI(T: Ty, RequiresExplicitLayout: false, MF: &MIRBuilder.getMF()))
1781 return MI;
1782 const MachineInstr *NewMI = createConstOrTypeAtFunctionEntry(
1783 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
1784 return MIRBuilder.buildInstr(Opcode).addDef(RegNo: createTypeVReg(MIRBuilder));
1785 });
1786 add(T: Ty, RequiresExplicitLayout: false, MI: NewMI);
1787 return NewMI;
1788}
1789
1790SPIRVTypeInst SPIRVGlobalRegistry::getOrCreateUnknownType(
1791 const Type *Ty, MachineIRBuilder &MIRBuilder, unsigned Opcode,
1792 const ArrayRef<MCOperand> Operands) {
1793 if (const MachineInstr *MI = findMI(T: Ty, RequiresExplicitLayout: false, MF: &MIRBuilder.getMF()))
1794 return MI;
1795 Register ResVReg = createTypeVReg(MIRBuilder);
1796 const MachineInstr *NewMI = createConstOrTypeAtFunctionEntry(
1797 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
1798 MachineInstrBuilder MIB = MIRBuilder.buildInstr(Opcode: SPIRV::UNKNOWN_type)
1799 .addDef(RegNo: ResVReg)
1800 .addImm(Val: Opcode);
1801 for (MCOperand Operand : Operands) {
1802 if (Operand.isReg()) {
1803 MIB.addUse(RegNo: Operand.getReg());
1804 } else if (Operand.isImm()) {
1805 MIB.addImm(Val: Operand.getImm());
1806 }
1807 }
1808 return MIB;
1809 });
1810 add(T: Ty, RequiresExplicitLayout: false, MI: NewMI);
1811 return NewMI;
1812}
1813
1814// Returns nullptr if unable to recognize SPIRV type name
1815SPIRVTypeInst SPIRVGlobalRegistry::getOrCreateSPIRVTypeByName(
1816 StringRef TypeStr, MachineIRBuilder &MIRBuilder, bool EmitIR,
1817 SPIRV::StorageClass::StorageClass SC,
1818 SPIRV::AccessQualifier::AccessQualifier AQ) {
1819 unsigned VecElts = 0;
1820 auto &Ctx = MIRBuilder.getMF().getFunction().getContext();
1821
1822 // Parse strings representing either a SPIR-V or OpenCL builtin type.
1823 if (hasBuiltinTypePrefix(Name: TypeStr))
1824 return getOrCreateSPIRVType(Ty: SPIRV::parseBuiltinTypeNameToTargetExtType(
1825 TypeName: TypeStr.str(), Context&: MIRBuilder.getContext()),
1826 MIRBuilder, AccessQual: AQ, ExplicitLayoutRequired: false, EmitIR: true);
1827
1828 // Parse type name in either "typeN" or "type vector[N]" format, where
1829 // N is the number of elements of the vector.
1830 Type *Ty;
1831
1832 Ty = parseBasicTypeName(TypeName&: TypeStr, Ctx);
1833 if (!Ty)
1834 // Unable to recognize SPIRV type name
1835 return nullptr;
1836
1837 SPIRVTypeInst SpirvTy = getOrCreateSPIRVType(Ty, MIRBuilder, AccessQual: AQ, ExplicitLayoutRequired: false, EmitIR: true);
1838
1839 // Handle "type*" or "type* vector[N]".
1840 if (TypeStr.consume_front(Prefix: "*"))
1841 SpirvTy = getOrCreateSPIRVPointerType(BaseType: Ty, MIRBuilder, SC);
1842
1843 // Handle "typeN*" or "type vector[N]*".
1844 bool IsPtrToVec = TypeStr.consume_back(Suffix: "*");
1845
1846 if (TypeStr.consume_front(Prefix: " vector[")) {
1847 TypeStr = TypeStr.substr(Start: 0, N: TypeStr.find(C: ']'));
1848 }
1849 TypeStr.getAsInteger(Radix: 10, Result&: VecElts);
1850 if (VecElts > 0)
1851 SpirvTy = getOrCreateSPIRVVectorType(BaseType: SpirvTy, NumElements: VecElts, MIRBuilder, EmitIR);
1852
1853 if (IsPtrToVec)
1854 SpirvTy = getOrCreateSPIRVPointerType(BaseType: SpirvTy, MIRBuilder, SC);
1855
1856 return SpirvTy;
1857}
1858
1859SPIRVTypeInst
1860SPIRVGlobalRegistry::getOrCreateSPIRVIntegerType(unsigned BitWidth,
1861 MachineIRBuilder &MIRBuilder) {
1862 return getOrCreateSPIRVType(
1863 Ty: IntegerType::get(C&: MIRBuilder.getMF().getFunction().getContext(), NumBits: BitWidth),
1864 MIRBuilder, AccessQual: SPIRV::AccessQualifier::ReadWrite, ExplicitLayoutRequired: false, EmitIR: true);
1865}
1866
1867SPIRVTypeInst
1868SPIRVGlobalRegistry::finishCreatingSPIRVType(const Type *LLVMTy,
1869 SPIRVTypeInst SpirvType) {
1870 assert(CurMF == SpirvType->getMF());
1871 VRegToTypeMap[CurMF][getSPIRVTypeID(SpirvType)] = SpirvType;
1872 SPIRVToLLVMType[SpirvType] = unifyPtrType(Ty: LLVMTy);
1873 return SpirvType;
1874}
1875
1876SPIRVTypeInst
1877SPIRVGlobalRegistry::getOrCreateSPIRVType(unsigned BitWidth, MachineInstr &I,
1878 const SPIRVInstrInfo &TII,
1879 unsigned SPIRVOPcode, Type *Ty) {
1880 if (const MachineInstr *MI = findMI(T: Ty, RequiresExplicitLayout: false, MF: CurMF))
1881 return MI;
1882 MachineBasicBlock &DepMBB = I.getMF()->front();
1883 MachineIRBuilder MIRBuilder(DepMBB, DepMBB.getFirstNonPHI());
1884 const MachineInstr *NewMI = createConstOrTypeAtFunctionEntry(
1885 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
1886 auto NewTypeMI = BuildMI(BB&: MIRBuilder.getMBB(), I&: *MIRBuilder.getInsertPt(),
1887 MIMD: MIRBuilder.getDL(), MCID: TII.get(Opcode: SPIRVOPcode))
1888 .addDef(RegNo: createTypeVReg(MRI&: CurMF->getRegInfo()))
1889 .addImm(Val: BitWidth);
1890 // Don't add Encoding to FP type
1891 if (!Ty->isFloatTy()) {
1892 return NewTypeMI.addImm(Val: 0);
1893 } else {
1894 return NewTypeMI;
1895 }
1896 });
1897 add(T: Ty, RequiresExplicitLayout: false, MI: NewMI);
1898 return finishCreatingSPIRVType(LLVMTy: Ty, SpirvType: NewMI);
1899}
1900
1901SPIRVTypeInst SPIRVGlobalRegistry::getOrCreateSPIRVIntegerType(
1902 unsigned BitWidth, MachineInstr &I, const SPIRVInstrInfo &TII) {
1903 // Maybe adjust bit width to keep DuplicateTracker consistent. Without
1904 // such an adjustment SPIRVGlobalRegistry::getOpTypeInt() could create, for
1905 // example, the same "OpTypeInt 8" type for a series of LLVM integer types
1906 // with number of bits less than 8, causing duplicate type definitions.
1907 if (BitWidth > 1)
1908 BitWidth = adjustOpTypeIntWidth(Width: BitWidth);
1909 Type *LLVMTy = IntegerType::get(C&: CurMF->getFunction().getContext(), NumBits: BitWidth);
1910 return getOrCreateSPIRVType(BitWidth, I, TII, SPIRVOPcode: SPIRV::OpTypeInt, Ty: LLVMTy);
1911}
1912
1913SPIRVTypeInst SPIRVGlobalRegistry::getOrCreateSPIRVFloatType(
1914 unsigned BitWidth, MachineInstr &I, const SPIRVInstrInfo &TII) {
1915 LLVMContext &Ctx = CurMF->getFunction().getContext();
1916 Type *LLVMTy;
1917 switch (BitWidth) {
1918 case 16:
1919 LLVMTy = Type::getHalfTy(C&: Ctx);
1920 break;
1921 case 32:
1922 LLVMTy = Type::getFloatTy(C&: Ctx);
1923 break;
1924 case 64:
1925 LLVMTy = Type::getDoubleTy(C&: Ctx);
1926 break;
1927 default:
1928 llvm_unreachable("Bit width is of unexpected size.");
1929 }
1930 return getOrCreateSPIRVType(BitWidth, I, TII, SPIRVOPcode: SPIRV::OpTypeFloat, Ty: LLVMTy);
1931}
1932
1933SPIRVTypeInst
1934SPIRVGlobalRegistry::getOrCreateSPIRVBoolType(MachineIRBuilder &MIRBuilder,
1935 bool EmitIR) {
1936 return getOrCreateSPIRVType(
1937 Ty: IntegerType::get(C&: MIRBuilder.getMF().getFunction().getContext(), NumBits: 1),
1938 MIRBuilder, AccessQual: SPIRV::AccessQualifier::ReadWrite, ExplicitLayoutRequired: false, EmitIR);
1939}
1940
1941SPIRVTypeInst
1942SPIRVGlobalRegistry::getOrCreateSPIRVBoolType(MachineInstr &I,
1943 const SPIRVInstrInfo &TII) {
1944 Type *Ty = IntegerType::get(C&: CurMF->getFunction().getContext(), NumBits: 1);
1945 if (const MachineInstr *MI = findMI(T: Ty, RequiresExplicitLayout: false, MF: CurMF))
1946 return MI;
1947 MachineBasicBlock &DepMBB = I.getMF()->front();
1948 MachineIRBuilder MIRBuilder(DepMBB, DepMBB.getFirstNonPHI());
1949 const MachineInstr *NewMI = createConstOrTypeAtFunctionEntry(
1950 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
1951 return BuildMI(BB&: MIRBuilder.getMBB(), I&: *MIRBuilder.getInsertPt(),
1952 MIMD: MIRBuilder.getDL(), MCID: TII.get(Opcode: SPIRV::OpTypeBool))
1953 .addDef(RegNo: createTypeVReg(MRI&: CurMF->getRegInfo()));
1954 });
1955 add(T: Ty, RequiresExplicitLayout: false, MI: NewMI);
1956 return finishCreatingSPIRVType(LLVMTy: Ty, SpirvType: NewMI);
1957}
1958
1959SPIRVTypeInst SPIRVGlobalRegistry::getOrCreateSPIRVVectorType(
1960 SPIRVTypeInst BaseType, unsigned NumElements, MachineIRBuilder &MIRBuilder,
1961 bool EmitIR) {
1962 return getOrCreateSPIRVType(
1963 Ty: FixedVectorType::get(ElementType: const_cast<Type *>(getTypeForSPIRVType(Ty: BaseType)),
1964 NumElts: NumElements),
1965 MIRBuilder, AccessQual: SPIRV::AccessQualifier::ReadWrite, ExplicitLayoutRequired: false, EmitIR);
1966}
1967
1968SPIRVTypeInst SPIRVGlobalRegistry::getOrCreateSPIRVVectorType(
1969 SPIRVTypeInst BaseType, unsigned NumElements, MachineInstr &I,
1970 const SPIRVInstrInfo &TII) {
1971 // At this point of time all 1-element vectors are resolved. Add assertion
1972 // to fire if anything changes.
1973 assert(NumElements >= 2 && "SPIR-V vectors must have at least 2 components");
1974 Type *Ty = FixedVectorType::get(
1975 ElementType: const_cast<Type *>(getTypeForSPIRVType(Ty: BaseType)), NumElts: NumElements);
1976 if (const MachineInstr *MI = findMI(T: Ty, RequiresExplicitLayout: false, MF: CurMF))
1977 return MI;
1978 MachineInstr *DepMI =
1979 const_cast<MachineInstr *>(static_cast<const MachineInstr *>(BaseType));
1980 MachineIRBuilder MIRBuilder(*DepMI->getParent(), DepMI->getIterator());
1981 const MachineInstr *NewMI = createConstOrTypeAtFunctionEntry(
1982 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
1983 return BuildMI(BB&: MIRBuilder.getMBB(), I&: *MIRBuilder.getInsertPt(),
1984 MIMD: MIRBuilder.getDL(), MCID: TII.get(Opcode: SPIRV::OpTypeVector))
1985 .addDef(RegNo: createTypeVReg(MRI&: CurMF->getRegInfo()))
1986 .addUse(RegNo: getSPIRVTypeID(SpirvType: BaseType))
1987 .addImm(Val: NumElements);
1988 });
1989 add(T: Ty, RequiresExplicitLayout: false, MI: NewMI);
1990 return finishCreatingSPIRVType(LLVMTy: Ty, SpirvType: NewMI);
1991}
1992
1993SPIRVTypeInst SPIRVGlobalRegistry::getOrCreateSPIRVPointerType(
1994 const Type *BaseType, MachineInstr &I,
1995 SPIRV::StorageClass::StorageClass SC) {
1996 MachineIRBuilder MIRBuilder(I);
1997 return getOrCreateSPIRVPointerType(BaseType, MIRBuilder, SC);
1998}
1999
2000SPIRVTypeInst SPIRVGlobalRegistry::getOrCreateSPIRVPointerType(
2001 const Type *BaseType, MachineIRBuilder &MIRBuilder,
2002 SPIRV::StorageClass::StorageClass SC) {
2003 if (BaseType->isFunctionTy() &&
2004 !cast<SPIRVSubtarget>(Val: MIRBuilder.getMF().getSubtarget())
2005 .canUseExtension(E: SPIRV::Extension::SPV_INTEL_function_pointers)) {
2006 const Function &F = MIRBuilder.getMF().getFunction();
2007 F.getContext().diagnose(
2008 DI: DiagnosticInfoUnsupported(F,
2009 "Function used as a data pointer requires "
2010 "SPV_INTEL_function_pointers extension",
2011 DebugLoc(), DS_Error));
2012 }
2013 // TODO: Need to check if EmitIr should always be true.
2014 SPIRVTypeInst SpirvBaseType = getOrCreateSPIRVType(
2015 Ty: BaseType, MIRBuilder, AccessQual: SPIRV::AccessQualifier::ReadWrite,
2016 ExplicitLayoutRequired: storageClassRequiresExplictLayout(SC), EmitIR: true);
2017 assert(SpirvBaseType);
2018 return getOrCreateSPIRVPointerTypeInternal(BaseType: SpirvBaseType, MIRBuilder, SC);
2019}
2020
2021SPIRVTypeInst SPIRVGlobalRegistry::changePointerStorageClass(
2022 SPIRVTypeInst PtrType, SPIRV::StorageClass::StorageClass SC,
2023 MachineInstr &I) {
2024 [[maybe_unused]] SPIRV::StorageClass::StorageClass OldSC =
2025 getPointerStorageClass(Type: PtrType);
2026 assert(storageClassRequiresExplictLayout(OldSC) ==
2027 storageClassRequiresExplictLayout(SC));
2028
2029 SPIRVTypeInst PointeeType = getPointeeType(PtrType);
2030 MachineIRBuilder MIRBuilder(I);
2031 return getOrCreateSPIRVPointerTypeInternal(BaseType: PointeeType, MIRBuilder, SC);
2032}
2033
2034SPIRVTypeInst SPIRVGlobalRegistry::getOrCreateSPIRVPointerType(
2035 SPIRVTypeInst BaseType, MachineIRBuilder &MIRBuilder,
2036 SPIRV::StorageClass::StorageClass SC) {
2037 const Type *LLVMType = getTypeForSPIRVType(Ty: BaseType);
2038 assert(!storageClassRequiresExplictLayout(SC));
2039 SPIRVTypeInst R = getOrCreateSPIRVPointerType(BaseType: LLVMType, MIRBuilder, SC);
2040 assert(
2041 getPointeeType(R) == BaseType &&
2042 "The base type was not correctly laid out for the given storage class.");
2043 return R;
2044}
2045
2046SPIRVTypeInst SPIRVGlobalRegistry::getOrCreateSPIRVPointerTypeInternal(
2047 SPIRVTypeInst BaseType, MachineIRBuilder &MIRBuilder,
2048 SPIRV::StorageClass::StorageClass SC) {
2049 const Type *PointerElementType = getTypeForSPIRVType(Ty: BaseType);
2050 unsigned AddressSpace = storageClassToAddressSpace(SC);
2051 if (const MachineInstr *MI = findMI(PointeeTy: PointerElementType, AddressSpace, MF: CurMF))
2052 return MI;
2053 Type *Ty = TypedPointerType::get(ElementType: const_cast<Type *>(PointerElementType),
2054 AddressSpace);
2055 const MachineInstr *NewMI = createConstOrTypeAtFunctionEntry(
2056 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
2057 return BuildMI(BB&: MIRBuilder.getMBB(), I: MIRBuilder.getInsertPt(),
2058 MIMD: MIRBuilder.getDebugLoc(),
2059 MCID: MIRBuilder.getTII().get(Opcode: SPIRV::OpTypePointer))
2060 .addDef(RegNo: createTypeVReg(MRI&: CurMF->getRegInfo()))
2061 .addImm(Val: static_cast<uint32_t>(SC))
2062 .addUse(RegNo: getSPIRVTypeID(SpirvType: BaseType));
2063 });
2064 add(PointeeTy: PointerElementType, AddressSpace, MI: NewMI);
2065 return finishCreatingSPIRVType(LLVMTy: Ty, SpirvType: NewMI);
2066}
2067
2068Register SPIRVGlobalRegistry::getOrCreateUndef(MachineInstr &I,
2069 SPIRVTypeInst SpvType,
2070 const SPIRVInstrInfo &TII) {
2071 UndefValue *UV =
2072 UndefValue::get(T: const_cast<Type *>(getTypeForSPIRVType(Ty: SpvType)));
2073 Register Res = find(V: UV, MF: CurMF);
2074 if (Res.isValid())
2075 return Res;
2076
2077 LLT LLTy = LLT::scalar(SizeInBits: 64);
2078 Res = CurMF->getRegInfo().createGenericVirtualRegister(Ty: LLTy);
2079 CurMF->getRegInfo().setRegClass(Reg: Res, RC: &SPIRV::iIDRegClass);
2080 assignSPIRVTypeToVReg(SpirvType: SpvType, VReg: Res, MF: *CurMF);
2081
2082 MachineInstr *DepMI =
2083 const_cast<MachineInstr *>(static_cast<const MachineInstr *>(SpvType));
2084 MachineIRBuilder MIRBuilder(*DepMI->getParent(), DepMI->getIterator());
2085 const MachineInstr *NewMI = createConstOrTypeAtFunctionEntry(
2086 MIRBuilder, Op: [&](MachineIRBuilder &MIRBuilder) {
2087 auto MIB = BuildMI(BB&: MIRBuilder.getMBB(), I&: *MIRBuilder.getInsertPt(),
2088 MIMD: MIRBuilder.getDL(), MCID: TII.get(Opcode: SPIRV::OpUndef))
2089 .addDef(RegNo: Res)
2090 .addUse(RegNo: getSPIRVTypeID(SpirvType: SpvType));
2091 const auto &ST = CurMF->getSubtarget();
2092 constrainSelectedInstRegOperands(I&: *MIB, TII: *ST.getInstrInfo(),
2093 TRI: *ST.getRegisterInfo(),
2094 RBI: *ST.getRegBankInfo());
2095 return MIB;
2096 });
2097 add(V: UV, MI: NewMI);
2098 return Res;
2099}
2100
2101const TargetRegisterClass *
2102SPIRVGlobalRegistry::getRegClass(SPIRVTypeInst SpvType) const {
2103 unsigned Opcode = SpvType->getOpcode();
2104 switch (Opcode) {
2105 case SPIRV::OpTypeFloat:
2106 return &SPIRV::fIDRegClass;
2107 case SPIRV::OpTypePointer:
2108 return &SPIRV::pIDRegClass;
2109 case SPIRV::OpTypeVector: {
2110 SPIRVTypeInst ElemType = getScalarOrVectorComponentType(Type: SpvType);
2111 unsigned ElemOpcode = ElemType ? ElemType->getOpcode() : 0;
2112 if (ElemOpcode == SPIRV::OpTypeFloat)
2113 return &SPIRV::vfIDRegClass;
2114 if (ElemOpcode == SPIRV::OpTypePointer)
2115 return &SPIRV::vpIDRegClass;
2116 return &SPIRV::viIDRegClass;
2117 }
2118 }
2119 return &SPIRV::iIDRegClass;
2120}
2121
2122inline unsigned getAS(SPIRVTypeInst SpvType) {
2123 return storageClassToAddressSpace(
2124 SC: static_cast<SPIRV::StorageClass::StorageClass>(
2125 SpvType->getOperand(i: 1).getImm()));
2126}
2127
2128LLT SPIRVGlobalRegistry::getRegType(SPIRVTypeInst SpvType) const {
2129 unsigned Opcode = SpvType ? SpvType->getOpcode() : 0;
2130 switch (Opcode) {
2131 case SPIRV::OpTypeInt:
2132 case SPIRV::OpTypeFloat:
2133 case SPIRV::OpTypeBool:
2134 return LLT::scalar(SizeInBits: getScalarOrVectorBitWidth(Type: SpvType));
2135 case SPIRV::OpTypePointer:
2136 return LLT::pointer(AddressSpace: getAS(SpvType), SizeInBits: getPointerSize());
2137 case SPIRV::OpTypeVector: {
2138 SPIRVTypeInst ElemType = getScalarOrVectorComponentType(Type: SpvType);
2139 LLT ET;
2140 switch (ElemType ? ElemType->getOpcode() : 0) {
2141 case SPIRV::OpTypePointer:
2142 ET = LLT::pointer(AddressSpace: getAS(SpvType: ElemType), SizeInBits: getPointerSize());
2143 break;
2144 case SPIRV::OpTypeInt:
2145 case SPIRV::OpTypeFloat:
2146 case SPIRV::OpTypeBool:
2147 ET = LLT::scalar(SizeInBits: getScalarOrVectorBitWidth(Type: ElemType));
2148 break;
2149 default:
2150 ET = LLT::scalar(SizeInBits: 64);
2151 }
2152 return LLT::fixed_vector(NumElements: getScalarOrVectorComponentCount(Type: SpvType), ScalarTy: ET);
2153 }
2154 }
2155 return LLT::scalar(SizeInBits: 64);
2156}
2157
2158// Aliasing list MD contains several scope MD nodes whithin it. Each scope MD
2159// has a selfreference and an extra MD node for aliasing domain and also it
2160// can contain an optional string operand. Domain MD contains a self-reference
2161// with an optional string operand. Here we unfold the list, creating SPIR-V
2162// aliasing instructions.
2163// TODO: add support for an optional string operand.
2164MachineInstr *SPIRVGlobalRegistry::getOrAddMemAliasingINTELInst(
2165 MachineIRBuilder &MIRBuilder, const MDNode *AliasingListMD) {
2166 if (AliasingListMD->getNumOperands() == 0)
2167 return nullptr;
2168 if (auto L = AliasInstMDMap.find(Val: AliasingListMD); L != AliasInstMDMap.end())
2169 return L->second;
2170
2171 SmallVector<MachineInstr *> ScopeList;
2172 MachineRegisterInfo *MRI = MIRBuilder.getMRI();
2173 for (const MDOperand &MDListOp : AliasingListMD->operands()) {
2174 if (MDNode *ScopeMD = dyn_cast<MDNode>(Val: MDListOp)) {
2175 if (ScopeMD->getNumOperands() < 2)
2176 return nullptr;
2177 MDNode *DomainMD = dyn_cast<MDNode>(Val: ScopeMD->getOperand(I: 1));
2178 if (!DomainMD)
2179 return nullptr;
2180 auto *Domain = [&] {
2181 auto D = AliasInstMDMap.find(Val: DomainMD);
2182 if (D != AliasInstMDMap.end())
2183 return D->second;
2184 const Register Ret = MRI->createVirtualRegister(RegClass: &SPIRV::IDRegClass);
2185 auto MIB =
2186 MIRBuilder.buildInstr(Opcode: SPIRV::OpAliasDomainDeclINTEL).addDef(RegNo: Ret);
2187 return MIB.getInstr();
2188 }();
2189 AliasInstMDMap.insert(KV: std::make_pair(x&: DomainMD, y&: Domain));
2190 auto *Scope = [&] {
2191 auto S = AliasInstMDMap.find(Val: ScopeMD);
2192 if (S != AliasInstMDMap.end())
2193 return S->second;
2194 const Register Ret = MRI->createVirtualRegister(RegClass: &SPIRV::IDRegClass);
2195 auto MIB = MIRBuilder.buildInstr(Opcode: SPIRV::OpAliasScopeDeclINTEL)
2196 .addDef(RegNo: Ret)
2197 .addUse(RegNo: Domain->getOperand(i: 0).getReg());
2198 return MIB.getInstr();
2199 }();
2200 AliasInstMDMap.insert(KV: std::make_pair(x&: ScopeMD, y&: Scope));
2201 ScopeList.push_back(Elt: Scope);
2202 }
2203 }
2204
2205 const Register Ret = MRI->createVirtualRegister(RegClass: &SPIRV::IDRegClass);
2206 auto MIB =
2207 MIRBuilder.buildInstr(Opcode: SPIRV::OpAliasScopeListDeclINTEL).addDef(RegNo: Ret);
2208 for (auto *Scope : ScopeList)
2209 MIB.addUse(RegNo: Scope->getOperand(i: 0).getReg());
2210 auto List = MIB.getInstr();
2211 AliasInstMDMap.insert(KV: std::make_pair(x&: AliasingListMD, y&: List));
2212 return List;
2213}
2214
2215void SPIRVGlobalRegistry::buildMemAliasingOpDecorate(
2216 Register Reg, MachineIRBuilder &MIRBuilder, uint32_t Dec,
2217 const MDNode *AliasingListMD) {
2218 MachineInstr *AliasList =
2219 getOrAddMemAliasingINTELInst(MIRBuilder, AliasingListMD);
2220 if (!AliasList)
2221 return;
2222 MIRBuilder.buildInstr(Opcode: SPIRV::OpDecorateId)
2223 .addUse(RegNo: Reg)
2224 .addImm(Val: Dec)
2225 .addUse(RegNo: AliasList->getOperand(i: 0).getReg());
2226}
2227void SPIRVGlobalRegistry::replaceAllUsesWith(Value *Old, Value *New,
2228 bool DeleteOld) {
2229 Old->replaceAllUsesWith(V: New);
2230 updateIfExistDeducedElementType(OldVal: Old, NewVal: New, DeleteOld);
2231 updateIfExistAssignPtrTypeInstr(OldVal: Old, NewVal: New, DeleteOld);
2232}
2233
2234void SPIRVGlobalRegistry::buildAssignType(IRBuilder<> &B, Type *Ty,
2235 Value *Arg) {
2236 Value *OfType = getNormalizedPoisonValue(Ty);
2237 CallInst *AssignCI = nullptr;
2238 if (Arg->getType()->isAggregateType() && Ty->isAggregateType() &&
2239 allowEmitFakeUse(Arg)) {
2240 LLVMContext &Ctx = Arg->getContext();
2241 SmallVector<Metadata *, 2> ArgMDs{
2242 MDNode::get(Context&: Ctx, MDs: ValueAsMetadata::getConstant(C: OfType)),
2243 MDString::get(Context&: Ctx, Str: Arg->getName())};
2244 B.CreateIntrinsic(ID: Intrinsic::spv_value_md,
2245 Args: {MetadataAsValue::get(Context&: Ctx, MD: MDTuple::get(Context&: Ctx, MDs: ArgMDs))});
2246 AssignCI = B.CreateIntrinsicWithoutFolding(ID: Intrinsic::fake_use, Args: {Arg});
2247 } else {
2248 AssignCI = buildIntrWithMD(IntrID: Intrinsic::spv_assign_type, Types: {Arg->getType()},
2249 Arg: OfType, Arg2: Arg, Imms: {}, B);
2250 }
2251 addAssignPtrTypeInstr(Val: Arg, AssignPtrTyCI: AssignCI);
2252}
2253
2254void SPIRVGlobalRegistry::buildAssignPtr(IRBuilder<> &B, Type *ElemTy,
2255 Value *Arg) {
2256 Value *OfType = PoisonValue::get(T: ElemTy);
2257 CallInst *AssignPtrTyCI = findAssignPtrTypeInstr(Val: Arg);
2258 Function *CurrF =
2259 B.GetInsertBlock() ? B.GetInsertBlock()->getParent() : nullptr;
2260 if (AssignPtrTyCI == nullptr ||
2261 AssignPtrTyCI->getParent()->getParent() != CurrF) {
2262 AssignPtrTyCI = buildIntrWithMD(
2263 IntrID: Intrinsic::spv_assign_ptr_type, Types: {Arg->getType()}, Arg: OfType, Arg2: Arg,
2264 Imms: {B.getInt32(C: getPointerAddressSpace(T: Arg->getType()))}, B);
2265 addDeducedElementType(Val: AssignPtrTyCI, Ty: ElemTy);
2266 addDeducedElementType(Val: Arg, Ty: ElemTy);
2267 addAssignPtrTypeInstr(Val: Arg, AssignPtrTyCI);
2268 } else {
2269 updateAssignType(AssignCI: AssignPtrTyCI, Arg, OfType);
2270 }
2271}
2272
2273void SPIRVGlobalRegistry::updateAssignType(CallInst *AssignCI, Value *Arg,
2274 Value *OfType) {
2275 AssignCI->setArgOperand(i: 1, v: buildMD(Arg: OfType));
2276 if (cast<IntrinsicInst>(Val: AssignCI)->getIntrinsicID() !=
2277 Intrinsic::spv_assign_ptr_type)
2278 return;
2279
2280 // update association with the pointee type
2281 Type *ElemTy = OfType->getType();
2282 addDeducedElementType(Val: AssignCI, Ty: ElemTy);
2283 addDeducedElementType(Val: Arg, Ty: ElemTy);
2284}
2285
2286void SPIRVGlobalRegistry::addStructOffsetDecorations(
2287 Register Reg, StructType *Ty, MachineIRBuilder &MIRBuilder) {
2288 ArrayRef<TypeSize> Offsets = DL.getStructLayout(Ty)->getMemberOffsets();
2289 for (uint32_t I = 0; I < Ty->getNumElements(); ++I) {
2290 buildOpMemberDecorate(Reg, MIRBuilder, Dec: SPIRV::Decoration::Offset, Member: I,
2291 DecArgs: {static_cast<uint32_t>(Offsets[I])});
2292 }
2293}
2294
2295void SPIRVGlobalRegistry::addArrayStrideDecorations(
2296 Register Reg, Type *ElementType, MachineIRBuilder &MIRBuilder) {
2297 uint32_t SizeInBytes = DL.getTypeAllocSize(Ty: ElementType);
2298 buildOpDecorate(Reg, MIRBuilder, Dec: SPIRV::Decoration::ArrayStride,
2299 DecArgs: {SizeInBytes});
2300}
2301
2302bool SPIRVGlobalRegistry::hasBlockDecoration(SPIRVTypeInst Type) const {
2303 Register Def = getSPIRVTypeID(SpirvType: Type);
2304 for (const MachineInstr &Use :
2305 Type->getMF()->getRegInfo().use_instructions(Reg: Def)) {
2306 if (Use.getOpcode() != SPIRV::OpDecorate)
2307 continue;
2308
2309 if (Use.getOperand(i: 1).getImm() == SPIRV::Decoration::Block)
2310 return true;
2311 }
2312 return false;
2313}
2314