1//===- SPIRVBuiltins.cpp - SPIR-V Built-in Functions ------------*- 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 implements lowering builtin function calls and types using their
10// demangled names and TableGen records.
11//
12//===----------------------------------------------------------------------===//
13
14#include "SPIRVBuiltins.h"
15#include "SPIRV.h"
16#include "SPIRVSubtarget.h"
17#include "SPIRVUtils.h"
18#include "llvm/ADT/StringExtras.h"
19#include "llvm/ADT/StringTable.h"
20#include "llvm/Analysis/ValueTracking.h"
21#include "llvm/IR/DiagnosticInfo.h"
22#include "llvm/IR/IntrinsicsSPIRV.h"
23#include <regex>
24#include <string>
25#include <tuple>
26
27#define DEBUG_TYPE "spirv-builtins"
28
29namespace llvm {
30namespace SPIRV {
31#define GET_BuiltinGroup_DECL
32#include "SPIRVGenTables.inc"
33
34struct DemangledBuiltin {
35 StringTable::Offset Name;
36 InstructionSet::InstructionSet Set;
37 BuiltinGroup Group;
38 uint8_t MinNumArgs;
39 uint8_t MaxNumArgs;
40
41 StringRef name() const;
42};
43
44#define GET_DemangledBuiltins_DECL
45#define GET_DemangledBuiltins_IMPL
46
47struct IncomingCall {
48 const std::string BuiltinName;
49 const DemangledBuiltin *Builtin;
50
51 const Register ReturnRegister;
52 const SPIRVTypeInst ReturnType;
53 const SmallVectorImpl<Register> &Arguments;
54
55 IncomingCall(const std::string BuiltinName, const DemangledBuiltin *Builtin,
56 const Register ReturnRegister, SPIRVTypeInst ReturnType,
57 const SmallVectorImpl<Register> &Arguments)
58 : BuiltinName(std::move(BuiltinName)), Builtin(Builtin),
59 ReturnRegister(ReturnRegister), ReturnType(ReturnType),
60 Arguments(Arguments) {}
61
62 bool isSpirvOp() const { return BuiltinName.rfind(s: "__spirv_", pos: 0) == 0; }
63};
64
65struct NativeBuiltin {
66 StringTable::Offset Name;
67 InstructionSet::InstructionSet Set;
68 uint32_t Opcode;
69};
70
71#define GET_NativeBuiltins_DECL
72#define GET_NativeBuiltins_IMPL
73
74struct GroupBuiltin {
75 StringTable::Offset Name;
76 uint32_t Opcode;
77 uint32_t GroupOperation;
78 bool IsElect;
79 bool IsAllOrAny;
80 bool IsAllEqual;
81 bool IsInverseBallot;
82 bool IsBallotBitExtract;
83 bool IsLogical;
84 bool NoGroupOperation;
85 bool HasBoolArg;
86};
87
88#define GET_GroupBuiltins_DECL
89#define GET_GroupBuiltins_IMPL
90
91struct IntelSubgroupsBuiltin {
92 StringTable::Offset Name;
93 uint32_t Opcode;
94 bool IsBlock;
95 bool IsWrite;
96 bool IsMedia;
97};
98
99#define GET_IntelSubgroupsBuiltins_DECL
100#define GET_IntelSubgroupsBuiltins_IMPL
101
102struct AtomicFloatingBuiltin {
103 StringTable::Offset Name;
104 uint32_t Opcode;
105};
106
107#define GET_AtomicFloatingBuiltins_DECL
108#define GET_AtomicFloatingBuiltins_IMPL
109struct GroupUniformBuiltin {
110 StringTable::Offset Name;
111 uint32_t Opcode;
112 bool IsLogical;
113};
114
115#define GET_GroupUniformBuiltins_DECL
116#define GET_GroupUniformBuiltins_IMPL
117
118struct GetBuiltin {
119 StringTable::Offset Name;
120 InstructionSet::InstructionSet Set;
121 BuiltIn::BuiltIn Value;
122};
123
124using namespace BuiltIn;
125#define GET_GetBuiltins_DECL
126#define GET_GetBuiltins_IMPL
127
128struct ImageQueryBuiltin {
129 StringTable::Offset Name;
130 InstructionSet::InstructionSet Set;
131 uint32_t Component;
132};
133
134#define GET_ImageQueryBuiltins_DECL
135#define GET_ImageQueryBuiltins_IMPL
136
137struct IntegerDotProductBuiltin {
138 StringTable::Offset Name;
139 uint32_t Opcode;
140 bool IsSwapReq;
141};
142
143#define GET_IntegerDotProductBuiltins_DECL
144#define GET_IntegerDotProductBuiltins_IMPL
145
146struct ConvertBuiltin {
147 StringTable::Offset Name;
148 InstructionSet::InstructionSet Set;
149 bool IsDestinationSigned;
150 bool IsSaturated;
151 bool IsRounded;
152 bool IsBfloat16;
153 bool IsTF32;
154 FPRoundingMode::FPRoundingMode RoundingMode;
155};
156
157struct VectorLoadStoreBuiltin {
158 StringTable::Offset Name;
159 InstructionSet::InstructionSet Set;
160 uint32_t Number;
161 uint32_t ElementCount;
162 bool IsRounded;
163 FPRoundingMode::FPRoundingMode RoundingMode;
164};
165
166using namespace FPRoundingMode;
167#define GET_ConvertBuiltins_DECL
168#define GET_ConvertBuiltins_IMPL
169
170using namespace InstructionSet;
171#define GET_VectorLoadStoreBuiltins_DECL
172#define GET_VectorLoadStoreBuiltins_IMPL
173
174#define GET_CLMemoryScope_DECL
175#define GET_CLSamplerAddressingMode_DECL
176#define GET_CLMemoryFenceFlags_DECL
177#define GET_ExtendedBuiltins_DECL
178#include "SPIRVGenTables.inc"
179
180// Defined here to reference declarations from tablegen.
181StringRef DemangledBuiltin::name() const {
182 return getDemangledBuiltinStr(Offset: Name);
183}
184} // namespace SPIRV
185
186//===----------------------------------------------------------------------===//
187// Misc functions for looking up builtins and veryfying requirements using
188// TableGen records
189//===----------------------------------------------------------------------===//
190
191namespace SPIRV {
192/// Parses the name part of the demangled builtin call.
193std::string lookupBuiltinNameHelper(StringRef DemangledCall,
194 FPDecorationId *DecorationId) {
195 StringRef PassPrefix = "(anonymous namespace)::";
196 StringRef SpvPrefix = "__spv::";
197 std::string BuiltinName = DemangledCall.str();
198
199 // Check if the extracted name contains type information between angle
200 // brackets. If so, the builtin is an instantiated template - needs to have
201 // the information after angle brackets and return type removed.
202 std::size_t Pos = BuiltinName.find(s: ">(");
203 if (Pos != std::string::npos) {
204 BuiltinName = BuiltinName.substr(pos: 0, n: BuiltinName.rfind(c: '<', pos: Pos));
205 } else {
206 Pos = BuiltinName.find(c: '(');
207 if (Pos != std::string::npos)
208 BuiltinName = BuiltinName.substr(pos: 0, n: Pos);
209 }
210 BuiltinName = BuiltinName.substr(pos: BuiltinName.find_last_of(c: ' ') + 1);
211
212 // Itanium Demangler result may have "(anonymous namespace)::" or "__spv::"
213 // prefix.
214 if (BuiltinName.find(svt: PassPrefix) == 0)
215 BuiltinName = BuiltinName.substr(pos: PassPrefix.size());
216 else if (BuiltinName.find(svt: SpvPrefix) == 0)
217 BuiltinName = BuiltinName.substr(pos: SpvPrefix.size());
218
219 // Account for possible "__spirv_ocl_" prefix in SPIR-V friendly LLVM IR
220 if (BuiltinName.rfind(s: "__spirv_ocl_", pos: 0) == 0)
221 BuiltinName = BuiltinName.substr(pos: 12);
222
223 // Check if the extracted name begins with:
224 // - "__spirv_ImageSampleExplicitLod"
225 // - "__spirv_ImageRead"
226 // - "__spirv_ImageWrite"
227 // - "__spirv_ImageQuerySizeLod"
228 // - "__spirv_UDotKHR"
229 // - "__spirv_SDotKHR"
230 // - "__spirv_SUDotKHR"
231 // - "__spirv_SDotAccSatKHR"
232 // - "__spirv_UDotAccSatKHR"
233 // - "__spirv_SUDotAccSatKHR"
234 // - "__spirv_ReadClockKHR"
235 // - "__spirv_SubgroupBlockReadINTEL"
236 // - "__spirv_SubgroupImageBlockReadINTEL"
237 // - "__spirv_SubgroupImageMediaBlockReadINTEL"
238 // - "__spirv_SubgroupImageMediaBlockWriteINTEL"
239 // - "__spirv_Convert"
240 // - "__spirv_Round"
241 // - "__spirv_UConvert"
242 // - "__spirv_SConvert"
243 // - "__spirv_FConvert"
244 // - "__spirv_SatConvert"
245 // and maybe contains return type information at the end "_R<type>".
246 // If so, extract the plain builtin name without the type information.
247 static const std::regex SpvWithR(
248 "(__spirv_(ImageSampleExplicitLod|ImageRead|ImageWrite|ImageQuerySizeLod|"
249 "UDotKHR|"
250 "SDotKHR|SUDotKHR|SDotAccSatKHR|UDotAccSatKHR|SUDotAccSatKHR|"
251 "ReadClockKHR|SubgroupBlockReadINTEL|SubgroupImageBlockReadINTEL|"
252 "SubgroupImageMediaBlockReadINTEL|SubgroupImageMediaBlockWriteINTEL|"
253 "Convert|Round|"
254 "UConvert|SConvert|FConvert|SatConvert)[^_]*)(_R[^_]*_?(\\w+)?.*)?");
255 std::smatch Match;
256 if (std::regex_match(s: BuiltinName, m&: Match, re: SpvWithR) && Match.size() > 1) {
257 std::ssub_match SubMatch;
258 if (DecorationId && Match.size() > 3) {
259 SubMatch = Match[4];
260 *DecorationId = demangledPostfixToDecorationId(S: SubMatch.str());
261 }
262 SubMatch = Match[1];
263 BuiltinName = SubMatch.str();
264 }
265
266 return BuiltinName;
267}
268} // namespace SPIRV
269
270/// Looks up the demangled builtin call in the SPIRVBuiltins.td records using
271/// the provided \p DemangledCall and specified \p Set.
272///
273/// The lookup follows the following algorithm, returning the first successful
274/// match:
275/// 1. Search with the plain demangled name (expecting a 1:1 match).
276/// 2. Search with the prefix before or suffix after the demangled name
277/// signyfying the type of the first argument.
278///
279/// \returns Wrapper around the demangled call and found builtin definition.
280static std::unique_ptr<const SPIRV::IncomingCall>
281lookupBuiltin(StringRef DemangledCall,
282 SPIRV::InstructionSet::InstructionSet Set,
283 Register ReturnRegister, SPIRVTypeInst ReturnType,
284 const SmallVectorImpl<Register> &Arguments) {
285 std::string BuiltinName = SPIRV::lookupBuiltinNameHelper(DemangledCall);
286
287 SmallVector<StringRef, 10> BuiltinArgumentTypes;
288 StringRef BuiltinArgs =
289 DemangledCall.slice(Start: DemangledCall.find(C: '(') + 1, End: DemangledCall.find(C: ')'));
290 BuiltinArgs.split(A&: BuiltinArgumentTypes, Separator: ',', MaxSplit: -1, KeepEmpty: false);
291
292 // Look up the builtin in the defined set. Start with the plain demangled
293 // name, expecting a 1:1 match in the defined builtin set.
294 const SPIRV::DemangledBuiltin *Builtin;
295 if ((Builtin = SPIRV::lookupBuiltin(Name: BuiltinName, Set)))
296 return std::make_unique<SPIRV::IncomingCall>(
297 args&: BuiltinName, args&: Builtin, args&: ReturnRegister, args&: ReturnType, args: Arguments);
298
299 // If the initial look up was unsuccessful and the demangled call takes at
300 // least 1 argument, add a prefix or suffix signifying the type of the first
301 // argument and repeat the search.
302 if (BuiltinArgumentTypes.size() >= 1) {
303 char FirstArgumentType = BuiltinArgumentTypes[0][0];
304 // Prefix and suffix to be added to the builtin's name for lookup.
305 // For example, OpenCL "abs" taking an unsigned value has a prefix "u_",
306 // and "group_reduce_max" taking an unsigned value has a suffix "u".
307 StringRef Prefix;
308 StringRef Suffix;
309
310 switch (FirstArgumentType) {
311 // Unsigned:
312 case 'u':
313 if (Set == SPIRV::InstructionSet::OpenCL_std)
314 Prefix = "u_";
315 else if (Set == SPIRV::InstructionSet::GLSL_std_450)
316 Prefix = "u";
317 Suffix = "u";
318 break;
319 // Signed:
320 case 'c':
321 case 's':
322 case 'i':
323 case 'l':
324 if (Set == SPIRV::InstructionSet::OpenCL_std)
325 Prefix = "s_";
326 else if (Set == SPIRV::InstructionSet::GLSL_std_450)
327 Prefix = "s";
328 Suffix = "s";
329 break;
330 // Floating-point:
331 case 'f':
332 case 'd':
333 case 'h':
334 if (Set == SPIRV::InstructionSet::OpenCL_std ||
335 Set == SPIRV::InstructionSet::GLSL_std_450)
336 Prefix = "f";
337 Suffix = "f";
338 break;
339 }
340
341 // If argument-type name prefix was added, look up the builtin again.
342 if (!Prefix.empty() &&
343 (Builtin = SPIRV::lookupBuiltin(Name: (Prefix + BuiltinName).str(), Set)))
344 return std::make_unique<SPIRV::IncomingCall>(
345 args&: BuiltinName, args&: Builtin, args&: ReturnRegister, args&: ReturnType, args: Arguments);
346
347 if (!Suffix.empty() &&
348 (Builtin = SPIRV::lookupBuiltin(Name: (BuiltinName + Suffix).str(), Set)))
349 return std::make_unique<SPIRV::IncomingCall>(
350 args&: BuiltinName, args&: Builtin, args&: ReturnRegister, args&: ReturnType, args: Arguments);
351 }
352
353 // No builtin with such name was found in the set.
354 return nullptr;
355}
356
357static MachineInstr *getBlockStructInstr(Register ParamReg,
358 MachineRegisterInfo *MRI) {
359 // We expect ParamReg to be defined by G_ADDRSPACE_CAST with a source from
360 // G_GLOBAL_VALUE or spv_alloca. Returns the source instruction.
361 MachineInstr *MI = MRI->getUniqueVRegDef(Reg: ParamReg);
362 assert(MI->getOpcode() == TargetOpcode::G_ADDRSPACE_CAST &&
363 MI->getOperand(1).isReg());
364 Register BitcastReg = MI->getOperand(i: 1).getReg();
365 MachineInstr *BitcastMI = MRI->getUniqueVRegDef(Reg: BitcastReg);
366 assert(BitcastMI && "Definition for source reg not found.");
367 if (BitcastMI->getOpcode() == TargetOpcode::G_GLOBAL_VALUE ||
368 isSpvIntrinsic(MI: *BitcastMI, IntrinsicID: Intrinsic::spv_alloca))
369 return BitcastMI;
370 llvm_unreachable("getBlockStructInstr: unexpected instruction pattern");
371}
372
373// Return type of the instruction result from spv_assign_type intrinsic.
374// TODO: maybe unify with prelegalizer pass.
375static const Type *getMachineInstrType(MachineInstr *MI) {
376 MachineInstr *NextMI = MI->getNextNode();
377 if (!NextMI)
378 return nullptr;
379 if (isSpvIntrinsic(MI: *NextMI, IntrinsicID: Intrinsic::spv_assign_name))
380 if ((NextMI = NextMI->getNextNode()) == nullptr)
381 return nullptr;
382 Register ValueReg = MI->getOperand(i: 0).getReg();
383 if ((!isSpvIntrinsic(MI: *NextMI, IntrinsicID: Intrinsic::spv_assign_type) &&
384 !isSpvIntrinsic(MI: *NextMI, IntrinsicID: Intrinsic::spv_assign_ptr_type)) ||
385 NextMI->getOperand(i: 1).getReg() != ValueReg)
386 return nullptr;
387 Type *Ty = getMDOperandAsType(N: NextMI->getOperand(i: 2).getMetadata(), I: 0);
388 assert(Ty && "Type is expected");
389 return Ty;
390}
391
392static const Type *getBlockStructType(Register ParamReg,
393 MachineRegisterInfo *MRI) {
394 // In principle, this information should be passed to us from Clang via
395 // an elementtype attribute. However, said attribute requires that
396 // the function call be an intrinsic, which is not. Instead, we rely on being
397 // able to trace this to the declaration of a variable: OpenCL C specification
398 // section 6.12.5 should guarantee that we can do this.
399 MachineInstr *MI = getBlockStructInstr(ParamReg, MRI);
400 if (MI->getOpcode() == TargetOpcode::G_GLOBAL_VALUE)
401 return MI->getOperand(i: 1).getGlobal()->getValueType();
402 assert(isSpvIntrinsic(*MI, Intrinsic::spv_alloca) &&
403 "Blocks in OpenCL C must be traceable to allocation site");
404 return getMachineInstrType(MI);
405}
406
407//===----------------------------------------------------------------------===//
408// Helper functions for building misc instructions
409//===----------------------------------------------------------------------===//
410
411/// Helper function building either a resulting scalar or vector bool register
412/// depending on the expected \p ResultType.
413///
414/// \returns Tuple of the resulting register and its type.
415static std::tuple<Register, SPIRVTypeInst>
416buildBoolRegister(MachineIRBuilder &MIRBuilder, SPIRVTypeInst ResultType,
417 SPIRVGlobalRegistry *GR) {
418 LLT Type;
419 SPIRVTypeInst BoolType = GR->getOrCreateSPIRVBoolType(MIRBuilder, EmitIR: true);
420
421 if (isVectorType(SPVTy: ResultType)) {
422 unsigned VectorElements = GR->getScalarOrVectorComponentCount(Type: ResultType);
423 BoolType = GR->getOrCreateSPIRVVectorType(BaseType: BoolType, NumElements: VectorElements,
424 MIRBuilder, EmitIR: true);
425 const FixedVectorType *LLVMVectorType =
426 cast<FixedVectorType>(Val: GR->getTypeForSPIRVType(Ty: BoolType));
427 Type = LLT::vector(EC: LLVMVectorType->getElementCount(), ScalarSizeInBits: 1);
428 } else {
429 Type = LLT::scalar(SizeInBits: 1);
430 }
431
432 Register ResultRegister =
433 MIRBuilder.getMRI()->createGenericVirtualRegister(Ty: Type);
434 MIRBuilder.getMRI()->setRegClass(Reg: ResultRegister, RC: GR->getRegClass(SpvType: ResultType));
435 GR->assignSPIRVTypeToVReg(Type: BoolType, VReg: ResultRegister, MF: MIRBuilder.getMF());
436 return std::make_tuple(args&: ResultRegister, args&: BoolType);
437}
438
439/// Helper function for building either a vector or scalar select instruction
440/// depending on the expected \p ResultType.
441static bool buildSelectInst(MachineIRBuilder &MIRBuilder,
442 Register ReturnRegister, Register SourceRegister,
443 SPIRVTypeInst ReturnType, SPIRVGlobalRegistry *GR) {
444 Register TrueConst, FalseConst;
445
446 if (isVectorType(SPVTy: ReturnType)) {
447 unsigned Bits = GR->getScalarOrVectorBitWidth(Type: ReturnType);
448 uint64_t AllOnes = APInt::getAllOnes(numBits: Bits).getZExtValue();
449 TrueConst =
450 GR->getOrCreateConsIntVector(Val: AllOnes, MIRBuilder, SpvType: ReturnType, EmitIR: true);
451 FalseConst = GR->getOrCreateConsIntVector(Val: 0, MIRBuilder, SpvType: ReturnType, EmitIR: true);
452 } else {
453 TrueConst = GR->buildConstantInt(Val: 1, MIRBuilder, SpvType: ReturnType, EmitIR: true);
454 FalseConst = GR->buildConstantInt(Val: 0, MIRBuilder, SpvType: ReturnType, EmitIR: true);
455 }
456
457 return MIRBuilder.buildSelect(Res: ReturnRegister, Tst: SourceRegister, Op0: TrueConst,
458 Op1: FalseConst);
459}
460
461/// Helper function for building a load instruction loading into the
462/// \p DestinationReg.
463static Register buildLoadInst(SPIRVTypeInst BaseType, Register PtrRegister,
464 MachineIRBuilder &MIRBuilder,
465 SPIRVGlobalRegistry *GR,
466 Register DestinationReg = Register(0)) {
467 if (!DestinationReg.isValid())
468 DestinationReg = createVirtualRegister(SpvType: BaseType, GR, MIRBuilder);
469 // TODO: consider using correct address space and alignment (p0 is canonical
470 // type for selection though).
471 MachinePointerInfo PtrInfo = MachinePointerInfo();
472 MIRBuilder.buildLoad(Res: DestinationReg, Addr: PtrRegister, PtrInfo, Alignment: Align());
473 return DestinationReg;
474}
475
476/// Helper function for building a load instruction for loading a builtin global
477/// variable of \p BuiltinValue value.
478static Register buildBuiltinVariableLoad(
479 MachineIRBuilder &MIRBuilder, SPIRVTypeInst VariableType,
480 SPIRVGlobalRegistry *GR, SPIRV::BuiltIn::BuiltIn BuiltinValue, LLT LLType,
481 Register Reg = Register(0), bool isConst = true,
482 const std::optional<SPIRV::LinkageType::LinkageType> &LinkageTy = {
483 SPIRV::LinkageType::Import}) {
484 Register NewRegister =
485 MIRBuilder.getMRI()->createVirtualRegister(RegClass: &SPIRV::pIDRegClass);
486 MIRBuilder.getMRI()->setType(
487 VReg: NewRegister,
488 Ty: LLT::pointer(AddressSpace: storageClassToAddressSpace(SC: SPIRV::StorageClass::Function),
489 SizeInBits: GR->getPointerSize()));
490 SPIRVTypeInst PtrType = GR->getOrCreateSPIRVPointerType(
491 BaseType: VariableType, MIRBuilder, SC: SPIRV::StorageClass::Input);
492 GR->assignSPIRVTypeToVReg(Type: PtrType, VReg: NewRegister, MF: MIRBuilder.getMF());
493
494 // Set up the global OpVariable with the necessary builtin decorations.
495 Register Variable = GR->buildGlobalVariable(
496 Reg: NewRegister, BaseType: PtrType, Name: getLinkStringForBuiltIn(BuiltInValue: BuiltinValue), GV: nullptr,
497 Storage: SPIRV::StorageClass::Input, Init: nullptr, /* isConst= */ IsConst: isConst, LinkageType: LinkageTy,
498 MIRBuilder, IsInstSelector: false);
499
500 // Load the value from the global variable.
501 Register LoadedRegister =
502 buildLoadInst(BaseType: VariableType, PtrRegister: Variable, MIRBuilder, GR, DestinationReg: Reg);
503 MIRBuilder.getMRI()->setType(VReg: LoadedRegister, Ty: LLType);
504 return LoadedRegister;
505}
506
507/// Helper external function for assigning a SPIRV type to a register, ensuring
508/// the register class and type are set in MRI. Defined in
509/// SPIRVPreLegalizer.cpp.
510extern void updateRegType(Register Reg, Type *Ty, SPIRVTypeInst SpirvTy,
511 SPIRVGlobalRegistry *GR, MachineIRBuilder &MIB,
512 MachineRegisterInfo &MRI);
513
514// TODO: Move to TableGen.
515static SPIRV::MemorySemantics::MemorySemantics
516getSPIRVMemSemantics(std::memory_order MemOrder) {
517 switch (MemOrder) {
518 case std::memory_order_relaxed:
519 return SPIRV::MemorySemantics::None;
520 case std::memory_order_acquire:
521 return SPIRV::MemorySemantics::Acquire;
522 case std::memory_order_release:
523 return SPIRV::MemorySemantics::Release;
524 case std::memory_order_acq_rel:
525 return SPIRV::MemorySemantics::AcquireRelease;
526 case std::memory_order_seq_cst:
527 return SPIRV::MemorySemantics::SequentiallyConsistent;
528 default:
529 report_fatal_error(reason: "Unknown CL memory order");
530 }
531}
532
533static SPIRV::Scope::Scope getSPIRVScope(SPIRV::CLMemoryScope ClScope) {
534 switch (ClScope) {
535 case SPIRV::CLMemoryScope::memory_scope_work_item:
536 return SPIRV::Scope::Invocation;
537 case SPIRV::CLMemoryScope::memory_scope_work_group:
538 return SPIRV::Scope::Workgroup;
539 case SPIRV::CLMemoryScope::memory_scope_device:
540 return SPIRV::Scope::Device;
541 case SPIRV::CLMemoryScope::memory_scope_all_svm_devices:
542 return SPIRV::Scope::CrossDevice;
543 case SPIRV::CLMemoryScope::memory_scope_sub_group:
544 return SPIRV::Scope::Subgroup;
545 }
546 report_fatal_error(reason: "Unknown CL memory scope");
547}
548
549static Register buildConstantIntReg32(uint64_t Val,
550 MachineIRBuilder &MIRBuilder,
551 SPIRVGlobalRegistry *GR) {
552 return GR->buildConstantInt(
553 Val, MIRBuilder, SpvType: GR->getOrCreateSPIRVIntegerType(BitWidth: 32, MIRBuilder), EmitIR: true);
554}
555
556static Register buildScopeReg(Register CLScopeRegister,
557 SPIRV::Scope::Scope Scope,
558 MachineIRBuilder &MIRBuilder,
559 SPIRVGlobalRegistry *GR,
560 MachineRegisterInfo *MRI) {
561 if (CLScopeRegister.isValid()) {
562 auto CLScope =
563 static_cast<SPIRV::CLMemoryScope>(getIConstVal(ConstReg: CLScopeRegister, MRI));
564 Scope = getSPIRVScope(ClScope: CLScope);
565
566 if (CLScope == static_cast<unsigned>(Scope)) {
567 MRI->setRegClass(Reg: CLScopeRegister, RC: &SPIRV::iIDRegClass);
568 return CLScopeRegister;
569 }
570 }
571 return buildConstantIntReg32(Val: Scope, MIRBuilder, GR);
572}
573
574static void setRegClassIfNull(Register Reg, MachineRegisterInfo *MRI,
575 SPIRVGlobalRegistry *GR) {
576 if (MRI->getRegClassOrNull(Reg))
577 return;
578 SPIRVTypeInst SpvType = GR->getSPIRVTypeForVReg(VReg: Reg);
579 MRI->setRegClass(Reg,
580 RC: SpvType ? GR->getRegClass(SpvType) : &SPIRV::iIDRegClass);
581}
582
583/// Translates an OpenCL memory_order argument into the memory ordering part of
584/// the SPIR-V memory semantics.
585static SPIRV::MemorySemantics::MemorySemantics
586getMemOrdering(Register OrderRegister, MachineRegisterInfo *MRI) {
587 return getSPIRVMemSemantics(
588 MemOrder: static_cast<std::memory_order>(getIConstVal(ConstReg: OrderRegister, MRI)));
589}
590
591/// Combines the memory ordering with the storage-class part of the memory
592/// semantics into a constant register.
593static Register
594buildMemSemanticsReg(SPIRV::MemorySemantics::MemorySemantics Ordering,
595 unsigned StorageClassSem, MachineIRBuilder &MIRBuilder,
596 SPIRVGlobalRegistry *GR) {
597 const auto *ST =
598 static_cast<const SPIRVSubtarget *>(&MIRBuilder.getMF().getSubtarget());
599 return buildConstantIntReg32(
600 Val: getMemSemanticsWithStorageClass(TT: ST->getTargetTriple(), OrderSem: Ordering,
601 StorageClassSem),
602 MIRBuilder, GR);
603}
604
605static bool buildOpFromWrapper(MachineIRBuilder &MIRBuilder, unsigned Opcode,
606 const SPIRV::IncomingCall *Call,
607 Register TypeReg,
608 ArrayRef<uint32_t> ImmArgs = {}) {
609 auto MIB = MIRBuilder.buildInstr(Opcode);
610 if (TypeReg.isValid())
611 MIB.addDef(RegNo: Call->ReturnRegister).addUse(RegNo: TypeReg);
612 unsigned Sz = Call->Arguments.size() - ImmArgs.size();
613 for (unsigned i = 0; i < Sz; ++i)
614 MIB.addUse(RegNo: Call->Arguments[i]);
615 for (uint32_t ImmArg : ImmArgs)
616 MIB.addImm(Val: ImmArg);
617 return true;
618}
619
620/// Helper function for translating atomic init to OpStore.
621static bool buildAtomicInitInst(const SPIRV::IncomingCall *Call,
622 MachineIRBuilder &MIRBuilder) {
623 if (Call->isSpirvOp())
624 return buildOpFromWrapper(MIRBuilder, Opcode: SPIRV::OpStore, Call, TypeReg: Register(0));
625
626 assert(Call->Arguments.size() == 2 &&
627 "Need 2 arguments for atomic init translation");
628 MIRBuilder.buildInstr(Opcode: SPIRV::OpStore)
629 .addUse(RegNo: Call->Arguments[0])
630 .addUse(RegNo: Call->Arguments[1]);
631 return true;
632}
633
634/// Helper function for building an atomic load instruction.
635static bool buildAtomicLoadInst(const SPIRV::IncomingCall *Call,
636 MachineIRBuilder &MIRBuilder,
637 SPIRVGlobalRegistry *GR) {
638 Register TypeReg = GR->getSPIRVTypeID(SpirvType: Call->ReturnType);
639 if (Call->isSpirvOp())
640 return buildOpFromWrapper(MIRBuilder, Opcode: SPIRV::OpAtomicLoad, Call, TypeReg);
641
642 // atomic_load_explicit(ptr, memory_order[, memory_scope]).
643 Register PtrRegister = Call->Arguments[0];
644 MachineRegisterInfo *MRI = MIRBuilder.getMRI();
645
646 const SPIRV::MemorySemantics::MemorySemantics Ordering =
647 Call->Arguments.size() >= 2
648 ? getMemOrdering(OrderRegister: Call->Arguments[1], MRI)
649 : SPIRV::MemorySemantics::SequentiallyConsistent;
650 const unsigned StorageClassSem =
651 getMemSemanticsForStorageClass(SC: GR->getPointerStorageClass(VReg: PtrRegister));
652 Register MemSemanticsReg =
653 buildMemSemanticsReg(Ordering, StorageClassSem, MIRBuilder, GR);
654
655 Register ScopeRegister = buildScopeReg(
656 CLScopeRegister: Call->Arguments.size() >= 3 ? Call->Arguments[2] : Register(),
657 Scope: SPIRV::Scope::Device, MIRBuilder, GR, MRI);
658
659 MIRBuilder.buildInstr(Opcode: SPIRV::OpAtomicLoad)
660 .addDef(RegNo: Call->ReturnRegister)
661 .addUse(RegNo: TypeReg)
662 .addUse(RegNo: PtrRegister)
663 .addUse(RegNo: ScopeRegister)
664 .addUse(RegNo: MemSemanticsReg);
665 return true;
666}
667
668/// Helper function for building an atomic store instruction.
669static bool buildAtomicStoreInst(const SPIRV::IncomingCall *Call,
670 MachineIRBuilder &MIRBuilder,
671 SPIRVGlobalRegistry *GR) {
672 if (Call->isSpirvOp())
673 return buildOpFromWrapper(MIRBuilder, Opcode: SPIRV::OpAtomicStore, Call,
674 TypeReg: Register(0));
675
676 MachineRegisterInfo *MRI = MIRBuilder.getMRI();
677 Register PtrRegister = Call->Arguments[0];
678 // atomic_store_explicit(ptr, value, memory_order[, memory_scope]).
679 const SPIRV::MemorySemantics::MemorySemantics Ordering =
680 Call->Arguments.size() >= 3
681 ? getMemOrdering(OrderRegister: Call->Arguments[2], MRI)
682 : SPIRV::MemorySemantics::SequentiallyConsistent;
683 const unsigned StorageClassSem =
684 getMemSemanticsForStorageClass(SC: GR->getPointerStorageClass(VReg: PtrRegister));
685 Register MemSemanticsReg =
686 buildMemSemanticsReg(Ordering, StorageClassSem, MIRBuilder, GR);
687 Register ScopeRegister = buildScopeReg(
688 CLScopeRegister: Call->Arguments.size() >= 4 ? Call->Arguments[3] : Register(),
689 Scope: SPIRV::Scope::Device, MIRBuilder, GR, MRI);
690 MIRBuilder.buildInstr(Opcode: SPIRV::OpAtomicStore)
691 .addUse(RegNo: PtrRegister)
692 .addUse(RegNo: ScopeRegister)
693 .addUse(RegNo: MemSemanticsReg)
694 .addUse(RegNo: Call->Arguments[1]);
695 return true;
696}
697
698/// Helper function for building an atomic compare-exchange instruction.
699static bool buildAtomicCompareExchangeInst(
700 const SPIRV::IncomingCall *Call, const SPIRV::DemangledBuiltin *Builtin,
701 unsigned Opcode, MachineIRBuilder &MIRBuilder, SPIRVGlobalRegistry *GR) {
702 if (Call->isSpirvOp())
703 return buildOpFromWrapper(MIRBuilder, Opcode, Call,
704 TypeReg: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
705
706 bool IsCmpxchg = Call->Builtin->name().contains(Other: "cmpxchg");
707 MachineRegisterInfo *MRI = MIRBuilder.getMRI();
708
709 Register ObjectPtr = Call->Arguments[0]; // Pointer (volatile A *object.)
710 Register ExpectedArg = Call->Arguments[1]; // Comparator (C* expected).
711 Register Desired = Call->Arguments[2]; // Value (C Desired).
712 SPIRVTypeInst SpvDesiredTy = GR->getSPIRVTypeForVReg(VReg: Desired);
713 LLT DesiredLLT = MRI->getType(Reg: Desired);
714
715 assert(GR->getSPIRVTypeForVReg(ObjectPtr).isPointer());
716 [[maybe_unused]] SPIRVTypeInst ExpectedTy =
717 GR->getSPIRVTypeForVReg(VReg: ExpectedArg);
718 assert(IsCmpxchg ? ExpectedTy->getOpcode() == SPIRV::OpTypeInt
719 : ExpectedTy.isPointer());
720 assert(GR->isScalarOfType(Desired, SPIRV::OpTypeInt));
721
722 SPIRVTypeInst SpvObjectPtrTy = GR->getSPIRVTypeForVReg(VReg: ObjectPtr);
723 assert((SpvObjectPtrTy->getOpcode() == SPIRV::OpTypeUntypedPointerKHR ||
724 SpvObjectPtrTy->getOperand(2).isReg()) &&
725 "SPIRV type is expected");
726 auto StorageClass = static_cast<SPIRV::StorageClass::StorageClass>(
727 SpvObjectPtrTy->getOperand(i: 1).getImm());
728 auto MemSemStorage = getMemSemanticsForStorageClass(SC: StorageClass);
729
730 Register MemSemEqualReg;
731 Register MemSemUnequalReg;
732 uint64_t MemSemEqual =
733 IsCmpxchg
734 ? SPIRV::MemorySemantics::None
735 : SPIRV::MemorySemantics::SequentiallyConsistent | MemSemStorage;
736 uint64_t MemSemUnequal =
737 IsCmpxchg
738 ? SPIRV::MemorySemantics::None
739 : SPIRV::MemorySemantics::SequentiallyConsistent | MemSemStorage;
740 if (Call->Arguments.size() >= 4) {
741 assert(Call->Arguments.size() >= 5 &&
742 "Need 5+ args for explicit atomic cmpxchg");
743 auto MemOrdEq =
744 static_cast<std::memory_order>(getIConstVal(ConstReg: Call->Arguments[3], MRI));
745 auto MemOrdNeq =
746 static_cast<std::memory_order>(getIConstVal(ConstReg: Call->Arguments[4], MRI));
747 MemSemEqual = getSPIRVMemSemantics(MemOrder: MemOrdEq) | MemSemStorage;
748 MemSemUnequal = getSPIRVMemSemantics(MemOrder: MemOrdNeq) | MemSemStorage;
749 if (static_cast<unsigned>(MemOrdEq) == MemSemEqual)
750 MemSemEqualReg = Call->Arguments[3];
751 if (static_cast<unsigned>(MemOrdNeq) == MemSemUnequal)
752 MemSemUnequalReg = Call->Arguments[4];
753 }
754 if (!MemSemEqualReg.isValid())
755 MemSemEqualReg = buildConstantIntReg32(Val: MemSemEqual, MIRBuilder, GR);
756 if (!MemSemUnequalReg.isValid())
757 MemSemUnequalReg = buildConstantIntReg32(Val: MemSemUnequal, MIRBuilder, GR);
758
759 Register ScopeReg;
760 auto Scope = IsCmpxchg ? SPIRV::Scope::Workgroup : SPIRV::Scope::Device;
761 if (Call->Arguments.size() >= 6) {
762 assert(Call->Arguments.size() == 6 &&
763 "Extra args for explicit atomic cmpxchg");
764 auto ClScope = static_cast<SPIRV::CLMemoryScope>(
765 getIConstVal(ConstReg: Call->Arguments[5], MRI));
766 Scope = getSPIRVScope(ClScope);
767 if (ClScope == static_cast<unsigned>(Scope))
768 ScopeReg = Call->Arguments[5];
769 }
770 if (!ScopeReg.isValid())
771 ScopeReg = buildConstantIntReg32(Val: Scope, MIRBuilder, GR);
772
773 Register Expected =
774 IsCmpxchg ? ExpectedArg
775 : buildLoadInst(BaseType: SpvDesiredTy, PtrRegister: ExpectedArg, MIRBuilder, GR);
776 MRI->setType(VReg: Expected, Ty: DesiredLLT);
777 Register Tmp = !IsCmpxchg ? MRI->createGenericVirtualRegister(Ty: DesiredLLT)
778 : Call->ReturnRegister;
779 if (!MRI->getRegClassOrNull(Reg: Tmp))
780 MRI->setRegClass(Reg: Tmp, RC: GR->getRegClass(SpvType: SpvDesiredTy));
781 GR->assignSPIRVTypeToVReg(Type: SpvDesiredTy, VReg: Tmp, MF: MIRBuilder.getMF());
782
783 MIRBuilder.buildInstr(Opcode)
784 .addDef(RegNo: Tmp)
785 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: SpvDesiredTy))
786 .addUse(RegNo: ObjectPtr)
787 .addUse(RegNo: ScopeReg)
788 .addUse(RegNo: MemSemEqualReg)
789 .addUse(RegNo: MemSemUnequalReg)
790 .addUse(RegNo: Desired)
791 .addUse(RegNo: Expected);
792 if (!IsCmpxchg) {
793 MIRBuilder.buildInstr(Opcode: SPIRV::OpStore).addUse(RegNo: ExpectedArg).addUse(RegNo: Tmp);
794 MIRBuilder.buildICmp(Pred: CmpInst::ICMP_EQ, Res: Call->ReturnRegister, Op0: Tmp, Op1: Expected);
795 }
796 return true;
797}
798
799/// Helper function for building atomic instructions.
800static bool buildAtomicRMWInst(const SPIRV::IncomingCall *Call, unsigned Opcode,
801 MachineIRBuilder &MIRBuilder,
802 SPIRVGlobalRegistry *GR) {
803 if (Call->isSpirvOp())
804 return buildOpFromWrapper(MIRBuilder, Opcode, Call,
805 TypeReg: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
806
807 MachineRegisterInfo *MRI = MIRBuilder.getMRI();
808 StringRef Name = Call->Builtin->name();
809 // The registry prefixes atomic_fetch_min/max with "s_" or "u_".
810 bool IsOCL20 =
811 Name.contains(Other: "atomic_fetch_") || Name.starts_with(Prefix: "atomic_exchange");
812 SPIRV::Scope::Scope DefaultScope =
813 IsOCL20 ? SPIRV::Scope::Device : SPIRV::Scope::Workgroup;
814 Register ScopeRegister =
815 Call->Arguments.size() >= 4 ? Call->Arguments[3] : Register();
816
817 assert(Call->Arguments.size() <= 4 &&
818 "Too many args for explicit atomic RMW");
819 ScopeRegister =
820 buildScopeReg(CLScopeRegister: ScopeRegister, Scope: DefaultScope, MIRBuilder, GR, MRI);
821
822 Register PtrRegister = Call->Arguments[0];
823 SPIRV::MemorySemantics::MemorySemantics Ordering =
824 IsOCL20 ? SPIRV::MemorySemantics::SequentiallyConsistent
825 : SPIRV::MemorySemantics::None;
826 unsigned StorageClassSem = SPIRV::MemorySemantics::None;
827 if (Call->Arguments.size() >= 3)
828 Ordering = getMemOrdering(OrderRegister: Call->Arguments[2], MRI);
829 if (IsOCL20 || Call->Arguments.size() >= 3)
830 StorageClassSem =
831 getMemSemanticsForStorageClass(SC: GR->getPointerStorageClass(VReg: PtrRegister));
832 Register MemSemanticsReg =
833 buildMemSemanticsReg(Ordering, StorageClassSem, MIRBuilder, GR);
834 Register ValueReg = Call->Arguments[1];
835 Register ValueTypeReg = GR->getSPIRVTypeID(SpirvType: Call->ReturnType);
836 // support cl_ext_float_atomics
837 if (Call->ReturnType->getOpcode() == SPIRV::OpTypeFloat) {
838 if (Opcode == SPIRV::OpAtomicIAdd) {
839 Opcode = SPIRV::OpAtomicFAddEXT;
840 } else if (Opcode == SPIRV::OpAtomicISub) {
841 // Translate OpAtomicISub applied to a floating type argument to
842 // OpAtomicFAddEXT with the negative value operand
843 Opcode = SPIRV::OpAtomicFAddEXT;
844 Register NegValueReg =
845 MRI->createGenericVirtualRegister(Ty: MRI->getType(Reg: ValueReg));
846 MRI->setRegClass(Reg: NegValueReg, RC: GR->getRegClass(SpvType: Call->ReturnType));
847 GR->assignSPIRVTypeToVReg(Type: Call->ReturnType, VReg: NegValueReg,
848 MF: MIRBuilder.getMF());
849 MIRBuilder.buildInstr(Opcode: TargetOpcode::G_FNEG)
850 .addDef(RegNo: NegValueReg)
851 .addUse(RegNo: ValueReg);
852 updateRegType(Reg: NegValueReg, Ty: nullptr, SpirvTy: Call->ReturnType, GR, MIB&: MIRBuilder,
853 MRI&: MIRBuilder.getMF().getRegInfo());
854 ValueReg = NegValueReg;
855 }
856 }
857 MIRBuilder.buildInstr(Opcode)
858 .addDef(RegNo: Call->ReturnRegister)
859 .addUse(RegNo: ValueTypeReg)
860 .addUse(RegNo: PtrRegister)
861 .addUse(RegNo: ScopeRegister)
862 .addUse(RegNo: MemSemanticsReg)
863 .addUse(RegNo: ValueReg);
864 return true;
865}
866
867/// Helper function for building an atomic floating-type instruction.
868static bool buildAtomicFloatingRMWInst(const SPIRV::IncomingCall *Call,
869 unsigned Opcode,
870 MachineIRBuilder &MIRBuilder,
871 SPIRVGlobalRegistry *GR) {
872 assert(Call->Arguments.size() == 4 &&
873 "Wrong number of atomic floating-type builtin");
874 Register PtrReg = Call->Arguments[0];
875 Register ScopeReg = Call->Arguments[1];
876 Register MemSemanticsReg = Call->Arguments[2];
877 Register ValueReg = Call->Arguments[3];
878 MIRBuilder.buildInstr(Opcode)
879 .addDef(RegNo: Call->ReturnRegister)
880 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType))
881 .addUse(RegNo: PtrReg)
882 .addUse(RegNo: ScopeReg)
883 .addUse(RegNo: MemSemanticsReg)
884 .addUse(RegNo: ValueReg);
885 return true;
886}
887
888/// Helper function for building atomic flag instructions (e.g.
889/// OpAtomicFlagTestAndSet).
890static bool buildAtomicFlagInst(const SPIRV::IncomingCall *Call,
891 unsigned Opcode, MachineIRBuilder &MIRBuilder,
892 SPIRVGlobalRegistry *GR) {
893 bool IsSet = Opcode == SPIRV::OpAtomicFlagTestAndSet;
894 Register TypeReg = GR->getSPIRVTypeID(SpirvType: Call->ReturnType);
895 if (Call->isSpirvOp())
896 return buildOpFromWrapper(MIRBuilder, Opcode, Call,
897 TypeReg: IsSet ? TypeReg : Register(0));
898
899 MachineRegisterInfo *MRI = MIRBuilder.getMRI();
900 Register PtrRegister = Call->Arguments[0];
901 SPIRV::MemorySemantics::MemorySemantics Ordering =
902 SPIRV::MemorySemantics::SequentiallyConsistent;
903 unsigned StorageClassSem = SPIRV::MemorySemantics::None;
904 if (Call->Arguments.size() >= 2) {
905 Ordering = getMemOrdering(OrderRegister: Call->Arguments[1], MRI);
906 StorageClassSem =
907 getMemSemanticsForStorageClass(SC: GR->getPointerStorageClass(VReg: PtrRegister));
908 }
909
910 assert((Opcode != SPIRV::OpAtomicFlagClear ||
911 (Ordering != SPIRV::MemorySemantics::Acquire &&
912 Ordering != SPIRV::MemorySemantics::AcquireRelease)) &&
913 "Invalid memory order argument!");
914
915 Register MemSemanticsReg =
916 buildMemSemanticsReg(Ordering, StorageClassSem, MIRBuilder, GR);
917
918 Register ScopeRegister = buildScopeReg(
919 CLScopeRegister: Call->Arguments.size() >= 3 ? Call->Arguments[2] : Register(),
920 Scope: SPIRV::Scope::Device, MIRBuilder, GR, MRI);
921
922 auto MIB = MIRBuilder.buildInstr(Opcode);
923 if (IsSet)
924 MIB.addDef(RegNo: Call->ReturnRegister).addUse(RegNo: TypeReg);
925
926 MIB.addUse(RegNo: PtrRegister).addUse(RegNo: ScopeRegister).addUse(RegNo: MemSemanticsReg);
927 return true;
928}
929
930static void reportUnsupported(MachineIRBuilder &MIRBuilder, const Twine &Msg) {
931 const Function &F = MIRBuilder.getMF().getFunction();
932 F.getContext().diagnose(
933 DI: DiagnosticInfoUnsupported(F, Msg, MIRBuilder.getDebugLoc()));
934}
935
936/// Helper function for building barriers, i.e., memory/control ordering
937/// operations.
938static bool buildBarrierInst(const SPIRV::IncomingCall *Call, unsigned Opcode,
939 MachineIRBuilder &MIRBuilder,
940 SPIRVGlobalRegistry *GR) {
941 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
942 const auto *ST =
943 static_cast<const SPIRVSubtarget *>(&MIRBuilder.getMF().getSubtarget());
944 if ((Opcode == SPIRV::OpControlBarrierArriveINTEL ||
945 Opcode == SPIRV::OpControlBarrierWaitINTEL) &&
946 !ST->canUseExtension(E: SPIRV::Extension::SPV_INTEL_split_barrier)) {
947 std::string DiagMsg = std::string(Builtin->name()) +
948 ": the builtin requires the following SPIR-V "
949 "extension: SPV_INTEL_split_barrier";
950 report_fatal_error(reason: DiagMsg.c_str(), gen_crash_diag: false);
951 }
952
953 if (Call->isSpirvOp())
954 return buildOpFromWrapper(MIRBuilder, Opcode, Call, TypeReg: Register(0));
955
956 MachineRegisterInfo *MRI = MIRBuilder.getMRI();
957 bool IsSubgroupBarrier = Builtin->name() == "sub_group_barrier";
958 if (IsSubgroupBarrier) {
959 // TODO: Support runtime flags and scopes for OpenCL barriers.
960 for (Register Arg : Call->Arguments) {
961 const MachineInstr *MI = getDefInstrMaybeConstant(ConstReg&: Arg, MRI);
962 if (!MI || MI->getOpcode() != TargetOpcode::G_CONSTANT) {
963 reportUnsupported(
964 MIRBuilder,
965 Msg: "sub_group_barrier with non-constant arguments is not supported");
966 return false;
967 }
968 }
969 }
970 unsigned MemFlags = getIConstVal(ConstReg: Call->Arguments[0], MRI);
971 unsigned MemSemantics = SPIRV::MemorySemantics::None;
972
973 if (MemFlags & SPIRV::CLK_LOCAL_MEM_FENCE)
974 MemSemantics |= SPIRV::MemorySemantics::WorkgroupMemory;
975
976 if (MemFlags & SPIRV::CLK_GLOBAL_MEM_FENCE)
977 MemSemantics |= SPIRV::MemorySemantics::CrossWorkgroupMemory;
978
979 if (MemFlags & SPIRV::CLK_IMAGE_MEM_FENCE)
980 MemSemantics |= SPIRV::MemorySemantics::ImageMemory;
981
982 if (Opcode == SPIRV::OpMemoryBarrier)
983 MemSemantics = getSPIRVMemSemantics(MemOrder: static_cast<std::memory_order>(
984 getIConstVal(ConstReg: Call->Arguments[1], MRI))) |
985 MemSemantics;
986 else if (Opcode == SPIRV::OpControlBarrierArriveINTEL)
987 MemSemantics |= SPIRV::MemorySemantics::Release;
988 else if (Opcode == SPIRV::OpControlBarrierWaitINTEL)
989 MemSemantics |= SPIRV::MemorySemantics::Acquire;
990 else
991 MemSemantics |= SPIRV::MemorySemantics::SequentiallyConsistent;
992
993 Register MemSemanticsReg =
994 MemFlags == MemSemantics
995 ? Call->Arguments[0]
996 : buildConstantIntReg32(Val: MemSemantics, MIRBuilder, GR);
997 Register ScopeReg;
998 SPIRV::Scope::Scope Scope =
999 IsSubgroupBarrier ? SPIRV::Scope::Subgroup : SPIRV::Scope::Workgroup;
1000 SPIRV::Scope::Scope MemScope = Scope;
1001 if (Call->Arguments.size() >= 2) {
1002 assert(
1003 ((Opcode != SPIRV::OpMemoryBarrier && Call->Arguments.size() == 2) ||
1004 (Opcode == SPIRV::OpMemoryBarrier && Call->Arguments.size() == 3)) &&
1005 "Extra args for explicitly scoped barrier");
1006 Register ScopeArg = (Opcode == SPIRV::OpMemoryBarrier) ? Call->Arguments[2]
1007 : Call->Arguments[1];
1008 SPIRV::CLMemoryScope CLScope =
1009 static_cast<SPIRV::CLMemoryScope>(getIConstVal(ConstReg: ScopeArg, MRI));
1010 MemScope = getSPIRVScope(ClScope: CLScope);
1011 if (Opcode == SPIRV::OpMemoryBarrier)
1012 Scope = MemScope;
1013 if (CLScope == static_cast<unsigned>(Scope))
1014 ScopeReg = Call->Arguments[1];
1015 }
1016
1017 if (!ScopeReg.isValid())
1018 ScopeReg = buildConstantIntReg32(Val: Scope, MIRBuilder, GR);
1019
1020 auto MIB = MIRBuilder.buildInstr(Opcode).addUse(RegNo: ScopeReg);
1021 if (Opcode != SPIRV::OpMemoryBarrier)
1022 MIB.addUse(RegNo: buildConstantIntReg32(Val: MemScope, MIRBuilder, GR));
1023 MIB.addUse(RegNo: MemSemanticsReg);
1024 return true;
1025}
1026
1027/// Helper function for building extended bit operations.
1028static bool buildExtendedBitOpsInst(const SPIRV::IncomingCall *Call,
1029 unsigned Opcode,
1030 MachineIRBuilder &MIRBuilder,
1031 SPIRVGlobalRegistry *GR) {
1032 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
1033 const auto *ST =
1034 static_cast<const SPIRVSubtarget *>(&MIRBuilder.getMF().getSubtarget());
1035 if ((Opcode == SPIRV::OpBitFieldInsert ||
1036 Opcode == SPIRV::OpBitFieldSExtract ||
1037 Opcode == SPIRV::OpBitFieldUExtract || Opcode == SPIRV::OpBitReverse) &&
1038 !ST->canUseExtension(E: SPIRV::Extension::SPV_KHR_bit_instructions)) {
1039 std::string DiagMsg = std::string(Builtin->name()) +
1040 ": the builtin requires the following SPIR-V "
1041 "extension: SPV_KHR_bit_instructions";
1042 report_fatal_error(reason: DiagMsg.c_str(), gen_crash_diag: false);
1043 }
1044
1045 // Generate SPIRV instruction accordingly.
1046 if (Call->isSpirvOp())
1047 return buildOpFromWrapper(MIRBuilder, Opcode, Call,
1048 TypeReg: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
1049
1050 auto MIB = MIRBuilder.buildInstr(Opcode)
1051 .addDef(RegNo: Call->ReturnRegister)
1052 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
1053 for (unsigned i = 0; i < Call->Arguments.size(); ++i)
1054 MIB.addUse(RegNo: Call->Arguments[i]);
1055
1056 return true;
1057}
1058
1059/// Helper function for building Intel's bindless image instructions.
1060static bool buildBindlessImageINTELInst(const SPIRV::IncomingCall *Call,
1061 unsigned Opcode,
1062 MachineIRBuilder &MIRBuilder,
1063 SPIRVGlobalRegistry *GR) {
1064 // Generate SPIRV instruction accordingly.
1065 if (Call->isSpirvOp())
1066 return buildOpFromWrapper(MIRBuilder, Opcode, Call,
1067 TypeReg: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
1068
1069 MIRBuilder.buildInstr(Opcode)
1070 .addDef(RegNo: Call->ReturnRegister)
1071 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType))
1072 .addUse(RegNo: Call->Arguments[0]);
1073
1074 return true;
1075}
1076
1077/// Helper function for building Intel's OpBitwiseFunctionINTEL instruction.
1078static bool buildTernaryBitwiseFunctionINTELInst(
1079 const SPIRV::IncomingCall *Call, unsigned Opcode,
1080 MachineIRBuilder &MIRBuilder, SPIRVGlobalRegistry *GR) {
1081 // Generate SPIRV instruction accordingly.
1082 if (Call->isSpirvOp())
1083 return buildOpFromWrapper(MIRBuilder, Opcode, Call,
1084 TypeReg: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
1085
1086 auto MIB = MIRBuilder.buildInstr(Opcode)
1087 .addDef(RegNo: Call->ReturnRegister)
1088 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
1089 for (unsigned i = 0; i < Call->Arguments.size(); ++i)
1090 MIB.addUse(RegNo: Call->Arguments[i]);
1091
1092 return true;
1093}
1094
1095static bool buildImageChannelDataTypeInst(const SPIRV::IncomingCall *Call,
1096 unsigned Opcode,
1097 MachineIRBuilder &MIRBuilder,
1098 SPIRVGlobalRegistry *GR) {
1099 if (Call->isSpirvOp())
1100 return buildOpFromWrapper(MIRBuilder, Opcode, Call,
1101 TypeReg: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
1102
1103 auto MIB = MIRBuilder.buildInstr(Opcode)
1104 .addDef(RegNo: Call->ReturnRegister)
1105 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
1106 for (unsigned i = 0; i < Call->Arguments.size(); ++i)
1107 MIB.addUse(RegNo: Call->Arguments[i]);
1108
1109 return true;
1110}
1111
1112/// Helper function for building Intel's 2d block io instructions.
1113static bool build2DBlockIOINTELInst(const SPIRV::IncomingCall *Call,
1114 unsigned Opcode,
1115 MachineIRBuilder &MIRBuilder,
1116 SPIRVGlobalRegistry *GR) {
1117 // Generate SPIRV instruction accordingly.
1118 if (Call->isSpirvOp())
1119 return buildOpFromWrapper(MIRBuilder, Opcode, Call, TypeReg: Register(0));
1120
1121 auto MIB = MIRBuilder.buildInstr(Opcode)
1122 .addDef(RegNo: Call->ReturnRegister)
1123 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
1124 for (unsigned i = 0; i < Call->Arguments.size(); ++i)
1125 MIB.addUse(RegNo: Call->Arguments[i]);
1126
1127 return true;
1128}
1129
1130static bool buildPipeInst(const SPIRV::IncomingCall *Call, unsigned Opcode,
1131 unsigned Scope, MachineIRBuilder &MIRBuilder,
1132 SPIRVGlobalRegistry *GR) {
1133 switch (Opcode) {
1134 case SPIRV::OpCommitReadPipe:
1135 case SPIRV::OpCommitWritePipe:
1136 return buildOpFromWrapper(MIRBuilder, Opcode, Call, TypeReg: Register(0));
1137 case SPIRV::OpGroupCommitReadPipe:
1138 case SPIRV::OpGroupCommitWritePipe:
1139 case SPIRV::OpGroupReserveReadPipePackets:
1140 case SPIRV::OpGroupReserveWritePipePackets: {
1141 Register ScopeConstReg =
1142 MIRBuilder.buildConstant(Res: LLT::scalar(SizeInBits: 32), Val: Scope).getReg(Idx: 0);
1143 MachineRegisterInfo *MRI = MIRBuilder.getMRI();
1144 MRI->setRegClass(Reg: ScopeConstReg, RC: &SPIRV::iIDRegClass);
1145 MachineInstrBuilder MIB;
1146 MIB = MIRBuilder.buildInstr(Opcode);
1147 // Add Return register and type.
1148 if (Opcode == SPIRV::OpGroupReserveReadPipePackets ||
1149 Opcode == SPIRV::OpGroupReserveWritePipePackets)
1150 MIB.addDef(RegNo: Call->ReturnRegister)
1151 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
1152
1153 MIB.addUse(RegNo: ScopeConstReg);
1154 for (unsigned int i = 0; i < Call->Arguments.size(); ++i)
1155 MIB.addUse(RegNo: Call->Arguments[i]);
1156
1157 return true;
1158 }
1159 default:
1160 return buildOpFromWrapper(MIRBuilder, Opcode, Call,
1161 TypeReg: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
1162 }
1163}
1164
1165static unsigned getNumComponentsForDim(SPIRV::Dim::Dim dim) {
1166 switch (dim) {
1167 case SPIRV::Dim::DIM_1D:
1168 case SPIRV::Dim::DIM_Buffer:
1169 return 1;
1170 case SPIRV::Dim::DIM_2D:
1171 case SPIRV::Dim::DIM_Cube:
1172 case SPIRV::Dim::DIM_Rect:
1173 return 2;
1174 case SPIRV::Dim::DIM_3D:
1175 return 3;
1176 default:
1177 report_fatal_error(reason: "Cannot get num components for given Dim");
1178 }
1179}
1180
1181/// Helper function for obtaining the number of size components.
1182static unsigned getNumSizeComponents(SPIRVTypeInst imgType) {
1183 assert(imgType->getOpcode() == SPIRV::OpTypeImage);
1184 auto dim = static_cast<SPIRV::Dim::Dim>(imgType->getOperand(i: 2).getImm());
1185 unsigned numComps = getNumComponentsForDim(dim);
1186 bool arrayed = imgType->getOperand(i: 4).getImm() == 1;
1187 return arrayed ? numComps + 1 : numComps;
1188}
1189
1190static bool builtinMayNeedPromotionToVec(uint32_t BuiltinNumber) {
1191 switch (BuiltinNumber) {
1192 case SPIRV::OpenCLExtInst::s_min:
1193 case SPIRV::OpenCLExtInst::u_min:
1194 case SPIRV::OpenCLExtInst::s_max:
1195 case SPIRV::OpenCLExtInst::u_max:
1196 case SPIRV::OpenCLExtInst::fmax:
1197 case SPIRV::OpenCLExtInst::fmin:
1198 case SPIRV::OpenCLExtInst::fmax_common:
1199 case SPIRV::OpenCLExtInst::fmin_common:
1200 case SPIRV::OpenCLExtInst::s_clamp:
1201 case SPIRV::OpenCLExtInst::fclamp:
1202 case SPIRV::OpenCLExtInst::u_clamp:
1203 case SPIRV::OpenCLExtInst::mix:
1204 case SPIRV::OpenCLExtInst::step:
1205 case SPIRV::OpenCLExtInst::smoothstep:
1206 case SPIRV::OpenCLExtInst::ldexp:
1207 case SPIRV::OpenCLExtInst::pown:
1208 case SPIRV::OpenCLExtInst::rootn:
1209 return true;
1210 default:
1211 break;
1212 }
1213 return false;
1214}
1215
1216//===----------------------------------------------------------------------===//
1217// Implementation functions for each builtin group
1218//===----------------------------------------------------------------------===//
1219
1220static SmallVector<Register>
1221getBuiltinCallArguments(const SPIRV::IncomingCall *Call, uint32_t BuiltinNumber,
1222 MachineIRBuilder &MIRBuilder, SPIRVGlobalRegistry *GR) {
1223
1224 Register ReturnTypeId = GR->getSPIRVTypeID(SpirvType: Call->ReturnType);
1225 unsigned ResultElementCount =
1226 GR->getScalarOrVectorComponentCount(VReg: ReturnTypeId);
1227 bool MayNeedPromotionToVec =
1228 builtinMayNeedPromotionToVec(BuiltinNumber) && ResultElementCount > 1;
1229
1230 if (!MayNeedPromotionToVec)
1231 return {Call->Arguments.begin(), Call->Arguments.end()};
1232
1233 SmallVector<Register> Arguments;
1234 for (Register Argument : Call->Arguments) {
1235 Register VecArg = Argument;
1236 SPIRVTypeInst ArgumentType = GR->getSPIRVTypeForVReg(VReg: Argument);
1237 if (GR->getScalarOrVectorComponentCount(Type: ArgumentType) == 1 &&
1238 ArgumentType != Call->ReturnType) {
1239 SPIRVTypeInst VecType = GR->getOrCreateSPIRVVectorType(
1240 BaseType: ArgumentType, NumElements: ResultElementCount, MIRBuilder, /*EmitIR=*/true);
1241 VecArg = createVirtualRegister(SpvType: VecType, GR, MIRBuilder);
1242 Register VecTypeId = GR->getSPIRVTypeID(SpirvType: VecType);
1243 auto VecSplat = MIRBuilder.buildInstr(Opcode: SPIRV::OpCompositeConstruct)
1244 .addDef(RegNo: VecArg)
1245 .addUse(RegNo: VecTypeId);
1246 for (unsigned I = 0; I != ResultElementCount; ++I)
1247 VecSplat.addUse(RegNo: Argument);
1248 }
1249 Arguments.push_back(Elt: VecArg);
1250 }
1251 return Arguments;
1252}
1253
1254static bool generateExtInst(const SPIRV::IncomingCall *Call,
1255 MachineIRBuilder &MIRBuilder,
1256 SPIRVGlobalRegistry *GR, const CallBase &CB) {
1257 // Lookup the extended instruction number in the TableGen records.
1258 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
1259 uint32_t Number =
1260 SPIRV::lookupExtendedBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Number;
1261 // fmin_common and fmax_common are now deprecated, and we should use fmin and
1262 // fmax with NotInf and NotNaN flags instead. Keep original number to add
1263 // later the NoNans and NoInfs flags.
1264 uint32_t OrigNumber = Number;
1265 const SPIRVSubtarget &ST =
1266 cast<SPIRVSubtarget>(Val: MIRBuilder.getMF().getSubtarget());
1267 if (ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_float_controls2) &&
1268 (Number == SPIRV::OpenCLExtInst::fmin_common ||
1269 Number == SPIRV::OpenCLExtInst::fmax_common)) {
1270 Number = (Number == SPIRV::OpenCLExtInst::fmin_common)
1271 ? SPIRV::OpenCLExtInst::fmin
1272 : SPIRV::OpenCLExtInst::fmax;
1273 }
1274
1275 // ExtInst prefetch cannot take an untyped pointer, so emit
1276 // OpUntypedPrefetchKHR with Num Bytes = num elements * element byte size.
1277 if (Number == SPIRV::OpenCLExtInst::prefetch && Call->Arguments.size() >= 2) {
1278 Register PtrReg = Call->Arguments[0];
1279 SPIRVTypeInst PtrTy = GR->getSPIRVTypeForVReg(VReg: PtrReg);
1280 if (PtrTy && PtrTy->getOpcode() == SPIRV::OpTypeUntypedPointerKHR) {
1281 MachineRegisterInfo *MRI = MIRBuilder.getMRI();
1282 Register NumElems = Call->Arguments[1];
1283 SPIRVTypeInst SizeTy = GR->getSPIRVTypeForVReg(VReg: NumElems);
1284 assert(SizeTy && "Expected a type for the number of elements");
1285 unsigned ElemBytes = GR->getDeducedPointeeByteSize(PtrVal: CB.getArgOperand(i: 0));
1286 Register NumBytes = NumElems;
1287 // A byte sized element already makes the element count a byte count. A
1288 // size of 0 means the element type could not be deduced, which the typed
1289 // lowering resolves to i8, so treat it as a single byte here as well.
1290 if (ElemBytes > 1) {
1291 Register ElemBytesReg = GR->buildConstantInt(Val: ElemBytes, MIRBuilder,
1292 SpvType: SizeTy, /*EmitIR=*/true);
1293 Register Mul =
1294 MRI->createGenericVirtualRegister(Ty: MRI->getType(Reg: NumElems));
1295 MRI->setRegClass(Reg: Mul, RC: GR->getRegClass(SpvType: SizeTy));
1296 GR->assignSPIRVTypeToVReg(Type: SizeTy, VReg: Mul, MF: MIRBuilder.getMF());
1297 MIRBuilder.buildInstr(Opcode: TargetOpcode::G_MUL)
1298 .addDef(RegNo: Mul)
1299 .addUse(RegNo: NumElems)
1300 .addUse(RegNo: ElemBytesReg);
1301 updateRegType(Reg: Mul, /*Ty=*/nullptr, SpirvTy: SizeTy, GR, MIB&: MIRBuilder, MRI&: *MRI);
1302 NumBytes = Mul;
1303 }
1304 MIRBuilder.buildInstr(Opcode: SPIRV::OpUntypedPrefetchKHR)
1305 .addUse(RegNo: PtrReg)
1306 .addUse(RegNo: NumBytes);
1307 return true;
1308 }
1309 }
1310
1311 Register ReturnTypeId = GR->getSPIRVTypeID(SpirvType: Call->ReturnType);
1312 SmallVector<Register> Arguments =
1313 getBuiltinCallArguments(Call, BuiltinNumber: Number, MIRBuilder, GR);
1314
1315 MachineInstrBuilder MIB;
1316 if (ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_fma) &&
1317 Number == SPIRV::OpenCLExtInst::fma) {
1318 // Use the SPIR-V fma instruction instead of the OpenCL extended
1319 // instruction if the extension is available.
1320 MIB = MIRBuilder.buildInstr(Opcode: SPIRV::OpFmaKHR)
1321 .addDef(RegNo: Call->ReturnRegister)
1322 .addUse(RegNo: ReturnTypeId);
1323 } else {
1324 // Build extended instruction.
1325 MIB = MIRBuilder.buildInstr(Opcode: SPIRV::OpExtInst)
1326 .addDef(RegNo: Call->ReturnRegister)
1327 .addUse(RegNo: ReturnTypeId)
1328 .addImm(Val: static_cast<uint32_t>(SPIRV::InstructionSet::OpenCL_std))
1329 .addImm(Val: Number);
1330 }
1331
1332 for (Register Argument : Arguments)
1333 MIB.addUse(RegNo: Argument);
1334
1335 MIB.getInstr()->copyIRFlags(I: CB);
1336 if (OrigNumber == SPIRV::OpenCLExtInst::fmin_common ||
1337 OrigNumber == SPIRV::OpenCLExtInst::fmax_common) {
1338 // Add NoNans and NoInfs flags to fmin/fmax instruction.
1339 MIB.getInstr()->setFlag(MachineInstr::MIFlag::FmNoNans);
1340 MIB.getInstr()->setFlag(MachineInstr::MIFlag::FmNoInfs);
1341 }
1342
1343 // Derive fast-math flags from nofpclass attributes on the called function.
1344 // FPFastMathMode decoration is valid on ExtInst in Kernel environments
1345 // (SPIR-V core) or with SPV_KHR_float_controls2 for any environment.
1346 if (ST.isKernel() ||
1347 ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_float_controls2)) {
1348 if (const Function *F = CB.getCalledFunction()) {
1349 bool AddNoNan = CB.getRetNoFPClass() & fcNan;
1350 bool AddNoInf = CB.getRetNoFPClass() & fcInf;
1351 FunctionType *FTy = F->getFunctionType();
1352 for (unsigned I = 0, E = FTy->getNumParams();
1353 I != E && (AddNoNan || AddNoInf); ++I) {
1354 if (!FTy->getParamType(i: I)->isFloatingPointTy())
1355 continue;
1356 FPClassTest ArgTest = CB.getParamNoFPClass(i: I);
1357 AddNoNan = AddNoNan && ArgTest & fcNan;
1358 AddNoInf = AddNoInf && ArgTest & fcInf;
1359 }
1360 if (AddNoNan)
1361 MIB.getInstr()->setFlag(MachineInstr::MIFlag::FmNoNans);
1362 if (AddNoInf)
1363 MIB.getInstr()->setFlag(MachineInstr::MIFlag::FmNoInfs);
1364 }
1365 }
1366
1367 return true;
1368}
1369
1370static bool generateRelationalInst(const SPIRV::IncomingCall *Call,
1371 MachineIRBuilder &MIRBuilder,
1372 SPIRVGlobalRegistry *GR) {
1373 // Lookup the instruction opcode in the TableGen records.
1374 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
1375 unsigned Opcode =
1376 SPIRV::lookupNativeBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Opcode;
1377
1378 Register CompareRegister;
1379 SPIRVTypeInst RelationType = nullptr;
1380 std::tie(args&: CompareRegister, args&: RelationType) =
1381 buildBoolRegister(MIRBuilder, ResultType: Call->ReturnType, GR);
1382
1383 // OpAny/OpAll require a boolean vector input, but OpenCL any()/all()
1384 // builtins receive integer vectors. Convert via OpINotEqual against zero.
1385 SmallVector<Register> Arguments(Call->Arguments.begin(),
1386 Call->Arguments.end());
1387 if ((Opcode == SPIRV::OpAny || Opcode == SPIRV::OpAll) &&
1388 !GR->isScalarOrVectorOfType(VReg: Arguments[0], TypeOpcode: SPIRV::OpTypeBool)) {
1389 SPIRVTypeInst ArgType = GR->getSPIRVTypeForVReg(VReg: Arguments[0]);
1390 unsigned NumElts = GR->getScalarOrVectorComponentCount(Type: ArgType);
1391 SPIRVTypeInst BoolVecTy = GR->getOrCreateSPIRVVectorType(
1392 BaseType: GR->getOrCreateSPIRVBoolType(MIRBuilder, /*EmitIR=*/true), NumElements: NumElts,
1393 MIRBuilder, /*EmitIR=*/true);
1394 Register ZeroReg =
1395 GR->getOrCreateConsIntVector(Val: uint64_t(0), MIRBuilder, SpvType: ArgType,
1396 /*EmitIR=*/true);
1397 Register BoolVecReg = createVirtualRegister(SpvType: BoolVecTy, GR, MIRBuilder);
1398 MIRBuilder.buildInstr(Opcode: SPIRV::OpINotEqual)
1399 .addDef(RegNo: BoolVecReg)
1400 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: BoolVecTy))
1401 .addUse(RegNo: Arguments[0])
1402 .addUse(RegNo: ZeroReg);
1403 Arguments[0] = BoolVecReg;
1404 }
1405
1406 // Build relational instruction.
1407 auto MIB = MIRBuilder.buildInstr(Opcode)
1408 .addDef(RegNo: CompareRegister)
1409 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: RelationType));
1410
1411 for (auto Argument : Arguments)
1412 MIB.addUse(RegNo: Argument);
1413
1414 // Build select instruction.
1415 return buildSelectInst(MIRBuilder, ReturnRegister: Call->ReturnRegister, SourceRegister: CompareRegister,
1416 ReturnType: Call->ReturnType, GR);
1417}
1418
1419static bool generateGroupInst(const SPIRV::IncomingCall *Call,
1420 MachineIRBuilder &MIRBuilder,
1421 SPIRVGlobalRegistry *GR) {
1422 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
1423 const SPIRV::GroupBuiltin *GroupBuiltin =
1424 SPIRV::lookupGroupBuiltin(Name: Builtin->name());
1425
1426 MachineRegisterInfo *MRI = MIRBuilder.getMRI();
1427 if (Call->isSpirvOp()) {
1428 if (GroupBuiltin->NoGroupOperation) {
1429 SmallVector<uint32_t, 1> ImmArgs;
1430 if (GroupBuiltin->Opcode ==
1431 SPIRV::OpSubgroupMatrixMultiplyAccumulateINTEL &&
1432 Call->Arguments.size() > 4)
1433 ImmArgs.push_back(Elt: getIConstVal(ConstReg: Call->Arguments[4], MRI));
1434 return buildOpFromWrapper(MIRBuilder, Opcode: GroupBuiltin->Opcode, Call,
1435 TypeReg: GR->getSPIRVTypeID(SpirvType: Call->ReturnType), ImmArgs);
1436 }
1437
1438 // Group Operation is a literal
1439 Register GroupOpReg = Call->Arguments[1];
1440 const MachineInstr *MI = getDefInstrMaybeConstant(ConstReg&: GroupOpReg, MRI);
1441 if (!MI || MI->getOpcode() != TargetOpcode::G_CONSTANT)
1442 report_fatal_error(
1443 reason: "Group Operation parameter must be an integer constant");
1444 uint64_t GrpOp = MI->getOperand(i: 1).getCImm()->getValue().getZExtValue();
1445 Register ScopeReg = Call->Arguments[0];
1446 auto MIB = MIRBuilder.buildInstr(Opcode: GroupBuiltin->Opcode)
1447 .addDef(RegNo: Call->ReturnRegister)
1448 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType))
1449 .addUse(RegNo: ScopeReg)
1450 .addImm(Val: GrpOp);
1451 for (unsigned i = 2; i < Call->Arguments.size(); ++i)
1452 MIB.addUse(RegNo: Call->Arguments[i]);
1453 return true;
1454 }
1455
1456 Register Arg0;
1457 if (GroupBuiltin->HasBoolArg) {
1458 SPIRVTypeInst BoolType = GR->getOrCreateSPIRVBoolType(MIRBuilder, EmitIR: true);
1459 Register BoolReg = Call->Arguments[0];
1460 SPIRVTypeInst BoolRegType = GR->getSPIRVTypeForVReg(VReg: BoolReg);
1461 if (!BoolRegType)
1462 report_fatal_error(reason: "Can't find a register's type definition");
1463 MachineInstr *ArgInstruction = getDefInstrMaybeConstant(ConstReg&: BoolReg, MRI);
1464 if (ArgInstruction->getOpcode() == TargetOpcode::G_CONSTANT) {
1465 if (BoolRegType->getOpcode() != SPIRV::OpTypeBool)
1466 Arg0 = GR->buildConstantInt(Val: getIConstVal(ConstReg: BoolReg, MRI) != 0, MIRBuilder,
1467 SpvType: BoolType, EmitIR: true);
1468 } else {
1469 if (BoolRegType->getOpcode() == SPIRV::OpTypeInt) {
1470 Arg0 = MRI->createGenericVirtualRegister(Ty: LLT::scalar(SizeInBits: 1));
1471 MRI->setRegClass(Reg: Arg0, RC: &SPIRV::iIDRegClass);
1472 GR->assignSPIRVTypeToVReg(Type: BoolType, VReg: Arg0, MF: MIRBuilder.getMF());
1473 MIRBuilder.buildICmp(
1474 Pred: CmpInst::ICMP_NE, Res: Arg0, Op0: BoolReg,
1475 Op1: GR->buildConstantInt(Val: 0, MIRBuilder, SpvType: BoolRegType, EmitIR: true));
1476 updateRegType(Reg: Arg0, Ty: nullptr, SpirvTy: BoolType, GR, MIB&: MIRBuilder,
1477 MRI&: MIRBuilder.getMF().getRegInfo());
1478 } else if (BoolRegType->getOpcode() != SPIRV::OpTypeBool) {
1479 report_fatal_error(reason: "Expect a boolean argument");
1480 }
1481 // if BoolReg is a boolean register, we don't need to do anything
1482 }
1483 }
1484
1485 Register GroupResultRegister = Call->ReturnRegister;
1486 SPIRVTypeInst GroupResultType = Call->ReturnType;
1487
1488 // TODO: maybe we need to check whether the result type is already boolean
1489 // and in this case do not insert select instruction.
1490 const bool HasBoolReturnTy =
1491 GroupBuiltin->IsElect || GroupBuiltin->IsAllOrAny ||
1492 GroupBuiltin->IsAllEqual || GroupBuiltin->IsLogical ||
1493 GroupBuiltin->IsInverseBallot || GroupBuiltin->IsBallotBitExtract;
1494
1495 if (HasBoolReturnTy)
1496 std::tie(args&: GroupResultRegister, args&: GroupResultType) =
1497 buildBoolRegister(MIRBuilder, ResultType: Call->ReturnType, GR);
1498
1499 auto Scope = Builtin->name().starts_with(Prefix: "sub_group")
1500 ? SPIRV::Scope::Subgroup
1501 : SPIRV::Scope::Workgroup;
1502 Register ScopeRegister = buildConstantIntReg32(Val: Scope, MIRBuilder, GR);
1503
1504 Register VecReg;
1505 if (GroupBuiltin->Opcode == SPIRV::OpGroupBroadcast &&
1506 Call->Arguments.size() > 2) {
1507 // For OpGroupBroadcast "LocalId must be an integer datatype. It must be a
1508 // scalar, a vector with 2 components, or a vector with 3 components.",
1509 // meaning that we must create a vector from the function arguments if
1510 // it's a work_group_broadcast(val, local_id_x, local_id_y) or
1511 // work_group_broadcast(val, local_id_x, local_id_y, local_id_z) call.
1512 Register ElemReg = Call->Arguments[1];
1513 SPIRVTypeInst ElemType = GR->getSPIRVTypeForVReg(VReg: ElemReg);
1514 if (!ElemType || ElemType->getOpcode() != SPIRV::OpTypeInt)
1515 report_fatal_error(reason: "Expect an integer <LocalId> argument");
1516 unsigned VecLen = Call->Arguments.size() - 1;
1517 VecReg = MRI->createGenericVirtualRegister(
1518 Ty: LLT::fixed_vector(NumElements: VecLen, ScalarTy: MRI->getType(Reg: ElemReg)));
1519 MRI->setRegClass(Reg: VecReg, RC: &SPIRV::viIDRegClass);
1520 SPIRVTypeInst VecType =
1521 GR->getOrCreateSPIRVVectorType(BaseType: ElemType, NumElements: VecLen, MIRBuilder, EmitIR: true);
1522 GR->assignSPIRVTypeToVReg(Type: VecType, VReg: VecReg, MF: MIRBuilder.getMF());
1523 auto MIB =
1524 MIRBuilder.buildInstr(Opcode: TargetOpcode::G_BUILD_VECTOR).addDef(RegNo: VecReg);
1525 for (unsigned i = 1; i < Call->Arguments.size(); i++) {
1526 MIB.addUse(RegNo: Call->Arguments[i]);
1527 setRegClassIfNull(Reg: Call->Arguments[i], MRI, GR);
1528 }
1529 updateRegType(Reg: VecReg, Ty: nullptr, SpirvTy: VecType, GR, MIB&: MIRBuilder,
1530 MRI&: MIRBuilder.getMF().getRegInfo());
1531 }
1532
1533 // Build work/sub group instruction.
1534 auto MIB = MIRBuilder.buildInstr(Opcode: GroupBuiltin->Opcode)
1535 .addDef(RegNo: GroupResultRegister)
1536 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: GroupResultType))
1537 .addUse(RegNo: ScopeRegister);
1538
1539 if (!GroupBuiltin->NoGroupOperation)
1540 MIB.addImm(Val: GroupBuiltin->GroupOperation);
1541 if (Call->Arguments.size() > 0) {
1542 MIB.addUse(RegNo: Arg0.isValid() ? Arg0 : Call->Arguments[0]);
1543 setRegClassIfNull(Reg: Call->Arguments[0], MRI, GR);
1544 if (VecReg.isValid())
1545 MIB.addUse(RegNo: VecReg);
1546 else
1547 for (unsigned i = 1; i < Call->Arguments.size(); i++)
1548 MIB.addUse(RegNo: Call->Arguments[i]);
1549 }
1550
1551 // Build select instruction.
1552 if (HasBoolReturnTy)
1553 buildSelectInst(MIRBuilder, ReturnRegister: Call->ReturnRegister, SourceRegister: GroupResultRegister,
1554 ReturnType: Call->ReturnType, GR);
1555 return true;
1556}
1557
1558static bool generateIntelSubgroupsInst(const SPIRV::IncomingCall *Call,
1559 MachineIRBuilder &MIRBuilder,
1560 SPIRVGlobalRegistry *GR) {
1561 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
1562 MachineFunction &MF = MIRBuilder.getMF();
1563 const auto *ST = static_cast<const SPIRVSubtarget *>(&MF.getSubtarget());
1564 const SPIRV::IntelSubgroupsBuiltin *IntelSubgroups =
1565 SPIRV::lookupIntelSubgroupsBuiltin(Name: Builtin->name());
1566
1567 if (IntelSubgroups->IsMedia &&
1568 !ST->canUseExtension(E: SPIRV::Extension::SPV_INTEL_media_block_io)) {
1569 std::string DiagMsg = std::string(Builtin->name()) +
1570 ": the builtin requires the following SPIR-V "
1571 "extension: SPV_INTEL_media_block_io";
1572 report_fatal_error(reason: DiagMsg.c_str(), gen_crash_diag: false);
1573 } else if (!IntelSubgroups->IsMedia &&
1574 !ST->canUseExtension(E: SPIRV::Extension::SPV_INTEL_subgroups)) {
1575 std::string DiagMsg = std::string(Builtin->name()) +
1576 ": the builtin requires the following SPIR-V "
1577 "extension: SPV_INTEL_subgroups";
1578 report_fatal_error(reason: DiagMsg.c_str(), gen_crash_diag: false);
1579 }
1580
1581 uint32_t OpCode = IntelSubgroups->Opcode;
1582 if (Call->isSpirvOp()) {
1583 bool IsSet = OpCode != SPIRV::OpSubgroupBlockWriteINTEL &&
1584 OpCode != SPIRV::OpSubgroupImageBlockWriteINTEL &&
1585 OpCode != SPIRV::OpSubgroupImageMediaBlockWriteINTEL;
1586 return buildOpFromWrapper(MIRBuilder, Opcode: OpCode, Call,
1587 TypeReg: IsSet ? GR->getSPIRVTypeID(SpirvType: Call->ReturnType)
1588 : Register(0));
1589 }
1590
1591 if (IntelSubgroups->IsBlock) {
1592 // Minimal number or arguments set in TableGen records is 1
1593 if (SPIRVTypeInst Arg0Type = GR->getSPIRVTypeForVReg(VReg: Call->Arguments[0])) {
1594 if (Arg0Type->getOpcode() == SPIRV::OpTypeImage) {
1595 // TODO: add required validation from the specification:
1596 // "'Image' must be an object whose type is OpTypeImage with a 'Sampled'
1597 // operand of 0 or 2. If the 'Sampled' operand is 2, then some
1598 // dimensions require a capability."
1599 switch (OpCode) {
1600 case SPIRV::OpSubgroupBlockReadINTEL:
1601 OpCode = SPIRV::OpSubgroupImageBlockReadINTEL;
1602 break;
1603 case SPIRV::OpSubgroupBlockWriteINTEL:
1604 OpCode = SPIRV::OpSubgroupImageBlockWriteINTEL;
1605 break;
1606 }
1607 }
1608 }
1609 }
1610
1611 // TODO: opaque pointers types should be eventually resolved in such a way
1612 // that validation of block read is enabled with respect to the following
1613 // specification requirement:
1614 // "'Result Type' may be a scalar or vector type, and its component type must
1615 // be equal to the type pointed to by 'Ptr'."
1616 // For example, function parameter type should not be default i8 pointer, but
1617 // depend on the result type of the instruction where it is used as a pointer
1618 // argument of OpSubgroupBlockReadINTEL
1619
1620 // Build Intel subgroups instruction
1621 MachineInstrBuilder MIB =
1622 IntelSubgroups->IsWrite
1623 ? MIRBuilder.buildInstr(Opcode: OpCode)
1624 : MIRBuilder.buildInstr(Opcode: OpCode)
1625 .addDef(RegNo: Call->ReturnRegister)
1626 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
1627 for (size_t i = 0; i < Call->Arguments.size(); ++i)
1628 MIB.addUse(RegNo: Call->Arguments[i]);
1629 return true;
1630}
1631
1632static bool generateGroupUniformInst(const SPIRV::IncomingCall *Call,
1633 MachineIRBuilder &MIRBuilder,
1634 SPIRVGlobalRegistry *GR) {
1635 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
1636 MachineFunction &MF = MIRBuilder.getMF();
1637 const auto *ST = static_cast<const SPIRVSubtarget *>(&MF.getSubtarget());
1638 if (!ST->canUseExtension(
1639 E: SPIRV::Extension::SPV_KHR_uniform_group_instructions)) {
1640 std::string DiagMsg = std::string(Builtin->name()) +
1641 ": the builtin requires the following SPIR-V "
1642 "extension: SPV_KHR_uniform_group_instructions";
1643 report_fatal_error(reason: DiagMsg.c_str(), gen_crash_diag: false);
1644 }
1645 const SPIRV::GroupUniformBuiltin *GroupUniform =
1646 SPIRV::lookupGroupUniformBuiltin(Name: Builtin->name());
1647 MachineRegisterInfo *MRI = MIRBuilder.getMRI();
1648
1649 Register GroupResultReg = Call->ReturnRegister;
1650 Register ScopeReg = Call->Arguments[0];
1651 Register ValueReg = Call->Arguments[2];
1652
1653 // Group Operation
1654 Register ConstGroupOpReg = Call->Arguments[1];
1655 const MachineInstr *Const = getDefInstrMaybeConstant(ConstReg&: ConstGroupOpReg, MRI);
1656 if (!Const || Const->getOpcode() != TargetOpcode::G_CONSTANT)
1657 report_fatal_error(
1658 reason: "expect a constant group operation for a uniform group instruction",
1659 gen_crash_diag: false);
1660 const MachineOperand &ConstOperand = Const->getOperand(i: 1);
1661 if (!ConstOperand.isCImm())
1662 report_fatal_error(reason: "uniform group instructions: group operation must be an "
1663 "integer constant",
1664 gen_crash_diag: false);
1665
1666 auto MIB = MIRBuilder.buildInstr(Opcode: GroupUniform->Opcode)
1667 .addDef(RegNo: GroupResultReg)
1668 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType))
1669 .addUse(RegNo: ScopeReg);
1670 addNumImm(Imm: ConstOperand.getCImm()->getValue(), MIB);
1671 MIB.addUse(RegNo: ValueReg);
1672
1673 return true;
1674}
1675
1676static bool generateKernelClockInst(const SPIRV::IncomingCall *Call,
1677 MachineIRBuilder &MIRBuilder,
1678 SPIRVGlobalRegistry *GR) {
1679 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
1680 MachineFunction &MF = MIRBuilder.getMF();
1681 const auto *ST = static_cast<const SPIRVSubtarget *>(&MF.getSubtarget());
1682 if (!ST->canUseExtension(E: SPIRV::Extension::SPV_KHR_shader_clock)) {
1683 std::string DiagMsg = std::string(Builtin->name()) +
1684 ": the builtin requires the following SPIR-V "
1685 "extension: SPV_KHR_shader_clock";
1686 report_fatal_error(reason: DiagMsg.c_str(), gen_crash_diag: false);
1687 }
1688
1689 Register ResultReg = Call->ReturnRegister;
1690
1691 if (Builtin->name() == "__spirv_ReadClockKHR") {
1692 MIRBuilder.buildInstr(Opcode: SPIRV::OpReadClockKHR)
1693 .addDef(RegNo: ResultReg)
1694 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType))
1695 .addUse(RegNo: Call->Arguments[0]);
1696 } else {
1697 // Deduce the `Scope` operand from the builtin function name.
1698 SPIRV::Scope::Scope ScopeArg =
1699 StringSwitch<SPIRV::Scope::Scope>(Builtin->name())
1700 .EndsWith(S: "device", Value: SPIRV::Scope::Scope::Device)
1701 .EndsWith(S: "work_group", Value: SPIRV::Scope::Scope::Workgroup)
1702 .EndsWith(S: "sub_group", Value: SPIRV::Scope::Scope::Subgroup);
1703 Register ScopeReg = buildConstantIntReg32(Val: ScopeArg, MIRBuilder, GR);
1704
1705 MIRBuilder.buildInstr(Opcode: SPIRV::OpReadClockKHR)
1706 .addDef(RegNo: ResultReg)
1707 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType))
1708 .addUse(RegNo: ScopeReg);
1709 }
1710
1711 return true;
1712}
1713
1714// These queries ask for a single size_t result for a given dimension index,
1715// e.g. size_t get_global_id(uint dimindex). In SPIR-V, the builtins
1716// corresponding to these values are all vec3 types, so we need to extract the
1717// correct index or return DefaultValue (0 or 1 depending on the query). We also
1718// handle extending or truncating in case size_t does not match the expected
1719// result type's bitwidth.
1720//
1721// For a constant index >= 3 we generate:
1722// %res = OpConstant %SizeT DefaultValue
1723//
1724// For other indices we generate:
1725// %g = OpVariable %ptr_V3_SizeT Input
1726// OpDecorate %g BuiltIn XXX
1727// OpDecorate %g LinkageAttributes "__spirv_BuiltInXXX"
1728// OpDecorate %g Constant
1729// %loadedVec = OpLoad %V3_SizeT %g
1730//
1731// Then, if the index is constant < 3, we generate:
1732// %res = OpCompositeExtract %SizeT %loadedVec idx
1733// If the index is dynamic, we generate:
1734// %tmp = OpVectorExtractDynamic %SizeT %loadedVec %idx
1735// %cmp = OpULessThan %bool %idx %const_3
1736// %res = OpSelect %SizeT %cmp %tmp %const_<DefaultValue>
1737//
1738// If the bitwidth of %res does not match the expected return type, we add an
1739// extend or truncate.
1740static bool genWorkgroupQuery(const SPIRV::IncomingCall *Call,
1741 MachineIRBuilder &MIRBuilder,
1742 SPIRVGlobalRegistry *GR,
1743 SPIRV::BuiltIn::BuiltIn BuiltinValue,
1744 uint64_t DefaultValue) {
1745 Register IndexRegister = Call->Arguments[0];
1746 const unsigned ResultWidth = Call->ReturnType->getOperand(i: 1).getImm();
1747 const unsigned PointerSize = GR->getPointerSize();
1748 const SPIRVTypeInst PointerSizeType =
1749 GR->getOrCreateSPIRVIntegerType(BitWidth: PointerSize, MIRBuilder);
1750 MachineRegisterInfo *MRI = MIRBuilder.getMRI();
1751 auto IndexInstruction = getDefInstrMaybeConstant(ConstReg&: IndexRegister, MRI);
1752
1753 // Set up the final register to do truncation or extension on at the end.
1754 Register ToTruncate = Call->ReturnRegister;
1755
1756 // If the index is constant, we can statically determine if it is in range.
1757 bool IsConstantIndex =
1758 IndexInstruction->getOpcode() == TargetOpcode::G_CONSTANT;
1759
1760 // If it's out of range (max dimension is 3), we can just return the constant
1761 // default value (0 or 1 depending on which query function).
1762 if (IsConstantIndex && getIConstVal(ConstReg: IndexRegister, MRI) >= 3) {
1763 Register DefaultReg = Call->ReturnRegister;
1764 if (PointerSize != ResultWidth) {
1765 DefaultReg = MRI->createGenericVirtualRegister(Ty: LLT::scalar(SizeInBits: PointerSize));
1766 MRI->setRegClass(Reg: DefaultReg, RC: &SPIRV::iIDRegClass);
1767 GR->assignSPIRVTypeToVReg(Type: PointerSizeType, VReg: DefaultReg,
1768 MF: MIRBuilder.getMF());
1769 ToTruncate = DefaultReg;
1770 }
1771 auto NewRegister =
1772 GR->buildConstantInt(Val: DefaultValue, MIRBuilder, SpvType: PointerSizeType, EmitIR: true);
1773 MIRBuilder.buildCopy(Res: DefaultReg, Op: NewRegister);
1774 } else { // If it could be in range, we need to load from the given builtin.
1775 auto Vec3Ty =
1776 GR->getOrCreateSPIRVVectorType(BaseType: PointerSizeType, NumElements: 3, MIRBuilder, EmitIR: true);
1777 Register LoadedVector =
1778 buildBuiltinVariableLoad(MIRBuilder, VariableType: Vec3Ty, GR, BuiltinValue,
1779 LLType: LLT::fixed_vector(NumElements: 3, ScalarSizeInBits: PointerSize));
1780 // Set up the vreg to extract the result to (possibly a new temporary one).
1781 Register Extracted = Call->ReturnRegister;
1782 if (!IsConstantIndex || PointerSize != ResultWidth) {
1783 Extracted = MRI->createGenericVirtualRegister(Ty: LLT::scalar(SizeInBits: PointerSize));
1784 MRI->setRegClass(Reg: Extracted, RC: &SPIRV::iIDRegClass);
1785 GR->assignSPIRVTypeToVReg(Type: PointerSizeType, VReg: Extracted, MF: MIRBuilder.getMF());
1786 }
1787 // Use Intrinsic::spv_extractelt so dynamic vs static extraction is
1788 // handled later: extr = spv_extractelt LoadedVector, IndexRegister.
1789 MachineInstrBuilder ExtractInst = MIRBuilder.buildIntrinsic(
1790 ID: Intrinsic::spv_extractelt, Res: ArrayRef<Register>{Extracted}, HasSideEffects: true, isConvergent: false);
1791 ExtractInst.addUse(RegNo: LoadedVector).addUse(RegNo: IndexRegister);
1792
1793 // If the index is dynamic, need check if it's < 3, and then use a select.
1794 if (!IsConstantIndex) {
1795 updateRegType(Reg: Extracted, Ty: nullptr, SpirvTy: PointerSizeType, GR, MIB&: MIRBuilder, MRI&: *MRI);
1796
1797 auto IndexType = GR->getSPIRVTypeForVReg(VReg: IndexRegister);
1798 auto BoolType = GR->getOrCreateSPIRVBoolType(MIRBuilder, EmitIR: true);
1799
1800 Register CompareRegister =
1801 MRI->createGenericVirtualRegister(Ty: LLT::scalar(SizeInBits: 1));
1802 MRI->setRegClass(Reg: CompareRegister, RC: &SPIRV::iIDRegClass);
1803 GR->assignSPIRVTypeToVReg(Type: BoolType, VReg: CompareRegister, MF: MIRBuilder.getMF());
1804
1805 // Use G_ICMP to check if idxVReg < 3.
1806 MIRBuilder.buildICmp(
1807 Pred: CmpInst::ICMP_ULT, Res: CompareRegister, Op0: IndexRegister,
1808 Op1: GR->buildConstantInt(Val: 3, MIRBuilder, SpvType: IndexType, EmitIR: true));
1809
1810 // Get constant for the default value (0 or 1 depending on which
1811 // function).
1812 Register DefaultRegister =
1813 GR->buildConstantInt(Val: DefaultValue, MIRBuilder, SpvType: PointerSizeType, EmitIR: true);
1814
1815 // Get a register for the selection result (possibly a new temporary one).
1816 Register SelectionResult = Call->ReturnRegister;
1817 if (PointerSize != ResultWidth) {
1818 SelectionResult =
1819 MRI->createGenericVirtualRegister(Ty: LLT::scalar(SizeInBits: PointerSize));
1820 MRI->setRegClass(Reg: SelectionResult, RC: &SPIRV::iIDRegClass);
1821 GR->assignSPIRVTypeToVReg(Type: PointerSizeType, VReg: SelectionResult,
1822 MF: MIRBuilder.getMF());
1823 }
1824 // Create the final G_SELECT to return the extracted value or the default.
1825 MIRBuilder.buildSelect(Res: SelectionResult, Tst: CompareRegister, Op0: Extracted,
1826 Op1: DefaultRegister);
1827 ToTruncate = SelectionResult;
1828 } else {
1829 ToTruncate = Extracted;
1830 }
1831 }
1832 // Alter the result's bitwidth if it does not match the SizeT value extracted.
1833 if (PointerSize != ResultWidth)
1834 MIRBuilder.buildZExtOrTrunc(Res: Call->ReturnRegister, Op: ToTruncate);
1835 return true;
1836}
1837
1838static bool generateBuiltinVar(const SPIRV::IncomingCall *Call,
1839 MachineIRBuilder &MIRBuilder,
1840 SPIRVGlobalRegistry *GR) {
1841 // Lookup the builtin variable record.
1842 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
1843 SPIRV::BuiltIn::BuiltIn Value =
1844 SPIRV::lookupGetBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Value;
1845
1846 if (Value == SPIRV::BuiltIn::GlobalInvocationId)
1847 return genWorkgroupQuery(Call, MIRBuilder, GR, BuiltinValue: Value, DefaultValue: 0);
1848
1849 // Build a load instruction for the builtin variable.
1850 unsigned BitWidth = GR->getScalarOrVectorBitWidth(Type: Call->ReturnType);
1851 LLT LLType;
1852 if (isVectorType(SPVTy: Call->ReturnType))
1853 LLType = LLT::fixed_vector(
1854 NumElements: GR->getScalarOrVectorComponentCount(Type: Call->ReturnType), ScalarSizeInBits: BitWidth);
1855 else
1856 LLType = LLT::scalar(SizeInBits: BitWidth);
1857
1858 return buildBuiltinVariableLoad(MIRBuilder, VariableType: Call->ReturnType, GR, BuiltinValue: Value,
1859 LLType, Reg: Call->ReturnRegister);
1860}
1861
1862static bool generateAtomicInst(const SPIRV::IncomingCall *Call,
1863 MachineIRBuilder &MIRBuilder,
1864 SPIRVGlobalRegistry *GR) {
1865 // Lookup the instruction opcode in the TableGen records.
1866 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
1867 unsigned Opcode =
1868 SPIRV::lookupNativeBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Opcode;
1869
1870 switch (Opcode) {
1871 case SPIRV::OpStore:
1872 return buildAtomicInitInst(Call, MIRBuilder);
1873 case SPIRV::OpAtomicLoad:
1874 return buildAtomicLoadInst(Call, MIRBuilder, GR);
1875 case SPIRV::OpAtomicStore:
1876 return buildAtomicStoreInst(Call, MIRBuilder, GR);
1877 case SPIRV::OpAtomicCompareExchange:
1878 case SPIRV::OpAtomicCompareExchangeWeak:
1879 return buildAtomicCompareExchangeInst(Call, Builtin, Opcode, MIRBuilder,
1880 GR);
1881 case SPIRV::OpAtomicIAdd:
1882 case SPIRV::OpAtomicISub:
1883 case SPIRV::OpAtomicOr:
1884 case SPIRV::OpAtomicXor:
1885 case SPIRV::OpAtomicAnd:
1886 case SPIRV::OpAtomicExchange:
1887 case SPIRV::OpAtomicSMax:
1888 case SPIRV::OpAtomicSMin:
1889 case SPIRV::OpAtomicUMax:
1890 case SPIRV::OpAtomicUMin:
1891 return buildAtomicRMWInst(Call, Opcode, MIRBuilder, GR);
1892 case SPIRV::OpMemoryBarrier:
1893 return buildBarrierInst(Call, Opcode: SPIRV::OpMemoryBarrier, MIRBuilder, GR);
1894 case SPIRV::OpAtomicFlagTestAndSet:
1895 case SPIRV::OpAtomicFlagClear:
1896 return buildAtomicFlagInst(Call, Opcode, MIRBuilder, GR);
1897 default:
1898 if (Call->isSpirvOp())
1899 return buildOpFromWrapper(MIRBuilder, Opcode, Call,
1900 TypeReg: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
1901 return false;
1902 }
1903}
1904
1905static bool generateAtomicFloatingInst(const SPIRV::IncomingCall *Call,
1906 MachineIRBuilder &MIRBuilder,
1907 SPIRVGlobalRegistry *GR) {
1908 // Lookup the instruction opcode in the TableGen records.
1909 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
1910 unsigned Opcode = SPIRV::lookupAtomicFloatingBuiltin(Name: Builtin->name())->Opcode;
1911
1912 switch (Opcode) {
1913 case SPIRV::OpAtomicFAddEXT:
1914 case SPIRV::OpAtomicFMinEXT:
1915 case SPIRV::OpAtomicFMaxEXT:
1916 return buildAtomicFloatingRMWInst(Call, Opcode, MIRBuilder, GR);
1917 default:
1918 return false;
1919 }
1920}
1921
1922static bool generateBarrierInst(const SPIRV::IncomingCall *Call,
1923 MachineIRBuilder &MIRBuilder,
1924 SPIRVGlobalRegistry *GR) {
1925 // Lookup the instruction opcode in the TableGen records.
1926 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
1927 unsigned Opcode =
1928 SPIRV::lookupNativeBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Opcode;
1929
1930 return buildBarrierInst(Call, Opcode, MIRBuilder, GR);
1931}
1932
1933static bool generateCastToPtrInst(const SPIRV::IncomingCall *Call,
1934 MachineIRBuilder &MIRBuilder,
1935 SPIRVGlobalRegistry *GR) {
1936 // Lookup the instruction opcode in the TableGen records.
1937 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
1938 unsigned Opcode =
1939 SPIRV::lookupNativeBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Opcode;
1940
1941 if (Opcode == SPIRV::OpGenericCastToPtrExplicit) {
1942 SPIRV::StorageClass::StorageClass ResSC =
1943 GR->getPointerStorageClass(VReg: Call->ReturnRegister);
1944 if (!isGenericCastablePtr(SC: ResSC))
1945 return false;
1946
1947 MIRBuilder.buildInstr(Opcode)
1948 .addDef(RegNo: Call->ReturnRegister)
1949 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType))
1950 .addUse(RegNo: Call->Arguments[0])
1951 .addImm(Val: ResSC);
1952 } else if (Opcode == SPIRV::OpGenericPtrMemSemantics) {
1953 if (GR->getPointerStorageClass(VReg: Call->Arguments[0]) !=
1954 SPIRV::StorageClass::Generic)
1955 return false;
1956
1957 // Shift the WorkgroupMemory/CrossWorkgroupMemory bits down to
1958 // CLK_LOCAL/GLOBAL_MEM_FENCE.
1959 MachineRegisterInfo *MRI = MIRBuilder.getMRI();
1960 SPIRVTypeInst RetTy = Call->ReturnType;
1961 Register SemReg =
1962 MRI->createGenericVirtualRegister(Ty: MRI->getType(Reg: Call->ReturnRegister));
1963 MRI->setRegClass(Reg: SemReg, RC: GR->getRegClass(SpvType: RetTy));
1964 GR->assignSPIRVTypeToVReg(Type: RetTy, VReg: SemReg, MF: MIRBuilder.getMF());
1965 MIRBuilder.buildInstr(Opcode)
1966 .addDef(RegNo: SemReg)
1967 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: RetTy))
1968 .addUse(RegNo: Call->Arguments[0]);
1969 Register ShiftReg =
1970 GR->buildConstantInt(Val: 8, MIRBuilder, SpvType: RetTy, /*EmitIR=*/true);
1971 MIRBuilder.buildInstr(Opcode: TargetOpcode::G_LSHR)
1972 .addDef(RegNo: Call->ReturnRegister)
1973 .addUse(RegNo: SemReg)
1974 .addUse(RegNo: ShiftReg);
1975 } else {
1976 MIRBuilder.buildInstr(Opcode: TargetOpcode::G_ADDRSPACE_CAST)
1977 .addDef(RegNo: Call->ReturnRegister)
1978 .addUse(RegNo: Call->Arguments[0]);
1979 }
1980 return true;
1981}
1982
1983static bool generateDotOrFMulInst(StringRef DemangledCall,
1984 const SPIRV::IncomingCall *Call,
1985 MachineIRBuilder &MIRBuilder,
1986 SPIRVGlobalRegistry *GR) {
1987 if (Call->isSpirvOp())
1988 return buildOpFromWrapper(MIRBuilder, Opcode: SPIRV::OpDot, Call,
1989 TypeReg: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
1990
1991 // Use OpDot only in case of vector args and OpFMul in case of scalar args.
1992 bool IsVec = isVectorType(SPVTy: GR->getSPIRVTypeForVReg(VReg: Call->Arguments[0]));
1993 uint32_t OC = IsVec ? SPIRV::OpDot : SPIRV::OpFMulS;
1994 bool IsSwapReq = false;
1995
1996 const auto *ST =
1997 static_cast<const SPIRVSubtarget *>(&MIRBuilder.getMF().getSubtarget());
1998 if (GR->isScalarOrVectorOfType(VReg: Call->ReturnRegister, TypeOpcode: SPIRV::OpTypeInt)) {
1999 if (!ST->canUseExtension(E: SPIRV::Extension::SPV_KHR_integer_dot_product) &&
2000 !ST->isAtLeastSPIRVVer(VerToCompareTo: VersionTuple(1, 6)))
2001 report_fatal_error(reason: Twine(Call->Builtin->name()) +
2002 ": the builtin requires the following SPIR-V "
2003 "extension: SPV_KHR_integer_dot_product",
2004 gen_crash_diag: false);
2005 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
2006 const SPIRV::IntegerDotProductBuiltin *IntDot =
2007 SPIRV::lookupIntegerDotProductBuiltin(Name: Builtin->name());
2008 if (IntDot) {
2009 OC = IntDot->Opcode;
2010 IsSwapReq = IntDot->IsSwapReq;
2011 } else if (IsVec) {
2012 // Handling "dot" and "dot_acc_sat" builtins which use vectors of
2013 // integers.
2014 LLVMContext &Ctx = MIRBuilder.getContext();
2015 SmallVector<StringRef, 10> TypeStrs;
2016 SPIRV::parseBuiltinTypeStr(BuiltinArgsTypeStrs&: TypeStrs, DemangledCall, Ctx);
2017 bool IsFirstSigned = TypeStrs[0].trim()[0] != 'u';
2018 bool IsSecondSigned = TypeStrs[1].trim()[0] != 'u';
2019
2020 if (Call->BuiltinName == "dot") {
2021 if (IsFirstSigned && IsSecondSigned)
2022 OC = SPIRV::OpSDot;
2023 else if (!IsFirstSigned && !IsSecondSigned)
2024 OC = SPIRV::OpUDot;
2025 else {
2026 OC = SPIRV::OpSUDot;
2027 if (!IsFirstSigned)
2028 IsSwapReq = true;
2029 }
2030 } else if (Call->BuiltinName == "dot_acc_sat") {
2031 if (IsFirstSigned && IsSecondSigned)
2032 OC = SPIRV::OpSDotAccSat;
2033 else if (!IsFirstSigned && !IsSecondSigned)
2034 OC = SPIRV::OpUDotAccSat;
2035 else {
2036 OC = SPIRV::OpSUDotAccSat;
2037 if (!IsFirstSigned)
2038 IsSwapReq = true;
2039 }
2040 }
2041 }
2042 }
2043
2044 MachineInstrBuilder MIB = MIRBuilder.buildInstr(Opcode: OC)
2045 .addDef(RegNo: Call->ReturnRegister)
2046 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
2047
2048 if (IsSwapReq) {
2049 MIB.addUse(RegNo: Call->Arguments[1]);
2050 MIB.addUse(RegNo: Call->Arguments[0]);
2051 // needed for dot_acc_sat* builtins
2052 for (size_t i = 2; i < Call->Arguments.size(); ++i)
2053 MIB.addUse(RegNo: Call->Arguments[i]);
2054 } else {
2055 for (size_t i = 0; i < Call->Arguments.size(); ++i)
2056 MIB.addUse(RegNo: Call->Arguments[i]);
2057 }
2058
2059 // Add Packed Vector Format for Integer dot product builtins if arguments are
2060 // scalar
2061 if (!IsVec && OC != SPIRV::OpFMulS)
2062 MIB.addImm(Val: SPIRV::PackedVectorFormat4x8Bit);
2063
2064 return true;
2065}
2066
2067static bool generateWaveInst(const SPIRV::IncomingCall *Call,
2068 MachineIRBuilder &MIRBuilder,
2069 SPIRVGlobalRegistry *GR) {
2070 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
2071 SPIRV::BuiltIn::BuiltIn Value =
2072 SPIRV::lookupGetBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Value;
2073
2074 // For now, we only support a single Wave intrinsic with a single return type.
2075 assert(Call->ReturnType->getOpcode() == SPIRV::OpTypeInt);
2076 LLT LLType = LLT::scalar(SizeInBits: GR->getScalarOrVectorBitWidth(Type: Call->ReturnType));
2077
2078 return buildBuiltinVariableLoad(
2079 MIRBuilder, VariableType: Call->ReturnType, GR, BuiltinValue: Value, LLType, Reg: Call->ReturnRegister,
2080 /* isConst= */ false, /* LinkageType= */ LinkageTy: std::nullopt);
2081}
2082
2083// Build a SPIR-V instruction with struct return via sret pointer:
2084// Res = Opcode RetType Op1 Op2
2085// OpStore SRetReg Res
2086static void buildSRetInst(unsigned Opcode, Register SRetReg, Register Op1Reg,
2087 Register Op2Reg, SPIRVTypeInst RetType,
2088 MachineIRBuilder &MIRBuilder,
2089 SPIRVGlobalRegistry *GR) {
2090 MachineRegisterInfo *MRI = MIRBuilder.getMRI();
2091 Register ResReg = MRI->createVirtualRegister(RegClass: &SPIRV::iIDRegClass);
2092 if (const TargetRegisterClass *DstRC = MRI->getRegClassOrNull(Reg: Op1Reg)) {
2093 MRI->setRegClass(Reg: ResReg, RC: DstRC);
2094 MRI->setType(VReg: ResReg, Ty: MRI->getType(Reg: Op1Reg));
2095 }
2096 GR->assignSPIRVTypeToVReg(Type: RetType, VReg: ResReg, MF: MIRBuilder.getMF());
2097 MIRBuilder.buildInstr(Opcode)
2098 .addDef(RegNo: ResReg)
2099 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: RetType))
2100 .addUse(RegNo: Op1Reg)
2101 .addUse(RegNo: Op2Reg);
2102 MIRBuilder.buildInstr(Opcode: SPIRV::OpStore).addUse(RegNo: SRetReg).addUse(RegNo: ResReg);
2103}
2104
2105// Find the pointee type of an sret pointer argument. A typed pointer gives us
2106// the type directly. An untyped one does not, so fall back to the element type
2107// we deduced for the matching IR argument, or null if there is nothing to fall
2108// back to.
2109static SPIRVTypeInst deduceSRetPointeeType(Register SRetReg,
2110 const Value *SRetArg,
2111 MachineIRBuilder &MIRBuilder,
2112 SPIRVGlobalRegistry *GR) {
2113 SPIRVTypeInst RetType = GR->getPointeeType(PtrType: GR->getSPIRVTypeForVReg(VReg: SRetReg));
2114 if (!RetType)
2115 if (Type *ElemTy = GR->findDeducedElementType(Val: SRetArg))
2116 RetType = GR->getOrCreateSPIRVType(
2117 Type: ElemTy, MIRBuilder, AQ: SPIRV::AccessQualifier::ReadWrite, EmitIR: false);
2118 return RetType;
2119}
2120
2121// We expect a builtin
2122// Name(ptr sret([RetType]) %result, Type %operand1, Type %operand1)
2123// where %result is a pointer to where the result of the builtin execution
2124// is to be stored, and generate the following instructions:
2125// Res = Opcode RetType Operand1 Operand1
2126// OpStore RetVariable Res
2127static bool generateICarryBorrowInst(const SPIRV::IncomingCall *Call,
2128 MachineIRBuilder &MIRBuilder,
2129 SPIRVGlobalRegistry *GR,
2130 const CallBase &CB) {
2131 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
2132 unsigned Opcode =
2133 SPIRV::lookupNativeBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Opcode;
2134
2135 Register SRetReg = Call->Arguments[0];
2136 SPIRVTypeInst RetType =
2137 deduceSRetPointeeType(SRetReg, SRetArg: CB.getArgOperand(i: 0), MIRBuilder, GR);
2138 if (!RetType)
2139 report_fatal_error(reason: "The first parameter must be a pointer");
2140 if (RetType->getOpcode() != SPIRV::OpTypeStruct)
2141 report_fatal_error(reason: "Expected struct type result for the arithmetic with "
2142 "overflow builtins");
2143
2144 SPIRVTypeInst OpType1 = GR->getSPIRVTypeForVReg(VReg: Call->Arguments[1]);
2145 SPIRVTypeInst OpType2 = GR->getSPIRVTypeForVReg(VReg: Call->Arguments[2]);
2146 if (!OpType1 || !OpType2 || OpType1 != OpType2)
2147 report_fatal_error(reason: "Operands must have the same type");
2148 if (isVectorType(SPVTy: OpType1))
2149 switch (Opcode) {
2150 case SPIRV::OpIAddCarryS:
2151 Opcode = SPIRV::OpIAddCarryV;
2152 break;
2153 case SPIRV::OpISubBorrowS:
2154 Opcode = SPIRV::OpISubBorrowV;
2155 break;
2156 }
2157
2158 buildSRetInst(Opcode, SRetReg, Op1Reg: Call->Arguments[1], Op2Reg: Call->Arguments[2],
2159 RetType, MIRBuilder, GR);
2160 return true;
2161}
2162
2163// We expect a builtin in one of two forms:
2164//
2165// (1) sret convention (3 arguments):
2166// void Name(ptr sret([RetType]) %result, Type %operand1, Type %operand2)
2167// => Res = Opcode RetType Operand1 Operand2
2168// OpStore %result Res
2169//
2170// (2) direct return convention (2 arguments):
2171// RetType Name(Type %operand1, Type %operand2)
2172// => Res = Opcode RetType Operand1 Operand2
2173//
2174// RetType is a struct with two members of the same type as the operands.
2175static bool generateMulExtendedInst(const SPIRV::IncomingCall *Call,
2176 MachineIRBuilder &MIRBuilder,
2177 SPIRVGlobalRegistry *GR,
2178 const CallBase &CB) {
2179 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
2180 unsigned Opcode =
2181 SPIRV::lookupNativeBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Opcode;
2182 assert((Opcode == SPIRV::OpUMulExtended || Opcode == SPIRV::OpSMulExtended) &&
2183 "Expected OpUMulExtended or OpSMulExtended");
2184
2185 const bool IsSret =
2186 !Call->ReturnType || Call->ReturnType->getOpcode() == SPIRV::OpTypeVoid;
2187 Register Op1Reg = IsSret ? Call->Arguments[1] : Call->Arguments[0];
2188 Register Op2Reg = IsSret ? Call->Arguments[2] : Call->Arguments[1];
2189
2190 SPIRVTypeInst RetType = nullptr;
2191 if (IsSret) {
2192 Register SRetReg = Call->Arguments[0];
2193 RetType =
2194 deduceSRetPointeeType(SRetReg, SRetArg: CB.getArgOperand(i: 0), MIRBuilder, GR);
2195 if (!RetType)
2196 report_fatal_error(reason: "The first parameter must be a pointer");
2197 } else {
2198 RetType = Call->ReturnType;
2199 }
2200
2201 if (!RetType || RetType->getOpcode() != SPIRV::OpTypeStruct)
2202 report_fatal_error(reason: "Expected struct type result for the extended "
2203 "multiplication builtins");
2204 if (RetType->getNumOperands() != 3)
2205 report_fatal_error(reason: "Expected struct with exactly two members for the "
2206 "extended multiplication builtins");
2207 SPIRVTypeInst Member0Type =
2208 GR->getSPIRVTypeForVReg(VReg: RetType->getOperand(i: 1).getReg());
2209 SPIRVTypeInst Member1Type =
2210 GR->getSPIRVTypeForVReg(VReg: RetType->getOperand(i: 2).getReg());
2211 if (!Member0Type || !Member1Type || Member0Type != Member1Type)
2212 report_fatal_error(reason: "Both struct members must be the same type");
2213
2214 SPIRVTypeInst OpType1 = GR->getSPIRVTypeForVReg(VReg: Op1Reg);
2215 SPIRVTypeInst OpType2 = GR->getSPIRVTypeForVReg(VReg: Op2Reg);
2216 if (!OpType1 || !OpType2 || OpType1 != OpType2)
2217 report_fatal_error(reason: "Operands must have the same type");
2218 if (OpType1 != Member0Type)
2219 report_fatal_error(reason: "Operand type must match the struct member type");
2220
2221 if (IsSret) {
2222 buildSRetInst(Opcode, SRetReg: Call->Arguments[0], Op1Reg, Op2Reg, RetType,
2223 MIRBuilder, GR);
2224 } else {
2225 MachineRegisterInfo *MRI = MIRBuilder.getMRI();
2226 Register ResReg = Call->ReturnRegister;
2227 if (const TargetRegisterClass *DstRC = MRI->getRegClassOrNull(Reg: Op1Reg)) {
2228 MRI->setRegClass(Reg: ResReg, RC: DstRC);
2229 }
2230 GR->assignSPIRVTypeToVReg(Type: RetType, VReg: ResReg, MF: MIRBuilder.getMF());
2231 MIRBuilder.buildInstr(Opcode)
2232 .addDef(RegNo: ResReg)
2233 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: RetType))
2234 .addUse(RegNo: Op1Reg)
2235 .addUse(RegNo: Op2Reg);
2236 }
2237 return true;
2238}
2239
2240static bool generateArithmeticInst(const SPIRV::IncomingCall *Call,
2241 MachineIRBuilder &MIRBuilder,
2242 SPIRVGlobalRegistry *GR) {
2243 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
2244 unsigned Opcode =
2245 SPIRV::lookupNativeBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Opcode;
2246
2247 auto MIB = MIRBuilder.buildInstr(Opcode)
2248 .addDef(RegNo: Call->ReturnRegister)
2249 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
2250 for (Register Arg : Call->Arguments)
2251 MIB.addUse(RegNo: Arg);
2252 return true;
2253}
2254
2255static bool generateGetQueryInst(const SPIRV::IncomingCall *Call,
2256 MachineIRBuilder &MIRBuilder,
2257 SPIRVGlobalRegistry *GR) {
2258 // Lookup the builtin record.
2259 SPIRV::BuiltIn::BuiltIn Value =
2260 SPIRV::lookupGetBuiltin(Name: Call->Builtin->name(), Set: Call->Builtin->Set)->Value;
2261 const bool IsDefaultOne = (Value == SPIRV::BuiltIn::GlobalSize ||
2262 Value == SPIRV::BuiltIn::NumWorkgroups ||
2263 Value == SPIRV::BuiltIn::WorkgroupSize ||
2264 Value == SPIRV::BuiltIn::EnqueuedWorkgroupSize);
2265 return genWorkgroupQuery(Call, MIRBuilder, GR, BuiltinValue: Value, DefaultValue: IsDefaultOne ? 1 : 0);
2266}
2267
2268static bool generateImageSizeQueryInst(const SPIRV::IncomingCall *Call,
2269 MachineIRBuilder &MIRBuilder,
2270 SPIRVGlobalRegistry *GR) {
2271 // Lookup the image size query component number in the TableGen records.
2272 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
2273 uint32_t Component =
2274 SPIRV::lookupImageQueryBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Component;
2275 // Query result may either be a vector or a scalar. If return type is not a
2276 // vector, expect only a single size component. Otherwise get the number of
2277 // expected components.
2278 unsigned NumExpectedRetComponents =
2279 GR->getScalarOrVectorComponentCount(Type: Call->ReturnType);
2280 // Get the actual number of query result/size components.
2281 SPIRVTypeInst ImgType = GR->getSPIRVTypeForVReg(VReg: Call->Arguments[0]);
2282 unsigned NumActualRetComponents = getNumSizeComponents(imgType: ImgType);
2283 Register QueryResult = Call->ReturnRegister;
2284 SPIRVTypeInst QueryResultType = Call->ReturnType;
2285 if (NumExpectedRetComponents != NumActualRetComponents) {
2286 unsigned Bitwidth = Call->ReturnType->getOpcode() == SPIRV::OpTypeInt
2287 ? Call->ReturnType->getOperand(i: 1).getImm()
2288 : 32;
2289 QueryResult = MIRBuilder.getMRI()->createGenericVirtualRegister(
2290 Ty: LLT::fixed_vector(NumElements: NumActualRetComponents, ScalarSizeInBits: Bitwidth));
2291 MIRBuilder.getMRI()->setRegClass(Reg: QueryResult, RC: &SPIRV::viIDRegClass);
2292 SPIRVTypeInst IntTy = GR->getOrCreateSPIRVIntegerType(BitWidth: Bitwidth, MIRBuilder);
2293 QueryResultType = GR->getOrCreateSPIRVVectorType(
2294 BaseType: IntTy, NumElements: NumActualRetComponents, MIRBuilder, EmitIR: true);
2295 GR->assignSPIRVTypeToVReg(Type: QueryResultType, VReg: QueryResult, MF: MIRBuilder.getMF());
2296 }
2297 bool IsDimBuf = ImgType->getOperand(i: 2).getImm() == SPIRV::Dim::DIM_Buffer;
2298 bool IsMultisampled = ImgType->getOperand(i: 5).getImm() != 0;
2299 bool UseQuerySize = IsDimBuf || IsMultisampled;
2300 unsigned Opcode =
2301 UseQuerySize ? SPIRV::OpImageQuerySize : SPIRV::OpImageQuerySizeLod;
2302 auto MIB = MIRBuilder.buildInstr(Opcode)
2303 .addDef(RegNo: QueryResult)
2304 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: QueryResultType))
2305 .addUse(RegNo: Call->Arguments[0]);
2306 if (!UseQuerySize)
2307 MIB.addUse(RegNo: buildConstantIntReg32(Val: 0, MIRBuilder, GR)); // Lod id.
2308 if (NumExpectedRetComponents == NumActualRetComponents)
2309 return true;
2310 if (NumExpectedRetComponents == 1) {
2311 // Only 1 component is expected, build OpCompositeExtract instruction.
2312 unsigned ExtractedComposite =
2313 Component == 3 ? NumActualRetComponents - 1 : Component;
2314 assert(ExtractedComposite < NumActualRetComponents &&
2315 "Invalid composite index!");
2316 Register TypeReg = GR->getSPIRVTypeID(SpirvType: Call->ReturnType);
2317 SPIRVTypeInst NewType = nullptr;
2318 if (isVectorType(SPVTy: QueryResultType)) {
2319 NewType = GR->getScalarOrVectorComponentType(Type: QueryResultType);
2320 Register NewTypeReg = GR->getSPIRVTypeID(SpirvType: NewType);
2321 if (TypeReg != NewTypeReg)
2322 TypeReg = NewTypeReg;
2323 else
2324 NewType = nullptr;
2325 }
2326 MIRBuilder.buildInstr(Opcode: SPIRV::OpCompositeExtract)
2327 .addDef(RegNo: Call->ReturnRegister)
2328 .addUse(RegNo: TypeReg)
2329 .addUse(RegNo: QueryResult)
2330 .addImm(Val: ExtractedComposite);
2331 if (NewType)
2332 updateRegType(Reg: Call->ReturnRegister, Ty: nullptr, SpirvTy: NewType, GR, MIB&: MIRBuilder,
2333 MRI&: MIRBuilder.getMF().getRegInfo());
2334 } else {
2335 // More than 1 component is expected, fill a new vector.
2336 auto MIB = MIRBuilder.buildInstr(Opcode: SPIRV::OpVectorShuffle)
2337 .addDef(RegNo: Call->ReturnRegister)
2338 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType))
2339 .addUse(RegNo: QueryResult)
2340 .addUse(RegNo: QueryResult);
2341 for (unsigned i = 0; i < NumExpectedRetComponents; ++i)
2342 MIB.addImm(Val: i < NumActualRetComponents ? i : 0xffffffff);
2343 }
2344 return true;
2345}
2346
2347static bool generateImageMiscQueryInst(const SPIRV::IncomingCall *Call,
2348 MachineIRBuilder &MIRBuilder,
2349 SPIRVGlobalRegistry *GR) {
2350 assert(Call->ReturnType->getOpcode() == SPIRV::OpTypeInt &&
2351 "Image samples query result must be of int type!");
2352
2353 // Lookup the instruction opcode in the TableGen records.
2354 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
2355 unsigned Opcode =
2356 SPIRV::lookupNativeBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Opcode;
2357
2358 Register Image = Call->Arguments[0];
2359 SPIRV::Dim::Dim ImageDimensionality = static_cast<SPIRV::Dim::Dim>(
2360 GR->getSPIRVTypeForVReg(VReg: Image)->getOperand(i: 2).getImm());
2361 (void)ImageDimensionality;
2362
2363 switch (Opcode) {
2364 case SPIRV::OpImageQuerySamples:
2365 assert(ImageDimensionality == SPIRV::Dim::DIM_2D &&
2366 "Image must be of 2D dimensionality");
2367 break;
2368 case SPIRV::OpImageQueryLevels:
2369 assert((ImageDimensionality == SPIRV::Dim::DIM_1D ||
2370 ImageDimensionality == SPIRV::Dim::DIM_2D ||
2371 ImageDimensionality == SPIRV::Dim::DIM_3D ||
2372 ImageDimensionality == SPIRV::Dim::DIM_Cube) &&
2373 "Image must be of 1D/2D/3D/Cube dimensionality");
2374 break;
2375 }
2376
2377 MIRBuilder.buildInstr(Opcode)
2378 .addDef(RegNo: Call->ReturnRegister)
2379 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType))
2380 .addUse(RegNo: Image);
2381 return true;
2382}
2383
2384// TODO: Move to TableGen.
2385static SPIRV::SamplerAddressingMode::SamplerAddressingMode
2386getSamplerAddressingModeFromBitmask(unsigned Bitmask) {
2387 switch (Bitmask & SPIRV::CLK_ADDRESS_MODE_MASK) {
2388 case SPIRV::CLK_ADDRESS_CLAMP:
2389 return SPIRV::SamplerAddressingMode::Clamp;
2390 case SPIRV::CLK_ADDRESS_CLAMP_TO_EDGE:
2391 return SPIRV::SamplerAddressingMode::ClampToEdge;
2392 case SPIRV::CLK_ADDRESS_REPEAT:
2393 return SPIRV::SamplerAddressingMode::Repeat;
2394 case SPIRV::CLK_ADDRESS_MIRRORED_REPEAT:
2395 return SPIRV::SamplerAddressingMode::RepeatMirrored;
2396 case SPIRV::CLK_ADDRESS_NONE:
2397 return SPIRV::SamplerAddressingMode::None;
2398 default:
2399 report_fatal_error(reason: "Unknown CL address mode");
2400 }
2401}
2402
2403static unsigned getSamplerParamFromBitmask(unsigned Bitmask) {
2404 return (Bitmask & SPIRV::CLK_NORMALIZED_COORDS_TRUE) ? 1 : 0;
2405}
2406
2407static SPIRV::SamplerFilterMode::SamplerFilterMode
2408getSamplerFilterModeFromBitmask(unsigned Bitmask) {
2409 if (Bitmask & SPIRV::CLK_FILTER_LINEAR)
2410 return SPIRV::SamplerFilterMode::Linear;
2411 if (Bitmask & SPIRV::CLK_FILTER_NEAREST)
2412 return SPIRV::SamplerFilterMode::Nearest;
2413 return SPIRV::SamplerFilterMode::Nearest;
2414}
2415
2416static bool generateReadImageInst(StringRef DemangledCall,
2417 const SPIRV::IncomingCall *Call,
2418 MachineIRBuilder &MIRBuilder,
2419 SPIRVGlobalRegistry *GR) {
2420 if (Call->isSpirvOp())
2421 return buildOpFromWrapper(MIRBuilder, Opcode: SPIRV::OpImageRead, Call,
2422 TypeReg: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
2423 Register Image = Call->Arguments[0];
2424 MachineRegisterInfo *MRI = MIRBuilder.getMRI();
2425 bool HasOclSampler = DemangledCall.contains_insensitive(Other: "ocl_sampler");
2426 bool HasMsaa = DemangledCall.contains_insensitive(Other: "msaa");
2427 if (HasOclSampler) {
2428 Register Sampler = Call->Arguments[1];
2429
2430 if (!GR->isScalarOfType(VReg: Sampler, TypeOpcode: SPIRV::OpTypeSampler) &&
2431 getDefInstrMaybeConstant(ConstReg&: Sampler, MRI)->getOperand(i: 1).isCImm()) {
2432 uint64_t SamplerMask = getIConstVal(ConstReg: Sampler, MRI);
2433 Sampler = GR->buildConstantSampler(
2434 Res: Register(), AddrMode: getSamplerAddressingModeFromBitmask(Bitmask: SamplerMask),
2435 Param: getSamplerParamFromBitmask(Bitmask: SamplerMask),
2436 FilerMode: getSamplerFilterModeFromBitmask(Bitmask: SamplerMask), MIRBuilder);
2437 }
2438 SPIRVTypeInst ImageType = GR->getSPIRVTypeForVReg(VReg: Image);
2439 SPIRVTypeInst SampledImageType =
2440 GR->getOrCreateOpTypeSampledImage(ImageType, MIRBuilder);
2441 Register SampledImage = MRI->createVirtualRegister(RegClass: &SPIRV::iIDRegClass);
2442
2443 MIRBuilder.buildInstr(Opcode: SPIRV::OpSampledImage)
2444 .addDef(RegNo: SampledImage)
2445 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: SampledImageType))
2446 .addUse(RegNo: Image)
2447 .addUse(RegNo: Sampler);
2448
2449 Register Lod = GR->buildConstantFP(Val: APFloat::getZero(Sem: APFloat::IEEEsingle()),
2450 MIRBuilder);
2451
2452 if (!isVectorType(SPVTy: Call->ReturnType)) {
2453 SPIRVTypeInst TempType =
2454 GR->getOrCreateSPIRVVectorType(BaseType: Call->ReturnType, NumElements: 4, MIRBuilder, EmitIR: true);
2455 Register TempRegister =
2456 MRI->createGenericVirtualRegister(Ty: GR->getRegType(SpvType: TempType));
2457 MRI->setRegClass(Reg: TempRegister, RC: GR->getRegClass(SpvType: TempType));
2458 GR->assignSPIRVTypeToVReg(Type: TempType, VReg: TempRegister, MF: MIRBuilder.getMF());
2459 MIRBuilder.buildInstr(Opcode: SPIRV::OpImageSampleExplicitLod)
2460 .addDef(RegNo: TempRegister)
2461 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: TempType))
2462 .addUse(RegNo: SampledImage)
2463 .addUse(RegNo: Call->Arguments[2]) // Coordinate.
2464 .addImm(Val: SPIRV::ImageOperand::Lod)
2465 .addUse(RegNo: Lod);
2466 MIRBuilder.buildInstr(Opcode: SPIRV::OpCompositeExtract)
2467 .addDef(RegNo: Call->ReturnRegister)
2468 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType))
2469 .addUse(RegNo: TempRegister)
2470 .addImm(Val: 0);
2471 } else {
2472 MIRBuilder.buildInstr(Opcode: SPIRV::OpImageSampleExplicitLod)
2473 .addDef(RegNo: Call->ReturnRegister)
2474 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType))
2475 .addUse(RegNo: SampledImage)
2476 .addUse(RegNo: Call->Arguments[2]) // Coordinate.
2477 .addImm(Val: SPIRV::ImageOperand::Lod)
2478 .addUse(RegNo: Lod);
2479 }
2480 } else if (HasMsaa) {
2481 MIRBuilder.buildInstr(Opcode: SPIRV::OpImageRead)
2482 .addDef(RegNo: Call->ReturnRegister)
2483 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType))
2484 .addUse(RegNo: Image)
2485 .addUse(RegNo: Call->Arguments[1]) // Coordinate.
2486 .addImm(Val: SPIRV::ImageOperand::Sample)
2487 .addUse(RegNo: Call->Arguments[2]);
2488 } else {
2489 MIRBuilder.buildInstr(Opcode: SPIRV::OpImageRead)
2490 .addDef(RegNo: Call->ReturnRegister)
2491 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType))
2492 .addUse(RegNo: Image)
2493 .addUse(RegNo: Call->Arguments[1]); // Coordinate.
2494 }
2495 return true;
2496}
2497
2498static bool generateWriteImageInst(const SPIRV::IncomingCall *Call,
2499 MachineIRBuilder &MIRBuilder,
2500 SPIRVGlobalRegistry *GR) {
2501 if (Call->isSpirvOp())
2502 return buildOpFromWrapper(MIRBuilder, Opcode: SPIRV::OpImageWrite, Call,
2503 TypeReg: Register(0));
2504 MIRBuilder.buildInstr(Opcode: SPIRV::OpImageWrite)
2505 .addUse(RegNo: Call->Arguments[0]) // Image.
2506 .addUse(RegNo: Call->Arguments[1]) // Coordinate.
2507 .addUse(RegNo: Call->Arguments[2]); // Texel.
2508 return true;
2509}
2510
2511static bool generateSampleImageInst(StringRef DemangledCall,
2512 const SPIRV::IncomingCall *Call,
2513 MachineIRBuilder &MIRBuilder,
2514 SPIRVGlobalRegistry *GR) {
2515 MachineRegisterInfo *MRI = MIRBuilder.getMRI();
2516 if (Call->Builtin->name().contains_insensitive(
2517 Other: "__translate_sampler_initializer")) {
2518 // Build sampler literal.
2519 uint64_t Bitmask = getIConstVal(ConstReg: Call->Arguments[0], MRI);
2520 Register Sampler = GR->buildConstantSampler(
2521 Res: Call->ReturnRegister, AddrMode: getSamplerAddressingModeFromBitmask(Bitmask),
2522 Param: getSamplerParamFromBitmask(Bitmask),
2523 FilerMode: getSamplerFilterModeFromBitmask(Bitmask), MIRBuilder);
2524 return Sampler.isValid();
2525 } else if (Call->Builtin->name().contains_insensitive(
2526 Other: "__spirv_SampledImage")) {
2527 // Create OpSampledImage.
2528 Register Image = Call->Arguments[0];
2529 SPIRVTypeInst ImageType = GR->getSPIRVTypeForVReg(VReg: Image);
2530 SPIRVTypeInst SampledImageType =
2531 GR->getOrCreateOpTypeSampledImage(ImageType, MIRBuilder);
2532 Register SampledImage =
2533 Call->ReturnRegister.isValid()
2534 ? Call->ReturnRegister
2535 : MRI->createVirtualRegister(RegClass: &SPIRV::iIDRegClass);
2536 MIRBuilder.buildInstr(Opcode: SPIRV::OpSampledImage)
2537 .addDef(RegNo: SampledImage)
2538 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: SampledImageType))
2539 .addUse(RegNo: Image)
2540 .addUse(RegNo: Call->Arguments[1]); // Sampler.
2541 return true;
2542 } else if (Call->Builtin->name().contains_insensitive(
2543 Other: "__spirv_ImageSampleExplicitLod")) {
2544 // Sample an image using an explicit level of detail.
2545 std::string ReturnType = DemangledCall.str();
2546 if (DemangledCall.contains(Other: "_R")) {
2547 ReturnType = ReturnType.substr(pos: ReturnType.find(s: "_R") + 2);
2548 ReturnType = ReturnType.substr(pos: 0, n: ReturnType.find(c: '('));
2549 }
2550 SPIRVTypeInst Type = Call->ReturnType
2551 ? Call->ReturnType
2552 : SPIRVTypeInst(GR->getOrCreateSPIRVTypeByName(
2553 TypeStr: ReturnType, MIRBuilder, EmitIR: true));
2554 if (!Type) {
2555 std::string DiagMsg =
2556 "Unable to recognize SPIRV type name: " + ReturnType;
2557 report_fatal_error(reason: DiagMsg.c_str());
2558 }
2559 MIRBuilder.buildInstr(Opcode: SPIRV::OpImageSampleExplicitLod)
2560 .addDef(RegNo: Call->ReturnRegister)
2561 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Type))
2562 .addUse(RegNo: Call->Arguments[0]) // Image.
2563 .addUse(RegNo: Call->Arguments[1]) // Coordinate.
2564 .addImm(Val: SPIRV::ImageOperand::Lod)
2565 .addUse(RegNo: Call->Arguments[3]);
2566 return true;
2567 }
2568 return false;
2569}
2570
2571static bool generateSelectInst(const SPIRV::IncomingCall *Call,
2572 MachineIRBuilder &MIRBuilder) {
2573 const MachineRegisterInfo *MRI = MIRBuilder.getMRI();
2574 LLT ResTy = MRI->getType(Reg: Call->ReturnRegister);
2575 LLT CondTy = MRI->getType(Reg: Call->Arguments[0]);
2576 if (!ResTy.isVector() && CondTy.isVector())
2577 report_fatal_error(reason: "OpSelect with a scalar result requires a scalar "
2578 "boolean condition");
2579 MIRBuilder.buildSelect(Res: Call->ReturnRegister, Tst: Call->Arguments[0],
2580 Op0: Call->Arguments[1], Op1: Call->Arguments[2]);
2581 return true;
2582}
2583
2584static bool generateConstructInst(const SPIRV::IncomingCall *Call,
2585 MachineIRBuilder &MIRBuilder,
2586 SPIRVGlobalRegistry *GR) {
2587 createContinuedInstructions(MIRBuilder, Opcode: SPIRV::OpCompositeConstruct, MinWC: 3,
2588 ContinuedOpcode: SPIRV::OpCompositeConstructContinuedINTEL,
2589 Args: Call->Arguments, ReturnRegister: Call->ReturnRegister,
2590 TypeID: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
2591 return true;
2592}
2593
2594static bool generateCoopMatrInst(const SPIRV::IncomingCall *Call,
2595 MachineIRBuilder &MIRBuilder,
2596 SPIRVGlobalRegistry *GR) {
2597 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
2598 unsigned Opcode =
2599 SPIRV::lookupNativeBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Opcode;
2600 bool IsSet = Opcode != SPIRV::OpCooperativeMatrixStoreKHR &&
2601 Opcode != SPIRV::OpCooperativeMatrixStoreCheckedINTEL &&
2602 Opcode != SPIRV::OpCooperativeMatrixPrefetchINTEL;
2603 unsigned ArgSz = Call->Arguments.size();
2604 unsigned LiteralIdx = 0;
2605 switch (Opcode) {
2606 // Memory operand is optional and is literal.
2607 case SPIRV::OpCooperativeMatrixLoadKHR:
2608 LiteralIdx = ArgSz > 3 ? 3 : 0;
2609 break;
2610 case SPIRV::OpCooperativeMatrixStoreKHR:
2611 LiteralIdx = ArgSz > 4 ? 4 : 0;
2612 break;
2613 case SPIRV::OpCooperativeMatrixLoadCheckedINTEL:
2614 LiteralIdx = ArgSz > 7 ? 7 : 0;
2615 break;
2616 case SPIRV::OpCooperativeMatrixStoreCheckedINTEL:
2617 LiteralIdx = ArgSz > 8 ? 8 : 0;
2618 break;
2619 // Cooperative Matrix Operands operand is optional and is literal.
2620 case SPIRV::OpCooperativeMatrixMulAddKHR:
2621 LiteralIdx = ArgSz > 3 ? 3 : 0;
2622 break;
2623 };
2624
2625 SmallVector<uint32_t, 1> ImmArgs;
2626 MachineRegisterInfo *MRI = MIRBuilder.getMRI();
2627 if (Opcode == SPIRV::OpCooperativeMatrixPrefetchINTEL) {
2628 const uint32_t CacheLevel = getIConstVal(ConstReg: Call->Arguments[3], MRI);
2629 auto MIB = MIRBuilder.buildInstr(Opcode: SPIRV::OpCooperativeMatrixPrefetchINTEL)
2630 .addUse(RegNo: Call->Arguments[0]) // pointer
2631 .addUse(RegNo: Call->Arguments[1]) // rows
2632 .addUse(RegNo: Call->Arguments[2]) // columns
2633 .addImm(Val: CacheLevel) // cache level
2634 .addUse(RegNo: Call->Arguments[4]); // memory layout
2635 if (ArgSz > 5)
2636 MIB.addUse(RegNo: Call->Arguments[5]); // stride
2637 if (ArgSz > 6) {
2638 const uint32_t MemOp = getIConstVal(ConstReg: Call->Arguments[6], MRI);
2639 MIB.addImm(Val: MemOp); // memory operand
2640 }
2641 return true;
2642 }
2643 if (LiteralIdx > 0)
2644 ImmArgs.push_back(Elt: getIConstVal(ConstReg: Call->Arguments[LiteralIdx], MRI));
2645 Register TypeReg = GR->getSPIRVTypeID(SpirvType: Call->ReturnType);
2646 if (Opcode == SPIRV::OpCooperativeMatrixLengthKHR) {
2647 SPIRVTypeInst CoopMatrType = GR->getSPIRVTypeForVReg(VReg: Call->Arguments[0]);
2648 if (!CoopMatrType)
2649 report_fatal_error(reason: "Can't find a register's type definition");
2650 MIRBuilder.buildInstr(Opcode)
2651 .addDef(RegNo: Call->ReturnRegister)
2652 .addUse(RegNo: TypeReg)
2653 .addUse(RegNo: CoopMatrType->getOperand(i: 0).getReg());
2654 return true;
2655 }
2656 return buildOpFromWrapper(MIRBuilder, Opcode, Call,
2657 TypeReg: IsSet ? TypeReg : Register(0), ImmArgs);
2658}
2659
2660static bool generateSpecConstantInst(const SPIRV::IncomingCall *Call,
2661 MachineIRBuilder &MIRBuilder,
2662 SPIRVGlobalRegistry *GR) {
2663 // Lookup the instruction opcode in the TableGen records.
2664 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
2665 unsigned Opcode =
2666 SPIRV::lookupNativeBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Opcode;
2667 const MachineRegisterInfo *MRI = MIRBuilder.getMRI();
2668
2669 switch (Opcode) {
2670 case SPIRV::OpSpecConstant: {
2671 // Determine the constant MI.
2672 Register ConstRegister = Call->Arguments[1];
2673 const MachineInstr *Const = getDefInstrMaybeConstant(ConstReg&: ConstRegister, MRI);
2674 assert(Const &&
2675 (Const->getOpcode() == TargetOpcode::G_CONSTANT ||
2676 Const->getOpcode() == TargetOpcode::G_FCONSTANT) &&
2677 "Argument should be either an int or floating-point constant");
2678 // Determine the opcode and built the OpSpec MI.
2679 const MachineOperand &ConstOperand = Const->getOperand(i: 1);
2680 if (Call->ReturnType->getOpcode() == SPIRV::OpTypeBool) {
2681 assert(ConstOperand.isCImm() && "Int constant operand is expected");
2682 Opcode = ConstOperand.getCImm()->getValue().getZExtValue()
2683 ? SPIRV::OpSpecConstantTrue
2684 : SPIRV::OpSpecConstantFalse;
2685 }
2686 auto MIB = MIRBuilder.buildInstr(Opcode)
2687 .addDef(RegNo: Call->ReturnRegister)
2688 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
2689
2690 if (Call->ReturnType->getOpcode() != SPIRV::OpTypeBool) {
2691 if (Const->getOpcode() == TargetOpcode::G_CONSTANT)
2692 addNumImm(Imm: ConstOperand.getCImm()->getValue(), MIB);
2693 else
2694 addNumImm(Imm: ConstOperand.getFPImm()->getValueAPF().bitcastToAPInt(), MIB);
2695 }
2696 // Build the SpecID decoration.
2697 unsigned SpecId =
2698 static_cast<unsigned>(getIConstVal(ConstReg: Call->Arguments[0], MRI));
2699 buildOpDecorate(Reg: Call->ReturnRegister, MIRBuilder, Dec: SPIRV::Decoration::SpecId,
2700 DecArgs: {SpecId});
2701 return true;
2702 }
2703 case SPIRV::OpSpecConstantComposite: {
2704 createContinuedInstructions(MIRBuilder, Opcode, MinWC: 3,
2705 ContinuedOpcode: SPIRV::OpSpecConstantCompositeContinuedINTEL,
2706 Args: Call->Arguments, ReturnRegister: Call->ReturnRegister,
2707 TypeID: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
2708 return true;
2709 }
2710 default:
2711 return false;
2712 }
2713}
2714
2715static bool generateExtendedBitOpsInst(const SPIRV::IncomingCall *Call,
2716 MachineIRBuilder &MIRBuilder,
2717 SPIRVGlobalRegistry *GR) {
2718 // Lookup the instruction opcode in the TableGen records.
2719 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
2720 unsigned Opcode =
2721 SPIRV::lookupNativeBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Opcode;
2722
2723 return buildExtendedBitOpsInst(Call, Opcode, MIRBuilder, GR);
2724}
2725
2726static bool generateBindlessImageINTELInst(const SPIRV::IncomingCall *Call,
2727 MachineIRBuilder &MIRBuilder,
2728 SPIRVGlobalRegistry *GR) {
2729 // Lookup the instruction opcode in the TableGen records.
2730 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
2731 unsigned Opcode =
2732 SPIRV::lookupNativeBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Opcode;
2733
2734 return buildBindlessImageINTELInst(Call, Opcode, MIRBuilder, GR);
2735}
2736
2737static bool generateBlockingPipesInst(const SPIRV::IncomingCall *Call,
2738 MachineIRBuilder &MIRBuilder,
2739 SPIRVGlobalRegistry *GR) {
2740 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
2741 unsigned Opcode =
2742 SPIRV::lookupNativeBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Opcode;
2743 return buildOpFromWrapper(MIRBuilder, Opcode, Call, TypeReg: Register(0));
2744}
2745
2746static bool buildAPFixedPointInst(const SPIRV::IncomingCall *Call,
2747 unsigned Opcode, MachineIRBuilder &MIRBuilder,
2748 SPIRVGlobalRegistry *GR, const CallBase &CB) {
2749 MachineRegisterInfo *MRI = MIRBuilder.getMRI();
2750 SmallVector<uint32_t, 1> ImmArgs;
2751 Register InputReg = Call->Arguments[0];
2752 const Type *RetTy = GR->getTypeForSPIRVType(Ty: Call->ReturnType);
2753 bool IsSRet = RetTy->isVoidTy();
2754
2755 if (IsSRet) {
2756 const LLT ValTy = MRI->getType(Reg: InputReg);
2757 Register ActualRetValReg = MRI->createGenericVirtualRegister(Ty: ValTy);
2758 SPIRVTypeInst InstructionType =
2759 deduceSRetPointeeType(SRetReg: InputReg, SRetArg: CB.getArgOperand(i: 0), MIRBuilder, GR);
2760 InputReg = Call->Arguments[1];
2761 auto InputType = GR->getTypeForSPIRVType(Ty: GR->getSPIRVTypeForVReg(VReg: InputReg));
2762 Register PtrInputReg;
2763 if (InputType->getTypeID() == llvm::Type::TypeID::TypedPointerTyID) {
2764 LLT InputLLT = MRI->getType(Reg: InputReg);
2765 PtrInputReg = MRI->createGenericVirtualRegister(Ty: InputLLT);
2766 SPIRVTypeInst PtrType =
2767 GR->getPointeeType(PtrType: GR->getSPIRVTypeForVReg(VReg: InputReg));
2768 MachineMemOperand *MMO1 = MIRBuilder.getMF().getMachineMemOperand(
2769 PtrInfo: MachinePointerInfo(), F: MachineMemOperand::MOLoad,
2770 Size: InputLLT.getSizeInBytes(), BaseAlignment: Align(4));
2771 MIRBuilder.buildLoad(Res: PtrInputReg, Addr: InputReg, MMO&: *MMO1);
2772 MRI->setRegClass(Reg: PtrInputReg, RC: &SPIRV::iIDRegClass);
2773 GR->assignSPIRVTypeToVReg(Type: PtrType, VReg: PtrInputReg, MF: MIRBuilder.getMF());
2774 }
2775
2776 for (unsigned index = 2; index < 7; index++) {
2777 ImmArgs.push_back(Elt: getIConstVal(ConstReg: Call->Arguments[index], MRI));
2778 }
2779
2780 // Emit the instruction
2781 auto MIB = MIRBuilder.buildInstr(Opcode)
2782 .addDef(RegNo: ActualRetValReg)
2783 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: InstructionType));
2784 if (PtrInputReg)
2785 MIB.addUse(RegNo: PtrInputReg);
2786 else
2787 MIB.addUse(RegNo: InputReg);
2788
2789 for (uint32_t Imm : ImmArgs)
2790 MIB.addImm(Val: Imm);
2791 unsigned Size = ValTy.getSizeInBytes();
2792 // Store result to the pointer passed in Arg[0]
2793 MachineMemOperand *MMO = MIRBuilder.getMF().getMachineMemOperand(
2794 PtrInfo: MachinePointerInfo(), F: MachineMemOperand::MOStore, Size, BaseAlignment: Align(4));
2795 MRI->setRegClass(Reg: ActualRetValReg, RC: &SPIRV::pIDRegClass);
2796 MIRBuilder.buildStore(Val: ActualRetValReg, Addr: Call->Arguments[0], MMO&: *MMO);
2797 return true;
2798 } else {
2799 for (unsigned index = 1; index < 6; index++)
2800 ImmArgs.push_back(Elt: getIConstVal(ConstReg: Call->Arguments[index], MRI));
2801
2802 return buildOpFromWrapper(MIRBuilder, Opcode, Call,
2803 TypeReg: GR->getSPIRVTypeID(SpirvType: Call->ReturnType), ImmArgs);
2804 }
2805}
2806
2807static bool generateAPFixedPointInst(const SPIRV::IncomingCall *Call,
2808 MachineIRBuilder &MIRBuilder,
2809 SPIRVGlobalRegistry *GR,
2810 const CallBase &CB) {
2811 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
2812 unsigned Opcode =
2813 SPIRV::lookupNativeBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Opcode;
2814
2815 return buildAPFixedPointInst(Call, Opcode, MIRBuilder, GR, CB);
2816}
2817
2818static bool
2819generateTernaryBitwiseFunctionINTELInst(const SPIRV::IncomingCall *Call,
2820 MachineIRBuilder &MIRBuilder,
2821 SPIRVGlobalRegistry *GR) {
2822 // Lookup the instruction opcode in the TableGen records.
2823 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
2824 unsigned Opcode =
2825 SPIRV::lookupNativeBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Opcode;
2826
2827 return buildTernaryBitwiseFunctionINTELInst(Call, Opcode, MIRBuilder, GR);
2828}
2829
2830static bool generateImageChannelDataTypeInst(const SPIRV::IncomingCall *Call,
2831 MachineIRBuilder &MIRBuilder,
2832 SPIRVGlobalRegistry *GR) {
2833 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
2834 unsigned Opcode =
2835 SPIRV::lookupNativeBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Opcode;
2836
2837 return buildImageChannelDataTypeInst(Call, Opcode, MIRBuilder, GR);
2838}
2839
2840static bool generate2DBlockIOINTELInst(const SPIRV::IncomingCall *Call,
2841 MachineIRBuilder &MIRBuilder,
2842 SPIRVGlobalRegistry *GR) {
2843 // Lookup the instruction opcode in the TableGen records.
2844 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
2845 unsigned Opcode =
2846 SPIRV::lookupNativeBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Opcode;
2847
2848 return build2DBlockIOINTELInst(Call, Opcode, MIRBuilder, GR);
2849}
2850
2851static bool generatePipeInst(const SPIRV::IncomingCall *Call,
2852 MachineIRBuilder &MIRBuilder,
2853 SPIRVGlobalRegistry *GR) {
2854 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
2855 unsigned Opcode =
2856 SPIRV::lookupNativeBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Opcode;
2857
2858 unsigned Scope = SPIRV::Scope::Workgroup;
2859 if (Builtin->name().contains(Other: "sub_group"))
2860 Scope = SPIRV::Scope::Subgroup;
2861
2862 return buildPipeInst(Call, Opcode, Scope, MIRBuilder, GR);
2863}
2864
2865static bool generatePredicatedLoadStoreInst(const SPIRV::IncomingCall *Call,
2866 MachineIRBuilder &MIRBuilder,
2867 SPIRVGlobalRegistry *GR) {
2868 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
2869 unsigned Opcode =
2870 SPIRV::lookupNativeBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Opcode;
2871
2872 bool IsSet = Opcode != SPIRV::OpPredicatedStoreINTEL;
2873 unsigned ArgSz = Call->Arguments.size();
2874 SmallVector<uint32_t, 1> ImmArgs;
2875 MachineRegisterInfo *MRI = MIRBuilder.getMRI();
2876 // Memory operand is optional and is literal.
2877 if (ArgSz > 3)
2878 ImmArgs.push_back(Elt: getIConstVal(ConstReg: Call->Arguments[/*Literal index*/ 3], MRI));
2879
2880 Register TypeReg = GR->getSPIRVTypeID(SpirvType: Call->ReturnType);
2881 return buildOpFromWrapper(MIRBuilder, Opcode, Call,
2882 TypeReg: IsSet ? TypeReg : Register(0), ImmArgs);
2883}
2884
2885static bool buildNDRange(const SPIRV::IncomingCall *Call,
2886 MachineIRBuilder &MIRBuilder, SPIRVGlobalRegistry *GR,
2887 const CallBase &CB) {
2888 // The OpenCL ndrange_*D functions are overloaded and support 1D, 2D, and 3D
2889 // variants, accepting 1 to 3 arguments:
2890 // (global_work_size)
2891 // (global_work_size, local_work_size)
2892 // (global_work_offset, global_work_size, local_work_size)
2893 // Note: When all three arguments are provided, they are reordered compared
2894 // to the one- or two-argument form.
2895 //
2896 // The function may return data through an sret argument at position 0 (with
2897 // a void function return type). When present, all other argument indices are
2898 // adjusted accordingly.
2899 //
2900 // SPIR-V's OpBuildNDRange requires all three arguments (GlobalWorkSize,
2901 // LocalWorkSize, GlobalWorkOffset). For 1D kernels, the values are scalars;
2902 // for 2D/3D kernels, they are arrays of 2 or 3 elements. Missing arguments
2903 // default to zero.
2904 //
2905 // Calculate argument indices based on the number of arguments and presence
2906 // of sret:
2907 const unsigned NumCallArgs = Call->Arguments.size();
2908 const unsigned MaxCallArgs = Call->Builtin->MaxNumArgs;
2909 const unsigned IncorrectArgIdx = MaxCallArgs + 1;
2910
2911 const Type *RetTy = GR->getTypeForSPIRVType(Ty: Call->ReturnType);
2912 bool HasSRetArg = RetTy->isVoidTy();
2913
2914 const unsigned SRetArgIdx = HasSRetArg ? 0 : IncorrectArgIdx;
2915 const unsigned ArgBase = HasSRetArg ? 1 : 0;
2916 const unsigned MaxNDRangeArgs = 3;
2917 const unsigned NumNDRangeArgs = NumCallArgs - ArgBase;
2918
2919 const unsigned GlobalWorkSizeArgIdx =
2920 NumNDRangeArgs < MaxNDRangeArgs ? ArgBase : ArgBase + 1;
2921 const unsigned LocalWorkSizeArgIdx =
2922 (NumNDRangeArgs == 1)
2923 ? IncorrectArgIdx
2924 : (NumNDRangeArgs == MaxNDRangeArgs ? ArgBase + 2 : ArgBase + 1);
2925 const unsigned GlobalWorkOffsetArgIdx =
2926 NumNDRangeArgs == MaxNDRangeArgs ? ArgBase : IncorrectArgIdx;
2927
2928 // Each nd_range field is an array of <Dimension> integers matching the
2929 // address model width (32 or 64 bits).
2930 const unsigned AddressModelBits = GR->getPointerSize();
2931 assert(AddressModelBits == 64 || AddressModelBits == 32);
2932
2933 // The dimension is encoded in the function name as "ndrange_XD" where X is
2934 // 1, 2, or 3.
2935 unsigned Dimension = 0;
2936 Call->Builtin->name().substr(Start: 8, N: 1).getAsInteger(Radix: 10, Result&: Dimension);
2937 assert(Dimension <= 3 && Dimension >= 1);
2938
2939 // Determine the work size type based on the dimension. For missing arguments,
2940 // create a zero constant of the appropriate type.
2941 MachineFunction &MF = MIRBuilder.getMF();
2942 SPIRVTypeInst SpvFieldTy;
2943 Register ConstZero;
2944 if (Dimension == 1) {
2945 SpvFieldTy = GR->getSPIRVTypeForVReg(VReg: Call->Arguments[GlobalWorkSizeArgIdx]);
2946 assert(SpvFieldTy && SpvFieldTy->getOpcode() == SPIRV::OpTypeInt &&
2947 "Expected scalar integer type");
2948
2949 if (NumNDRangeArgs < MaxNDRangeArgs)
2950 ConstZero = GR->buildConstantInt(Val: 0, MIRBuilder, SpvType: SpvFieldTy, EmitIR: true);
2951 } else {
2952 Type *BaseTy =
2953 IntegerType::get(C&: MF.getFunction().getContext(), NumBits: AddressModelBits);
2954 Type *FieldTy = ArrayType::get(ElementType: BaseTy, NumElements: Dimension);
2955 SpvFieldTy = GR->getOrCreateSPIRVType(
2956 Type: FieldTy, MIRBuilder, AQ: SPIRV::AccessQualifier::ReadOnly, EmitIR: true);
2957
2958 if (NumNDRangeArgs < MaxNDRangeArgs) {
2959 auto InsertIt = MIRBuilder.getInsertPt();
2960 MachineBasicBlock &MBB = MIRBuilder.getMBB();
2961 MachineInstr &InsertMI = (InsertIt != MBB.end()) ? *InsertIt : MBB.back();
2962 const SPIRVSubtarget &ST = cast<SPIRVSubtarget>(Val: MF.getSubtarget());
2963 ConstZero = GR->getOrCreateConstIntArray(Val: 0, Num: Dimension, I&: InsertMI,
2964 SpvType: SpvFieldTy, TII: *ST.getInstrInfo());
2965 }
2966 }
2967
2968 MachineRegisterInfo *MRI = MIRBuilder.getMRI();
2969
2970 auto CreateDataRegister = [&](unsigned Idx) -> Register {
2971 Register Reg = (Idx == IncorrectArgIdx) ? ConstZero : Call->Arguments[Idx];
2972
2973 if (GR->getSPIRVTypeForVReg(VReg: Reg) == SpvFieldTy) {
2974 // Already has the correct type.
2975 return Reg;
2976 }
2977
2978 assert(GR->getSPIRVTypeForVReg(Reg).isPointer() &&
2979 "Only pointer types are supported for loading values");
2980
2981 Register Ptr = Reg;
2982
2983 Reg = MRI->createVirtualRegister(RegClass: &SPIRV::iIDRegClass);
2984 GR->assignSPIRVTypeToVReg(Type: SpvFieldTy, VReg: Reg, MF);
2985
2986 MIRBuilder.buildInstr(Opcode: SPIRV::OpLoad)
2987 .addDef(RegNo: Reg)
2988 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: SpvFieldTy))
2989 .addUse(RegNo: Ptr);
2990 return Reg;
2991 };
2992
2993 Register GlobalWorkSize = CreateDataRegister(GlobalWorkSizeArgIdx);
2994 Register LocalWorkSize = CreateDataRegister(LocalWorkSizeArgIdx);
2995 Register GlobalWorkOffset = CreateDataRegister(GlobalWorkOffsetArgIdx);
2996
2997 if (!HasSRetArg) {
2998 return MIRBuilder.buildInstr(Opcode: SPIRV::OpBuildNDRange)
2999 .addDef(RegNo: Call->ReturnRegister)
3000 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType))
3001 .addUse(RegNo: GlobalWorkSize)
3002 .addUse(RegNo: LocalWorkSize)
3003 .addUse(RegNo: GlobalWorkOffset);
3004 }
3005
3006 // When sret is used, store nd_range struct through the pointer in the first
3007 // argument.
3008 Register SRetReg = Call->Arguments[SRetArgIdx];
3009 SPIRVTypeInst SRetType = deduceSRetPointeeType(
3010 SRetReg, SRetArg: CB.getArgOperand(i: SRetArgIdx), MIRBuilder, GR);
3011
3012 Register TmpReg = MRI->createVirtualRegister(RegClass: &SPIRV::iIDRegClass);
3013 GR->assignSPIRVTypeToVReg(Type: SRetType, VReg: TmpReg, MF);
3014
3015 MIRBuilder.buildInstr(Opcode: SPIRV::OpBuildNDRange)
3016 .addDef(RegNo: TmpReg)
3017 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: SRetType))
3018 .addUse(RegNo: GlobalWorkSize)
3019 .addUse(RegNo: LocalWorkSize)
3020 .addUse(RegNo: GlobalWorkOffset);
3021 return MIRBuilder.buildInstr(Opcode: SPIRV::OpStore)
3022 .addUse(RegNo: Call->Arguments[SRetArgIdx])
3023 .addUse(RegNo: TmpReg);
3024}
3025
3026static void buildKernelInvokeOperands(
3027 const SPIRV::IncomingCall *Call, unsigned InvokeIdx, unsigned ParamIdx,
3028 MachineIRBuilder &MIRBuilder, SPIRVGlobalRegistry *GR, Register &InvokeReg,
3029 Register &ParamReg, Register &ParamSizeReg, Register &ParamAlignReg) {
3030 MachineRegisterInfo *MRI = MIRBuilder.getMRI();
3031 const DataLayout &DL = MIRBuilder.getDataLayout();
3032
3033 // Bypass the addrspacecast so Invoke references the function's <id>.
3034 MachineInstr *InvokeGlobalMI =
3035 getBlockStructInstr(ParamReg: Call->Arguments[InvokeIdx], MRI);
3036 assert(InvokeGlobalMI->getOpcode() == TargetOpcode::G_GLOBAL_VALUE);
3037 InvokeReg = InvokeGlobalMI->getOperand(i: 0).getReg();
3038 MRI->setRegClass(Reg: InvokeReg, RC: &SPIRV::pIDRegClass);
3039
3040 Register BlockLiteralReg = Call->Arguments[ParamIdx];
3041 const SPIRVTypeInst Int8Ty = GR->getOrCreateSPIRVIntegerType(BitWidth: 8, MIRBuilder);
3042 const SPIRVTypeInst Int8PtrGen = GR->getOrCreateSPIRVPointerType(
3043 BaseType: Int8Ty, MIRBuilder, SC: SPIRV::StorageClass::Generic);
3044 Type *PType = const_cast<Type *>(getBlockStructType(ParamReg: BlockLiteralReg, MRI));
3045
3046 ParamReg = createVirtualRegister(SpvType: Int8PtrGen, GR, MIRBuilder);
3047 MIRBuilder.buildInstr(Opcode: SPIRV::OpBitcast)
3048 .addDef(RegNo: ParamReg)
3049 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Int8PtrGen))
3050 .addUse(RegNo: BlockLiteralReg);
3051 // TODO: these numbers should be obtained from block literal structure.
3052 ParamSizeReg =
3053 buildConstantIntReg32(Val: DL.getTypeStoreSize(Ty: PType), MIRBuilder, GR);
3054 ParamAlignReg =
3055 buildConstantIntReg32(Val: DL.getPrefTypeAlign(Ty: PType).value(), MIRBuilder, GR);
3056}
3057
3058static bool buildKernelQuery(const SPIRV::IncomingCall *Call, unsigned Opcode,
3059 MachineIRBuilder &MIRBuilder,
3060 SPIRVGlobalRegistry *GR) {
3061 bool HasNDRange = Call->Builtin->name().contains(Other: "_ndrange_impl");
3062 unsigned InvokeIdx = HasNDRange ? 1 : 0;
3063 Register InvokeReg, ParamReg, ParamSizeReg, ParamAlignReg;
3064 buildKernelInvokeOperands(Call, InvokeIdx, ParamIdx: InvokeIdx + 1, MIRBuilder, GR,
3065 InvokeReg, ParamReg, ParamSizeReg, ParamAlignReg);
3066 auto MIB = MIRBuilder.buildInstr(Opcode)
3067 .addDef(RegNo: Call->ReturnRegister)
3068 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
3069 if (HasNDRange)
3070 MIB.addUse(RegNo: Call->Arguments[0]);
3071 MIB.addUse(RegNo: InvokeReg)
3072 .addUse(RegNo: ParamReg)
3073 .addUse(RegNo: ParamSizeReg)
3074 .addUse(RegNo: ParamAlignReg);
3075 return true;
3076}
3077
3078static bool buildEnqueueKernel(const SPIRV::IncomingCall *Call,
3079 MachineIRBuilder &MIRBuilder,
3080 SPIRVGlobalRegistry *GR) {
3081 // In this function there are three stages:
3082 // 1. prepare call indexes in order we expect them.
3083 // 2. process all arguments which requered preparation.
3084 // 3. create a SPIRV operator with arguments.
3085
3086 MachineRegisterInfo *MRI = MIRBuilder.getMRI();
3087 const SPIRVTypeInst Int32Ty = GR->getOrCreateSPIRVIntegerType(BitWidth: 32, MIRBuilder);
3088
3089 // 1. prepare call indexes in order we expect them.
3090 // Based on clang sources, clang/lib/CodeGen/CGBuiltin.cpp, BIenqueue_kernel,
3091 // We expect 4 different layouts of call arguments:
3092 // 1) No events, no vargs: {Queue, Flags, Range, Kernel, Block};
3093 // 2) No events, varargs: {Queue, Flags, Range, Kernel, Block, NumElem,
3094 // ElemPtr};
3095 // 3) events, no varargs: {Queue, Flags, Range, NumEvents,
3096 // EventWaitList, EventRet, Kernel, Block};
3097 // 4) events, varargs: {Queue,
3098 // Flags, Range, NumEvents, EventWaitList, EventRet, Kernel, Block,
3099 // NumElem, ElemPtr};
3100 //
3101 // We also may expect __spirv_EnqueueKernel
3102
3103 bool IsSpirvOp = Call->isSpirvOp();
3104 bool HasEvents = Call->Builtin->name().contains(Other: "_events") || IsSpirvOp;
3105 bool HasVarArgs = Call->Builtin->name().contains(Other: "_varargs") || IsSpirvOp;
3106
3107 const unsigned NumArgs = Call->Arguments.size();
3108 const unsigned BaseArgIdx = 0;
3109 const unsigned IncorrectIdx = NumArgs + 1;
3110
3111 const unsigned QueueIdx = BaseArgIdx;
3112 const unsigned FlagsIdx = BaseArgIdx + 1;
3113 const unsigned NDRangeIdx = BaseArgIdx + 2;
3114 const unsigned NumEventsIdx = HasEvents ? BaseArgIdx + 3 : IncorrectIdx;
3115 const unsigned WaitEventsIdx = HasEvents ? BaseArgIdx + 4 : IncorrectIdx;
3116 const unsigned RetEventIdx = HasEvents ? BaseArgIdx + 5 : IncorrectIdx;
3117 const unsigned InvokeIdx = BaseArgIdx + 3 + (HasEvents ? 3 : 0);
3118 const unsigned ParamIdx = BaseArgIdx + 4 + (HasEvents ? 3 : 0);
3119 const unsigned LocalSizeNumElemIdx =
3120 HasVarArgs ? (BaseArgIdx + 5 + (HasEvents ? 3 : 0)) : IncorrectIdx;
3121 const unsigned LocalSizeElemPtrIdx =
3122 HasVarArgs ? (BaseArgIdx + 6 + (HasEvents ? 3 : 0)) : IncorrectIdx;
3123
3124 [[maybe_unused]] const unsigned LastArgIdx =
3125 (BaseArgIdx + 4 + (HasEvents ? 3 : 0) + (HasVarArgs ? 2 : 0));
3126 assert(LastArgIdx < NumArgs && "Incorrect number arguments");
3127
3128 // 2. Process all arguments which requered preparation.
3129 // 2.1 Events - use Call arguments, or use dummy nulls in case of absence of
3130 // events
3131
3132 auto BuildDeviceEventNullPtr = [&]() {
3133 LLVMContext &Ctx = MIRBuilder.getMF().getFunction().getContext();
3134 Type *DeviceEventTy = TargetExtType::get(Context&: Ctx, Name: "spirv.DeviceEvent");
3135 SPIRVTypeInst DeviceEventPtrTy = GR->getOrCreateSPIRVPointerType(
3136 BaseType: DeviceEventTy, MIRBuilder, SC: SPIRV::StorageClass::Generic);
3137 return GR->getOrCreateConstNullPtr(MIRBuilder, SpvType: DeviceEventPtrTy);
3138 };
3139
3140 Register NumEventsReg;
3141 Register WaitEventsReg;
3142 Register RetEventReg;
3143 if (HasEvents) {
3144 auto IsNullEvent = [&](Register R) {
3145 MachineInstr *Def = getDefInstrMaybeConstant(ConstReg&: R, MRI);
3146 return Def->getOpcode() == TargetOpcode::G_CONSTANT &&
3147 Def->getOperand(i: 1).getCImm()->isZero();
3148 };
3149
3150 NumEventsReg = Call->Arguments[NumEventsIdx];
3151 WaitEventsReg = Call->Arguments[WaitEventsIdx];
3152 RetEventReg = Call->Arguments[RetEventIdx];
3153 if (IsNullEvent(WaitEventsReg))
3154 WaitEventsReg = BuildDeviceEventNullPtr();
3155 if (IsNullEvent(RetEventReg))
3156 RetEventReg = BuildDeviceEventNullPtr();
3157 } else {
3158 NumEventsReg = buildConstantIntReg32(Val: 0, MIRBuilder, GR);
3159 Register NullPtr = BuildDeviceEventNullPtr();
3160 WaitEventsReg = NullPtr;
3161 RetEventReg = NullPtr;
3162 }
3163
3164 Register InvokeReg, ParamReg, ParamSizeReg, ParamAlignReg;
3165 buildKernelInvokeOperands(Call, InvokeIdx, ParamIdx, MIRBuilder, GR,
3166 InvokeReg, ParamReg, ParamSizeReg, ParamAlignReg);
3167
3168 // 2.3 Local Size Array
3169 SmallVector<Register, 16> LocalSizes;
3170 if (HasVarArgs) {
3171 Register LocalSizeNumElem = Call->Arguments[LocalSizeNumElemIdx];
3172 MachineInstr *LocalSizeNumElemMI = MRI->getUniqueVRegDef(Reg: LocalSizeNumElem);
3173 const MachineOperand &ConstOp = LocalSizeNumElemMI->getOperand(i: 1);
3174 assert(LocalSizeNumElemMI->getOpcode() == TargetOpcode::G_CONSTANT &&
3175 ConstOp.isCImm() && "Expected constant immediate");
3176 uint64_t NumElem = ConstOp.getCImm()->getValue().getZExtValue();
3177
3178 Register LocalSizeArrayReg = Call->Arguments[LocalSizeElemPtrIdx];
3179
3180 for (unsigned i = 0; i < NumElem; ++i) {
3181 Register Reg = MRI->createVirtualRegister(RegClass: &SPIRV::pIDRegClass);
3182 auto GEPInst = MIRBuilder.buildIntrinsic(
3183 ID: Intrinsic::spv_gep, Res: ArrayRef<Register>{Reg}, HasSideEffects: true, isConvergent: false);
3184 GEPInst
3185 .addImm(Val: 0) // In bound.
3186 .addUse(RegNo: LocalSizeArrayReg) // Base pointer.
3187 .addUse(RegNo: buildConstantIntReg32(Val: 0, MIRBuilder, GR)) // Indices.
3188 .addUse(RegNo: buildConstantIntReg32(Val: i, MIRBuilder, GR));
3189 LocalSizes.push_back(Elt: Reg);
3190 }
3191 }
3192
3193 // 3. create a SPIRV operator with arguments.
3194 auto MIB = MIRBuilder.buildInstr(Opcode: SPIRV::OpEnqueueKernel)
3195 .addDef(RegNo: Call->ReturnRegister)
3196 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Int32Ty))
3197 .addUse(RegNo: Call->Arguments[QueueIdx])
3198 .addUse(RegNo: Call->Arguments[FlagsIdx])
3199 .addUse(RegNo: Call->Arguments[NDRangeIdx])
3200 .addUse(RegNo: NumEventsReg)
3201 .addUse(RegNo: WaitEventsReg)
3202 .addUse(RegNo: RetEventReg)
3203 .addUse(RegNo: InvokeReg)
3204 .addUse(RegNo: ParamReg)
3205 .addUse(RegNo: ParamSizeReg)
3206 .addUse(RegNo: ParamAlignReg);
3207 for (auto &LocalSize : LocalSizes)
3208 MIB.addUse(RegNo: LocalSize);
3209
3210 return true;
3211}
3212
3213static bool generateEnqueueInst(const SPIRV::IncomingCall *Call,
3214 MachineIRBuilder &MIRBuilder,
3215 SPIRVGlobalRegistry *GR, const CallBase &CB) {
3216 // Lookup the instruction opcode in the TableGen records.
3217 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
3218 unsigned Opcode =
3219 SPIRV::lookupNativeBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Opcode;
3220
3221 switch (Opcode) {
3222 case SPIRV::OpRetainEvent:
3223 case SPIRV::OpReleaseEvent:
3224 return MIRBuilder.buildInstr(Opcode).addUse(RegNo: Call->Arguments[0]);
3225 case SPIRV::OpCreateUserEvent:
3226 case SPIRV::OpGetDefaultQueue:
3227 return MIRBuilder.buildInstr(Opcode)
3228 .addDef(RegNo: Call->ReturnRegister)
3229 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
3230 case SPIRV::OpIsValidEvent:
3231 return MIRBuilder.buildInstr(Opcode)
3232 .addDef(RegNo: Call->ReturnRegister)
3233 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType))
3234 .addUse(RegNo: Call->Arguments[0]);
3235 case SPIRV::OpSetUserEventStatus:
3236 return MIRBuilder.buildInstr(Opcode)
3237 .addUse(RegNo: Call->Arguments[0])
3238 .addUse(RegNo: Call->Arguments[1]);
3239 case SPIRV::OpCaptureEventProfilingInfo:
3240 return MIRBuilder.buildInstr(Opcode)
3241 .addUse(RegNo: Call->Arguments[0])
3242 .addUse(RegNo: Call->Arguments[1])
3243 .addUse(RegNo: Call->Arguments[2]);
3244 case SPIRV::OpBuildNDRange:
3245 return buildNDRange(Call, MIRBuilder, GR, CB);
3246 case SPIRV::OpEnqueueKernel:
3247 return buildEnqueueKernel(Call, MIRBuilder, GR);
3248 case SPIRV::OpGetKernelNDrangeSubGroupCount:
3249 case SPIRV::OpGetKernelNDrangeMaxSubGroupSize:
3250 case SPIRV::OpGetKernelWorkGroupSize:
3251 case SPIRV::OpGetKernelPreferredWorkGroupSizeMultiple:
3252 return buildKernelQuery(Call, Opcode, MIRBuilder, GR);
3253 default:
3254 return false;
3255 }
3256}
3257
3258static bool generateAsyncCopy(const SPIRV::IncomingCall *Call,
3259 MachineIRBuilder &MIRBuilder,
3260 SPIRVGlobalRegistry *GR, const CallBase &CB) {
3261 // Lookup the instruction opcode in the TableGen records.
3262 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
3263 unsigned Opcode =
3264 SPIRV::lookupNativeBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Opcode;
3265
3266 bool IsSet = Opcode == SPIRV::OpGroupAsyncCopy;
3267 Register TypeReg = GR->getSPIRVTypeID(SpirvType: Call->ReturnType);
3268 if (Call->isSpirvOp())
3269 return buildOpFromWrapper(MIRBuilder, Opcode, Call,
3270 TypeReg: IsSet ? TypeReg : Register(0));
3271
3272 auto Scope = buildConstantIntReg32(Val: SPIRV::Scope::Workgroup, MIRBuilder, GR);
3273
3274 switch (Opcode) {
3275 case SPIRV::OpGroupAsyncCopy: {
3276 SPIRVTypeInst NewType =
3277 Call->ReturnType->getOpcode() == SPIRV::OpTypeEvent
3278 ? nullptr
3279 : GR->getOrCreateSPIRVTypeByName(TypeStr: "spirv.Event", MIRBuilder, EmitIR: true);
3280 Register TypeReg = GR->getSPIRVTypeID(SpirvType: NewType ? NewType : Call->ReturnType);
3281 unsigned NumArgs = Call->Arguments.size();
3282 Register EventReg = Call->Arguments[NumArgs - 1];
3283 SPIRVTypeInst EventType = GR->getSPIRVTypeForVReg(VReg: EventReg);
3284 if (!EventType || EventType->getOpcode() != SPIRV::OpTypeEvent) {
3285 MachineRegisterInfo *MRI = MIRBuilder.getMRI();
3286 Register ConstReg = EventReg;
3287 MachineInstr *Def = getDefInstrMaybeConstant(ConstReg, MRI);
3288 SPIRVTypeInst EventPointeeType =
3289 EventType && EventType->getOpcode() == SPIRV::OpTypePointer
3290 ? GR->getPointeeType(PtrType: EventType)
3291 : nullptr;
3292 if (Def->getOpcode() == TargetOpcode::G_CONSTANT &&
3293 Def->getOperand(i: 1).getCImm()->isZero()) {
3294 // Only substitute a null Event for the "ptr null" idiom, not for a
3295 // real event value that just is not typed as OpTypeEvent yet.
3296 SPIRVTypeInst EventTy = NewType ? NewType
3297 : GR->getOrCreateSPIRVTypeByName(
3298 TypeStr: "spirv.Event", MIRBuilder, EmitIR: true);
3299 Register EventTyReg = GR->getSPIRVTypeID(SpirvType: EventTy);
3300 Register NullEventReg = createVirtualRegister(SpvType: EventTy, GR, MIRBuilder);
3301 MIRBuilder.buildInstr(Opcode: SPIRV::OpConstantNull)
3302 .addDef(RegNo: NullEventReg)
3303 .addUse(RegNo: EventTyReg);
3304 EventReg = NullEventReg;
3305 } else if (EventPointeeType &&
3306 EventPointeeType->getOpcode() == SPIRV::OpTypeEvent) {
3307 // Dereference: a real event can end up typed as pointer-to-Event
3308 // after round-tripping through a stack slot under the legacy
3309 // opaque-ptr ocl_event ABI.
3310 Register EventTyReg = GR->getSPIRVTypeID(SpirvType: EventPointeeType);
3311 Register LoadedReg =
3312 createVirtualRegister(SpvType: EventPointeeType, GR, MIRBuilder);
3313 MIRBuilder.buildInstr(Opcode: SPIRV::OpLoad)
3314 .addDef(RegNo: LoadedReg)
3315 .addUse(RegNo: EventTyReg)
3316 .addUse(RegNo: EventReg);
3317 EventReg = LoadedReg;
3318 }
3319 }
3320 Register NumElemReg = Call->Arguments[2];
3321
3322 // Untyped pointers use OpUntypedGroupAsyncCopyKHR, which adds an explicit
3323 // Element Num Bytes operand.
3324 SPIRVTypeInst DestPtrTy = GR->getSPIRVTypeForVReg(VReg: Call->Arguments[0]);
3325 bool IsUntyped =
3326 DestPtrTy && DestPtrTy->getOpcode() == SPIRV::OpTypeUntypedPointerKHR;
3327 SPIRVTypeInst SizeTy = GR->getSPIRVTypeForVReg(VReg: NumElemReg);
3328 Register StrideReg =
3329 Call->Arguments.size() > 4
3330 ? Call->Arguments[3]
3331 : (IsUntyped ? GR->buildConstantInt(Val: 1, MIRBuilder, SpvType: SizeTy,
3332 /*EmitIR=*/true)
3333 : buildConstantIntReg32(Val: 1, MIRBuilder, GR));
3334
3335 auto MIB = MIRBuilder
3336 .buildInstr(Opcode: IsUntyped ? SPIRV::OpUntypedGroupAsyncCopyKHR
3337 : SPIRV::OpGroupAsyncCopy)
3338 .addDef(RegNo: Call->ReturnRegister)
3339 .addUse(RegNo: TypeReg)
3340 .addUse(RegNo: Scope)
3341 .addUse(RegNo: Call->Arguments[0])
3342 .addUse(RegNo: Call->Arguments[1]);
3343 if (IsUntyped) {
3344 // Element Num Bytes from the deduced element type of dest (or source).
3345 unsigned ElemBytes = GR->getDeducedPointeeByteSize(PtrVal: CB.getArgOperand(i: 0));
3346 if (!ElemBytes)
3347 ElemBytes = GR->getDeducedPointeeByteSize(PtrVal: CB.getArgOperand(i: 1));
3348 if (!ElemBytes)
3349 report_fatal_error(reason: "Could not deduce the element type of an untyped "
3350 "async copy pointer argument");
3351 MIB.addUse(RegNo: GR->buildConstantInt(Val: ElemBytes, MIRBuilder, SpvType: SizeTy,
3352 /*EmitIR=*/true));
3353 }
3354 MIB.addUse(RegNo: NumElemReg);
3355 MIB.addUse(RegNo: StrideReg);
3356 MIB.addUse(RegNo: EventReg);
3357 if (NewType)
3358 updateRegType(Reg: Call->ReturnRegister, /*Ty=*/nullptr, SpirvTy: NewType, GR,
3359 MIB&: MIRBuilder, MRI&: MIRBuilder.getMF().getRegInfo());
3360 return true;
3361 }
3362 case SPIRV::OpGroupWaitEvents:
3363 return MIRBuilder.buildInstr(Opcode)
3364 .addUse(RegNo: Scope)
3365 .addUse(RegNo: Call->Arguments[0])
3366 .addUse(RegNo: Call->Arguments[1]);
3367 default:
3368 return false;
3369 }
3370}
3371
3372// Same type S/U/FConvert are invalid. OpSatConvert* are valid and must stay.
3373static bool foldNoOpConvert(unsigned Opcode, const SPIRV::IncomingCall *Call,
3374 MachineIRBuilder &MIRBuilder,
3375 SPIRVGlobalRegistry *GR) {
3376 if (Opcode != SPIRV::OpSConvert && Opcode != SPIRV::OpUConvert &&
3377 Opcode != SPIRV::OpFConvert)
3378 return false;
3379 if (Call->Arguments.size() != 1 ||
3380 GR->getSPIRVTypeForVReg(VReg: Call->Arguments[0]) != Call->ReturnType)
3381 return false;
3382 MIRBuilder.buildCopy(Res: Call->ReturnRegister, Op: Call->Arguments[0]);
3383 return true;
3384}
3385
3386static bool generateConvertInst(StringRef DemangledCall,
3387 const SPIRV::IncomingCall *Call,
3388 MachineIRBuilder &MIRBuilder,
3389 SPIRVGlobalRegistry *GR) {
3390 // Lookup the conversion builtin in the TableGen records.
3391 const SPIRV::ConvertBuiltin *Builtin =
3392 SPIRV::lookupConvertBuiltin(Name: Call->Builtin->name(), Set: Call->Builtin->Set);
3393
3394 if (!Builtin && Call->isSpirvOp()) {
3395 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
3396 unsigned Opcode =
3397 SPIRV::lookupNativeBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Opcode;
3398 if (foldNoOpConvert(Opcode, Call, MIRBuilder, GR))
3399 return true;
3400 return buildOpFromWrapper(MIRBuilder, Opcode, Call,
3401 TypeReg: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
3402 }
3403
3404 assert(Builtin && "Conversion builtin not found.");
3405
3406 std::string NeedExtMsg; // no errors if empty
3407 bool IsRightComponentsNumber = true; // check if input/output accepts vectors
3408 unsigned Opcode = SPIRV::OpNop;
3409 if (GR->isScalarOrVectorOfType(VReg: Call->Arguments[0], TypeOpcode: SPIRV::OpTypeInt)) {
3410 // Int -> ...
3411 bool IsSourceSigned =
3412 DemangledCall[DemangledCall.find_first_of(C: '(') + 1] != 'u';
3413 if (GR->isScalarOrVectorOfType(VReg: Call->ReturnRegister, TypeOpcode: SPIRV::OpTypeInt)) {
3414 // Int -> Int
3415 if (Builtin->IsSaturated)
3416 Opcode = Builtin->IsDestinationSigned ? SPIRV::OpSatConvertUToS
3417 : SPIRV::OpSatConvertSToU;
3418 else
3419 Opcode = IsSourceSigned ? SPIRV::OpSConvert : SPIRV::OpUConvert;
3420 } else if (GR->isScalarOrVectorOfType(VReg: Call->ReturnRegister,
3421 TypeOpcode: SPIRV::OpTypeFloat)) {
3422 // Int -> Float
3423 if (Builtin->IsBfloat16) {
3424 const auto *ST = static_cast<const SPIRVSubtarget *>(
3425 &MIRBuilder.getMF().getSubtarget());
3426 if (!ST->canUseExtension(
3427 E: SPIRV::Extension::SPV_INTEL_bfloat16_conversion))
3428 NeedExtMsg = "SPV_INTEL_bfloat16_conversion";
3429 IsRightComponentsNumber =
3430 GR->getScalarOrVectorComponentCount(VReg: Call->Arguments[0]) ==
3431 GR->getScalarOrVectorComponentCount(VReg: Call->ReturnRegister);
3432 Opcode = SPIRV::OpConvertBF16ToFINTEL;
3433 } else {
3434 Opcode = IsSourceSigned ? SPIRV::OpConvertSToF : SPIRV::OpConvertUToF;
3435 }
3436 }
3437 } else if (GR->isScalarOrVectorOfType(VReg: Call->Arguments[0],
3438 TypeOpcode: SPIRV::OpTypeFloat)) {
3439 // Float -> ...
3440 if (GR->isScalarOrVectorOfType(VReg: Call->ReturnRegister, TypeOpcode: SPIRV::OpTypeInt)) {
3441 // Float -> Int
3442 if (Builtin->IsBfloat16) {
3443 const auto *ST = static_cast<const SPIRVSubtarget *>(
3444 &MIRBuilder.getMF().getSubtarget());
3445 if (!ST->canUseExtension(
3446 E: SPIRV::Extension::SPV_INTEL_bfloat16_conversion))
3447 NeedExtMsg = "SPV_INTEL_bfloat16_conversion";
3448 IsRightComponentsNumber =
3449 GR->getScalarOrVectorComponentCount(VReg: Call->Arguments[0]) ==
3450 GR->getScalarOrVectorComponentCount(VReg: Call->ReturnRegister);
3451 Opcode = SPIRV::OpConvertFToBF16INTEL;
3452 } else {
3453 Opcode = Builtin->IsDestinationSigned ? SPIRV::OpConvertFToS
3454 : SPIRV::OpConvertFToU;
3455 }
3456 } else if (GR->isScalarOrVectorOfType(VReg: Call->ReturnRegister,
3457 TypeOpcode: SPIRV::OpTypeFloat)) {
3458 if (Builtin->IsTF32) {
3459 const auto *ST = static_cast<const SPIRVSubtarget *>(
3460 &MIRBuilder.getMF().getSubtarget());
3461 if (!ST->canUseExtension(
3462 E: SPIRV::Extension::SPV_INTEL_tensor_float32_conversion))
3463 NeedExtMsg = "SPV_INTEL_tensor_float32_conversion";
3464 IsRightComponentsNumber =
3465 GR->getScalarOrVectorComponentCount(VReg: Call->Arguments[0]) ==
3466 GR->getScalarOrVectorComponentCount(VReg: Call->ReturnRegister);
3467 Opcode = SPIRV::OpRoundFToTF32INTEL;
3468 } else {
3469 // Float -> Float
3470 Opcode = SPIRV::OpFConvert;
3471 }
3472 }
3473 }
3474
3475 StringRef BuiltinName = SPIRV::getConvertBuiltinStr(Offset: Builtin->Name);
3476 if (!NeedExtMsg.empty()) {
3477 std::string DiagMsg = std::string(BuiltinName) +
3478 ": the builtin requires the following SPIR-V "
3479 "extension: " +
3480 NeedExtMsg;
3481 report_fatal_error(reason: DiagMsg.c_str(), gen_crash_diag: false);
3482 }
3483 if (!IsRightComponentsNumber) {
3484 std::string DiagMsg =
3485 std::string(BuiltinName) +
3486 ": result and argument must have the same number of components";
3487 report_fatal_error(reason: DiagMsg.c_str(), gen_crash_diag: false);
3488 }
3489 assert(Opcode != SPIRV::OpNop &&
3490 "Conversion between the types not implemented!");
3491
3492 // Must run before the decorations below: a folded conversion has none.
3493 if (foldNoOpConvert(Opcode, Call, MIRBuilder, GR))
3494 return true;
3495
3496 if (Builtin->IsSaturated)
3497 buildOpDecorate(Reg: Call->ReturnRegister, MIRBuilder,
3498 Dec: SPIRV::Decoration::SaturatedConversion, DecArgs: {});
3499
3500 if (Builtin->IsRounded) {
3501 bool AnyTypeIsFloat =
3502 GR->isScalarOrVectorOfType(VReg: Call->ReturnRegister, TypeOpcode: SPIRV::OpTypeFloat) ||
3503 GR->isScalarOrVectorOfType(VReg: Call->Arguments[0], TypeOpcode: SPIRV::OpTypeFloat);
3504
3505 // Rounding mode decorations are only valid for floating point types.
3506 // Conversion builtins from integer to integer are equivalent to their
3507 // non-rounded counterparts.
3508 if (AnyTypeIsFloat) {
3509 buildOpDecorate(Reg: Call->ReturnRegister, MIRBuilder,
3510 Dec: SPIRV::Decoration::FPRoundingMode,
3511 DecArgs: {(unsigned)Builtin->RoundingMode});
3512 }
3513 }
3514
3515 MIRBuilder.buildInstr(Opcode)
3516 .addDef(RegNo: Call->ReturnRegister)
3517 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType))
3518 .addUse(RegNo: Call->Arguments[0]);
3519 return true;
3520}
3521
3522static bool generateVectorLoadStoreInst(const SPIRV::IncomingCall *Call,
3523 MachineIRBuilder &MIRBuilder,
3524 SPIRVGlobalRegistry *GR) {
3525 // Lookup the vector load/store builtin in the TableGen records.
3526 const SPIRV::VectorLoadStoreBuiltin *Builtin =
3527 SPIRV::lookupVectorLoadStoreBuiltin(Name: Call->Builtin->name(),
3528 Set: Call->Builtin->Set);
3529 // Build extended instruction.
3530 auto MIB =
3531 MIRBuilder.buildInstr(Opcode: SPIRV::OpExtInst)
3532 .addDef(RegNo: Call->ReturnRegister)
3533 .addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType))
3534 .addImm(Val: static_cast<uint32_t>(SPIRV::InstructionSet::OpenCL_std))
3535 .addImm(Val: Builtin->Number);
3536 for (auto Argument : Call->Arguments)
3537 MIB.addUse(RegNo: Argument);
3538 StringRef BuiltinName = SPIRV::getVectorLoadStoreBuiltinStr(Offset: Builtin->Name);
3539 if (BuiltinName.contains(Other: "load") && Builtin->ElementCount > 1)
3540 MIB.addImm(Val: Builtin->ElementCount);
3541
3542 // Rounding mode should be passed as a last argument in the MI for builtins
3543 // like "vstorea_halfn_r".
3544 if (Builtin->IsRounded)
3545 MIB.addImm(Val: static_cast<uint32_t>(Builtin->RoundingMode));
3546 return true;
3547}
3548
3549static bool generateAFPInst(const SPIRV::IncomingCall *Call,
3550 MachineIRBuilder &MIRBuilder,
3551 SPIRVGlobalRegistry *GR) {
3552 const auto *Builtin = Call->Builtin;
3553 auto *MRI = MIRBuilder.getMRI();
3554 unsigned Opcode =
3555 SPIRV::lookupNativeBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Opcode;
3556 const Type *RetTy = GR->getTypeForSPIRVType(Ty: Call->ReturnType);
3557 bool IsVoid = RetTy->isVoidTy();
3558 auto MIB = MIRBuilder.buildInstr(Opcode);
3559 Register DestReg;
3560 if (IsVoid) {
3561 LLT PtrTy = MRI->getType(Reg: Call->Arguments[0]);
3562 DestReg = MRI->createGenericVirtualRegister(Ty: PtrTy);
3563 MRI->setRegClass(Reg: DestReg, RC: &SPIRV::pIDRegClass);
3564 SPIRVTypeInst PointeeTy =
3565 GR->getPointeeType(PtrType: GR->getSPIRVTypeForVReg(VReg: Call->Arguments[0]));
3566 MIB.addDef(RegNo: DestReg);
3567 MIB.addUse(RegNo: GR->getSPIRVTypeID(SpirvType: PointeeTy));
3568 } else {
3569 MIB.addDef(RegNo: Call->ReturnRegister);
3570 MIB.addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
3571 }
3572 for (unsigned i = IsVoid ? 1 : 0; i < Call->Arguments.size(); ++i) {
3573 Register Arg = Call->Arguments[i];
3574 MachineInstr *DefMI = MRI->getUniqueVRegDef(Reg: Arg);
3575 if (DefMI->getOpcode() == TargetOpcode::G_CONSTANT &&
3576 DefMI->getOperand(i: 1).isCImm()) {
3577 MIB.addImm(Val: getIConstVal(ConstReg: Arg, MRI));
3578 } else {
3579 MIB.addUse(RegNo: Arg);
3580 }
3581 }
3582 if (IsVoid) {
3583 LLT PtrTy = MRI->getType(Reg: Call->Arguments[0]);
3584 MachineMemOperand *MMO = MIRBuilder.getMF().getMachineMemOperand(
3585 PtrInfo: MachinePointerInfo(), F: MachineMemOperand::MOStore,
3586 Size: PtrTy.getSizeInBytes(), BaseAlignment: Align(4));
3587 MIRBuilder.buildStore(Val: DestReg, Addr: Call->Arguments[0], MMO&: *MMO);
3588 }
3589 return true;
3590}
3591
3592static bool generateLoadStoreInst(const SPIRV::IncomingCall *Call,
3593 MachineIRBuilder &MIRBuilder,
3594 SPIRVGlobalRegistry *GR) {
3595 // Lookup the instruction opcode in the TableGen records.
3596 const SPIRV::DemangledBuiltin *Builtin = Call->Builtin;
3597 unsigned Opcode =
3598 SPIRV::lookupNativeBuiltin(Name: Builtin->name(), Set: Builtin->Set)->Opcode;
3599 bool IsLoad = Opcode == SPIRV::OpLoad;
3600 // Build the instruction.
3601 auto MIB = MIRBuilder.buildInstr(Opcode);
3602 if (IsLoad) {
3603 MIB.addDef(RegNo: Call->ReturnRegister);
3604 MIB.addUse(RegNo: GR->getSPIRVTypeID(SpirvType: Call->ReturnType));
3605 }
3606 // Add a pointer to the value to load/store.
3607 MIB.addUse(RegNo: Call->Arguments[0]);
3608 MachineRegisterInfo *MRI = MIRBuilder.getMRI();
3609 // Add a value to store.
3610 if (!IsLoad)
3611 MIB.addUse(RegNo: Call->Arguments[1]);
3612 // Add optional memory attributes and an alignment.
3613 unsigned NumArgs = Call->Arguments.size();
3614 if ((IsLoad && NumArgs >= 2) || NumArgs >= 3)
3615 MIB.addImm(Val: getIConstVal(ConstReg: Call->Arguments[IsLoad ? 1 : 2], MRI));
3616 if ((IsLoad && NumArgs >= 3) || NumArgs >= 4)
3617 MIB.addImm(Val: getIConstVal(ConstReg: Call->Arguments[IsLoad ? 2 : 3], MRI));
3618 return true;
3619}
3620
3621namespace SPIRV {
3622// Try to find a builtin function attributes by a demangled function name and
3623// return a tuple <builtin group, op code, ext instruction number>, or a special
3624// tuple value <-1, 0, 0> if the builtin function is not found.
3625// Not all builtin functions are supported, only those with a ready-to-use op
3626// code or instruction number defined in TableGen.
3627// TODO: consider a major rework of mapping demangled calls into a builtin
3628// functions to unify search and decrease number of individual cases.
3629std::tuple<int, unsigned, unsigned>
3630mapBuiltinToOpcode(StringRef DemangledCall,
3631 SPIRV::InstructionSet::InstructionSet Set) {
3632 Register Reg;
3633 SmallVector<Register> Args;
3634 std::unique_ptr<const IncomingCall> Call =
3635 lookupBuiltin(DemangledCall, Set, ReturnRegister: Reg, ReturnType: nullptr, Arguments: Args);
3636 if (!Call)
3637 return std::make_tuple(args: -1, args: 0, args: 0);
3638
3639 switch (Call->Builtin->Group) {
3640 case SPIRV::Relational:
3641 case SPIRV::Atomic:
3642 case SPIRV::Barrier:
3643 case SPIRV::CastToPtr:
3644 case SPIRV::ImageMiscQuery:
3645 case SPIRV::SpecConstant:
3646 case SPIRV::Enqueue:
3647 case SPIRV::AsyncCopy:
3648 case SPIRV::LoadStore:
3649 case SPIRV::CoopMatr:
3650 case SPIRV::Arithmetic:
3651 if (const auto *R = SPIRV::lookupNativeBuiltin(Name: Call->Builtin->name(),
3652 Set: Call->Builtin->Set))
3653 return std::make_tuple(args: Call->Builtin->Group, args: R->Opcode, args: 0);
3654 break;
3655 case SPIRV::Extended:
3656 if (const auto *R = SPIRV::lookupExtendedBuiltin(Name: Call->Builtin->name(),
3657 Set: Call->Builtin->Set))
3658 return std::make_tuple(args: Call->Builtin->Group, args: 0, args: R->Number);
3659 break;
3660 case SPIRV::VectorLoadStore:
3661 if (const auto *R = SPIRV::lookupVectorLoadStoreBuiltin(
3662 Name: Call->Builtin->name(), Set: Call->Builtin->Set))
3663 return std::make_tuple(args: SPIRV::Extended, args: 0, args: R->Number);
3664 break;
3665 case SPIRV::Group:
3666 if (const auto *R = SPIRV::lookupGroupBuiltin(Name: Call->Builtin->name()))
3667 return std::make_tuple(args: Call->Builtin->Group, args: R->Opcode, args: 0);
3668 break;
3669 case SPIRV::AtomicFloating:
3670 if (const auto *R =
3671 SPIRV::lookupAtomicFloatingBuiltin(Name: Call->Builtin->name()))
3672 return std::make_tuple(args: Call->Builtin->Group, args: R->Opcode, args: 0);
3673 break;
3674 case SPIRV::IntelSubgroups:
3675 if (const auto *R =
3676 SPIRV::lookupIntelSubgroupsBuiltin(Name: Call->Builtin->name()))
3677 return std::make_tuple(args: Call->Builtin->Group, args: R->Opcode, args: 0);
3678 break;
3679 case SPIRV::GroupUniform:
3680 if (const auto *R = SPIRV::lookupGroupUniformBuiltin(Name: Call->Builtin->name()))
3681 return std::make_tuple(args: Call->Builtin->Group, args: R->Opcode, args: 0);
3682 break;
3683 case SPIRV::IntegerDot:
3684 if (const auto *R =
3685 SPIRV::lookupIntegerDotProductBuiltin(Name: Call->Builtin->name()))
3686 return std::make_tuple(args: Call->Builtin->Group, args: R->Opcode, args: 0);
3687 break;
3688 case SPIRV::WriteImage:
3689 return std::make_tuple(args: Call->Builtin->Group, args: SPIRV::OpImageWrite, args: 0);
3690 case SPIRV::Select:
3691 return std::make_tuple(args: Call->Builtin->Group, args: TargetOpcode::G_SELECT, args: 0);
3692 case SPIRV::Construct:
3693 return std::make_tuple(args: Call->Builtin->Group, args: SPIRV::OpCompositeConstruct,
3694 args: 0);
3695 case SPIRV::KernelClock:
3696 return std::make_tuple(args: Call->Builtin->Group, args: SPIRV::OpReadClockKHR, args: 0);
3697 default:
3698 return std::make_tuple(args: -1, args: 0, args: 0);
3699 }
3700 return std::make_tuple(args: -1, args: 0, args: 0);
3701}
3702
3703bool isBuiltin(StringRef DemangledCall,
3704 SPIRV::InstructionSet::InstructionSet Set) {
3705 SmallVector<Register> Args;
3706 return lookupBuiltin(DemangledCall, Set, ReturnRegister: Register(), ReturnType: nullptr, Arguments: Args) !=
3707 nullptr;
3708}
3709
3710/// Checks that scalar/vector numeric arguments of \p Call match the types
3711/// implied by their mangling in \p DemangledCall. Pointers and opaque
3712/// builtin types (images, samplers, pipes, etc.) are not validated here, as
3713/// mangling does not enforce their exact spelling.
3714///
3715/// \returns false if a numeric argument's SPIR-V type disagrees with the
3716/// type implied by the mangled name, true otherwise.
3717static bool demangledArgTypesMatchIR(const SPIRV::IncomingCall *Call,
3718 StringRef DemangledCall,
3719 SPIRVGlobalRegistry *GR, LLVMContext &Ctx,
3720 const CallBase &CB) {
3721 if (Call->isSpirvOp())
3722 return true;
3723
3724 SmallVector<StringRef, 10> ArgTypeStrs;
3725 if (!SPIRV::parseBuiltinTypeStr(BuiltinArgsTypeStrs&: ArgTypeStrs, DemangledCall, Ctx))
3726 return true;
3727
3728 unsigned ArgBase = CB.hasStructRetAttr() ? 1 : 0;
3729 if (Call->Arguments.size() < ArgBase)
3730 return true;
3731 unsigned NumMangledArgs = Call->Arguments.size() - ArgBase;
3732 unsigned NumArgsToCheck =
3733 std::min<unsigned>(a: NumMangledArgs, b: ArgTypeStrs.size());
3734 for (unsigned ArgIdx = 0; ArgIdx < NumArgsToCheck; ++ArgIdx) {
3735 StringRef ArgTypeStr = ArgTypeStrs[ArgIdx].trim();
3736 // Opaque/builtin OpenCL and SPIR-V types (images, samplers, pipes,
3737 // reserve_id, etc.) are not validated here, as mangling does not enforce
3738 // their exact spelling, and some builtin type names have no TableGen
3739 // record and would otherwise abort compilation when parsed.
3740 if (hasBuiltinTypePrefix(Name: ArgTypeStr))
3741 continue;
3742
3743 Type *ExpectedType = SPIRV::parseBuiltinCallArgumentType(TypeStr: ArgTypeStr, Ctx);
3744 if (!ExpectedType || ExpectedType->isVoidTy() ||
3745 ExpectedType->isPointerTy() || ExpectedType->isTargetExtTy())
3746 continue;
3747
3748 SPIRVTypeInst ArgType =
3749 GR->getSPIRVTypeForVReg(VReg: Call->Arguments[ArgIdx + ArgBase]);
3750 if (!ArgType)
3751 continue;
3752 unsigned ArgTypeOpcode = ArgType->getOpcode();
3753 if (ArgTypeOpcode != SPIRV::OpTypeInt &&
3754 ArgTypeOpcode != SPIRV::OpTypeFloat &&
3755 ArgTypeOpcode != SPIRV::OpTypeBool &&
3756 ArgTypeOpcode != SPIRV::OpTypeVector)
3757 continue;
3758
3759 auto *ExpectedVecType = dyn_cast<VectorType>(Val: ExpectedType);
3760 Type *ExpectedScalarType =
3761 ExpectedVecType ? ExpectedVecType->getElementType() : ExpectedType;
3762 SPIRVTypeInst ArgScalarType = GR->getScalarOrVectorComponentType(Type: ArgType);
3763 if (!ArgScalarType)
3764 continue;
3765
3766 bool ExpectedIsInt = ExpectedScalarType->isIntegerTy();
3767 unsigned ArgOpcode = ArgScalarType->getOpcode();
3768 bool ArgIsInt =
3769 ArgOpcode == SPIRV::OpTypeInt || ArgOpcode == SPIRV::OpTypeBool;
3770
3771 if (ExpectedIsInt != ArgIsInt)
3772 return false;
3773
3774 unsigned ExpectedElts =
3775 ExpectedVecType ? ExpectedVecType->getElementCount().getFixedValue()
3776 : 1;
3777 if (ExpectedElts != GR->getScalarOrVectorComponentCount(Type: ArgType))
3778 return false;
3779 }
3780 return true;
3781}
3782
3783std::optional<bool> lowerBuiltin(StringRef DemangledCall,
3784 SPIRV::InstructionSet::InstructionSet Set,
3785 MachineIRBuilder &MIRBuilder,
3786 const Register OrigRet, const Type *OrigRetTy,
3787 const SmallVectorImpl<Register> &Args,
3788 SPIRVGlobalRegistry *GR, const CallBase &CB) {
3789 LLVM_DEBUG(dbgs() << "Lowering builtin call: " << DemangledCall << "\n");
3790
3791 // Lookup the builtin in the TableGen records.
3792 SPIRVTypeInst SpvType = GR->getSPIRVTypeForVReg(VReg: OrigRet);
3793 assert(SpvType && "Inconsistent return register: expected valid type info");
3794 std::unique_ptr<const IncomingCall> Call =
3795 lookupBuiltin(DemangledCall, Set, ReturnRegister: OrigRet, ReturnType: SpvType, Arguments: Args);
3796
3797 if (!Call) {
3798 LLVM_DEBUG(dbgs() << "Builtin record was not found!\n");
3799 return std::nullopt;
3800 }
3801
3802 // Check if the provided args meet the builtin requirements. If not, treat
3803 // the call as a regular function call rather than crashing.
3804 if (Args.size() < Call->Builtin->MinNumArgs) {
3805 LLVM_DEBUG(dbgs() << "Too few arguments for builtin " << DemangledCall
3806 << ": expected at least " << Call->Builtin->MinNumArgs
3807 << ", got " << Args.size()
3808 << "; treating as a normal function\n");
3809 return std::nullopt;
3810 }
3811 if (Call->Builtin->MaxNumArgs && Args.size() > Call->Builtin->MaxNumArgs) {
3812 LLVM_DEBUG(dbgs() << "Too many arguments for builtin " << DemangledCall
3813 << ": expected at most " << Call->Builtin->MaxNumArgs
3814 << ", got " << Args.size()
3815 << "; treating as a normal function\n");
3816 return std::nullopt;
3817 }
3818
3819 // Check that argument types match what the mangling implies. If not
3820 // (e.g. broken mangling), treat the call as a regular function call
3821 // rather than crashing.
3822 if (!demangledArgTypesMatchIR(Call: Call.get(), DemangledCall, GR,
3823 Ctx&: MIRBuilder.getContext(), CB)) {
3824 LLVM_DEBUG(dbgs() << "Argument types do not match mangled types for "
3825 << "builtin " << DemangledCall
3826 << "; treating as a normal function\n");
3827 return std::nullopt;
3828 }
3829
3830 // Match the builtin with implementation based on the grouping.
3831 switch (Call->Builtin->Group) {
3832 case SPIRV::Extended:
3833 return generateExtInst(Call: Call.get(), MIRBuilder, GR, CB);
3834 case SPIRV::Relational:
3835 return generateRelationalInst(Call: Call.get(), MIRBuilder, GR);
3836 case SPIRV::Group:
3837 return generateGroupInst(Call: Call.get(), MIRBuilder, GR);
3838 case SPIRV::Variable:
3839 return generateBuiltinVar(Call: Call.get(), MIRBuilder, GR);
3840 case SPIRV::Atomic:
3841 return generateAtomicInst(Call: Call.get(), MIRBuilder, GR);
3842 case SPIRV::AtomicFloating:
3843 return generateAtomicFloatingInst(Call: Call.get(), MIRBuilder, GR);
3844 case SPIRV::Barrier:
3845 return generateBarrierInst(Call: Call.get(), MIRBuilder, GR);
3846 case SPIRV::CastToPtr:
3847 return generateCastToPtrInst(Call: Call.get(), MIRBuilder, GR);
3848 case SPIRV::Dot:
3849 case SPIRV::IntegerDot:
3850 return generateDotOrFMulInst(DemangledCall, Call: Call.get(), MIRBuilder, GR);
3851 case SPIRV::Wave:
3852 return generateWaveInst(Call: Call.get(), MIRBuilder, GR);
3853 case SPIRV::ICarryBorrow:
3854 return generateICarryBorrowInst(Call: Call.get(), MIRBuilder, GR, CB);
3855 case SPIRV::MulExtended:
3856 return generateMulExtendedInst(Call: Call.get(), MIRBuilder, GR, CB);
3857 case SPIRV::Arithmetic:
3858 return generateArithmeticInst(Call: Call.get(), MIRBuilder, GR);
3859 case SPIRV::GetQuery:
3860 return generateGetQueryInst(Call: Call.get(), MIRBuilder, GR);
3861 case SPIRV::ImageSizeQuery:
3862 return generateImageSizeQueryInst(Call: Call.get(), MIRBuilder, GR);
3863 case SPIRV::ImageMiscQuery:
3864 return generateImageMiscQueryInst(Call: Call.get(), MIRBuilder, GR);
3865 case SPIRV::ReadImage:
3866 return generateReadImageInst(DemangledCall, Call: Call.get(), MIRBuilder, GR);
3867 case SPIRV::WriteImage:
3868 return generateWriteImageInst(Call: Call.get(), MIRBuilder, GR);
3869 case SPIRV::SampleImage:
3870 return generateSampleImageInst(DemangledCall, Call: Call.get(), MIRBuilder, GR);
3871 case SPIRV::Select:
3872 return generateSelectInst(Call: Call.get(), MIRBuilder);
3873 case SPIRV::Construct:
3874 return generateConstructInst(Call: Call.get(), MIRBuilder, GR);
3875 case SPIRV::SpecConstant:
3876 return generateSpecConstantInst(Call: Call.get(), MIRBuilder, GR);
3877 case SPIRV::Enqueue:
3878 return generateEnqueueInst(Call: Call.get(), MIRBuilder, GR, CB);
3879 case SPIRV::AsyncCopy:
3880 return generateAsyncCopy(Call: Call.get(), MIRBuilder, GR, CB);
3881 case SPIRV::Convert:
3882 return generateConvertInst(DemangledCall, Call: Call.get(), MIRBuilder, GR);
3883 case SPIRV::VectorLoadStore:
3884 return generateVectorLoadStoreInst(Call: Call.get(), MIRBuilder, GR);
3885 case SPIRV::LoadStore:
3886 return generateLoadStoreInst(Call: Call.get(), MIRBuilder, GR);
3887 case SPIRV::IntelSubgroups:
3888 return generateIntelSubgroupsInst(Call: Call.get(), MIRBuilder, GR);
3889 case SPIRV::GroupUniform:
3890 return generateGroupUniformInst(Call: Call.get(), MIRBuilder, GR);
3891 case SPIRV::KernelClock:
3892 return generateKernelClockInst(Call: Call.get(), MIRBuilder, GR);
3893 case SPIRV::CoopMatr:
3894 return generateCoopMatrInst(Call: Call.get(), MIRBuilder, GR);
3895 case SPIRV::ExtendedBitOps:
3896 return generateExtendedBitOpsInst(Call: Call.get(), MIRBuilder, GR);
3897 case SPIRV::BindlessINTEL:
3898 return generateBindlessImageINTELInst(Call: Call.get(), MIRBuilder, GR);
3899 case SPIRV::TernaryBitwiseINTEL:
3900 return generateTernaryBitwiseFunctionINTELInst(Call: Call.get(), MIRBuilder, GR);
3901 case SPIRV::Block2DLoadStore:
3902 return generate2DBlockIOINTELInst(Call: Call.get(), MIRBuilder, GR);
3903 case SPIRV::Pipe:
3904 return generatePipeInst(Call: Call.get(), MIRBuilder, GR);
3905 case SPIRV::PredicatedLoadStore:
3906 return generatePredicatedLoadStoreInst(Call: Call.get(), MIRBuilder, GR);
3907 case SPIRV::BlockingPipes:
3908 return generateBlockingPipesInst(Call: Call.get(), MIRBuilder, GR);
3909 case SPIRV::ArbitraryPrecisionFixedPoint:
3910 return generateAPFixedPointInst(Call: Call.get(), MIRBuilder, GR, CB);
3911 case SPIRV::ImageChannelDataTypes:
3912 return generateImageChannelDataTypeInst(Call: Call.get(), MIRBuilder, GR);
3913 case SPIRV::ArbitraryFloatingPoint:
3914 return generateAFPInst(Call: Call.get(), MIRBuilder, GR);
3915 }
3916 return false;
3917}
3918
3919Type *parseBuiltinCallArgumentType(StringRef TypeStr, LLVMContext &Ctx) {
3920 // Parse strings representing OpenCL builtin types.
3921 if (hasBuiltinTypePrefix(Name: TypeStr)) {
3922 // OpenCL builtin types in demangled call strings have the following format:
3923 // e.g. ocl_image2d_ro
3924 [[maybe_unused]] bool IsOCLBuiltinType = TypeStr.consume_front(Prefix: "ocl_");
3925 assert(IsOCLBuiltinType && "Invalid OpenCL builtin prefix");
3926
3927 // Check if this is pointer to a builtin type and not just pointer
3928 // representing a builtin type. In case it is a pointer to builtin type,
3929 // this will require additional handling in the method calling
3930 // parseBuiltinCallArgumentBaseType(...) as this function only retrieves the
3931 // base types.
3932 if (TypeStr.ends_with(Suffix: "*"))
3933 TypeStr = TypeStr.slice(Start: 0, End: TypeStr.find_first_of(Chars: " *"));
3934
3935 return parseBuiltinTypeNameToTargetExtType(TypeName: "opencl." + TypeStr.str() + "_t",
3936 Context&: Ctx);
3937 }
3938
3939 // Parse type name in either "typeN" or "type vector[N]" format, where
3940 // N is the number of elements of the vector.
3941 Type *BaseType;
3942 unsigned VecElts = 0;
3943
3944 BaseType = parseBasicTypeName(TypeName&: TypeStr, Ctx);
3945 if (!BaseType)
3946 // Unable to recognize SPIRV type name.
3947 return nullptr;
3948
3949 // Handle "typeN*" or "type vector[N]*".
3950 TypeStr.consume_back(Suffix: "*");
3951
3952 if (TypeStr.consume_front(Prefix: " vector["))
3953 TypeStr = TypeStr.substr(Start: 0, N: TypeStr.find(C: ']'));
3954
3955 TypeStr.getAsInteger(Radix: 10, Result&: VecElts);
3956 if (VecElts > 0)
3957 BaseType = VectorType::get(
3958 ElementType: BaseType->isVoidTy() ? Type::getInt8Ty(C&: Ctx) : BaseType, NumElements: VecElts, Scalable: false);
3959
3960 return BaseType;
3961}
3962
3963bool parseBuiltinTypeStr(SmallVector<StringRef, 10> &BuiltinArgsTypeStrs,
3964 StringRef DemangledCall, LLVMContext &Ctx) {
3965 auto Pos1 = DemangledCall.find(C: '(');
3966 if (Pos1 == StringRef::npos)
3967 return false;
3968 auto Pos2 = DemangledCall.find(C: ')');
3969 if (Pos2 == StringRef::npos || Pos1 > Pos2)
3970 return false;
3971 DemangledCall.slice(Start: Pos1 + 1, End: Pos2)
3972 .split(A&: BuiltinArgsTypeStrs, Separator: ',', MaxSplit: -1, KeepEmpty: false);
3973 return true;
3974}
3975
3976Type *parseBuiltinCallArgumentBaseType(StringRef DemangledCall, unsigned ArgIdx,
3977 LLVMContext &Ctx) {
3978 SmallVector<StringRef, 10> BuiltinArgsTypeStrs;
3979 parseBuiltinTypeStr(BuiltinArgsTypeStrs, DemangledCall, Ctx);
3980 if (ArgIdx >= BuiltinArgsTypeStrs.size())
3981 return nullptr;
3982 StringRef TypeStr = BuiltinArgsTypeStrs[ArgIdx].trim();
3983 return parseBuiltinCallArgumentType(TypeStr, Ctx);
3984}
3985
3986struct BuiltinType {
3987 StringTable::Offset Name;
3988 uint32_t Opcode;
3989};
3990
3991#define GET_BuiltinTypes_DECL
3992#define GET_BuiltinTypes_IMPL
3993
3994struct OpenCLType {
3995 StringTable::Offset Name;
3996 StringTable::Offset SpirvTypeLiteral;
3997};
3998
3999#define GET_OpenCLTypes_DECL
4000#define GET_OpenCLTypes_IMPL
4001
4002#include "SPIRVGenTables.inc"
4003} // namespace SPIRV
4004
4005//===----------------------------------------------------------------------===//
4006// Misc functions for parsing builtin types.
4007//===----------------------------------------------------------------------===//
4008
4009static Type *parseTypeString(StringRef Name, LLVMContext &Context) {
4010 if (Name.starts_with(Prefix: "void"))
4011 return Type::getVoidTy(C&: Context);
4012 else if (Name.starts_with(Prefix: "int") || Name.starts_with(Prefix: "uint"))
4013 return Type::getInt32Ty(C&: Context);
4014 else if (Name.starts_with(Prefix: "bfloat"))
4015 return Type::getBFloatTy(C&: Context);
4016 else if (Name.starts_with(Prefix: "float"))
4017 return Type::getFloatTy(C&: Context);
4018 else if (Name.starts_with(Prefix: "half"))
4019 return Type::getHalfTy(C&: Context);
4020 else if (Name.starts_with(Prefix: "double"))
4021 return Type::getDoubleTy(C&: Context);
4022 report_fatal_error(reason: "Unable to recognize type!");
4023}
4024
4025//===----------------------------------------------------------------------===//
4026// Implementation functions for builtin types.
4027//===----------------------------------------------------------------------===//
4028
4029static SPIRVTypeInst
4030getNonParameterizedType(const TargetExtType *ExtensionType,
4031 const SPIRV::BuiltinType *TypeRecord,
4032 MachineIRBuilder &MIRBuilder, SPIRVGlobalRegistry *GR) {
4033 unsigned Opcode = TypeRecord->Opcode;
4034 // Create or get an existing type from GlobalRegistry.
4035 return GR->getOrCreateOpTypeByOpcode(Ty: ExtensionType, MIRBuilder, Opcode);
4036}
4037
4038static SPIRVTypeInst getSamplerType(MachineIRBuilder &MIRBuilder,
4039 SPIRVGlobalRegistry *GR) {
4040 // Create or get an existing type from GlobalRegistry.
4041 return GR->getOrCreateOpTypeSampler(MIRBuilder);
4042}
4043
4044static SPIRVTypeInst getPipeType(const TargetExtType *ExtensionType,
4045 MachineIRBuilder &MIRBuilder,
4046 SPIRVGlobalRegistry *GR) {
4047 assert(ExtensionType->getNumIntParameters() == 1 &&
4048 "Invalid number of parameters for SPIR-V pipe builtin!");
4049 // Create or get an existing type from GlobalRegistry.
4050 return GR->getOrCreateOpTypePipe(MIRBuilder,
4051 AccQual: SPIRV::AccessQualifier::AccessQualifier(
4052 ExtensionType->getIntParameter(i: 0)));
4053}
4054
4055static SPIRVTypeInst getCoopMatrType(const TargetExtType *ExtensionType,
4056 MachineIRBuilder &MIRBuilder,
4057 SPIRVGlobalRegistry *GR) {
4058 assert(ExtensionType->getNumIntParameters() == 4 &&
4059 "Invalid number of parameters for SPIR-V coop matrices builtin!");
4060 assert(ExtensionType->getNumTypeParameters() == 1 &&
4061 "SPIR-V coop matrices builtin type must have a type parameter!");
4062 SPIRVTypeInst ElemType =
4063 GR->getOrCreateSPIRVType(Type: ExtensionType->getTypeParameter(i: 0), MIRBuilder,
4064 AQ: SPIRV::AccessQualifier::ReadWrite, EmitIR: true);
4065 // Create or get an existing type from GlobalRegistry.
4066 return GR->getOrCreateOpTypeCoopMatr(
4067 MIRBuilder, ExtensionType, ElemType, Scope: ExtensionType->getIntParameter(i: 0),
4068 Rows: ExtensionType->getIntParameter(i: 1), Columns: ExtensionType->getIntParameter(i: 2),
4069 Use: ExtensionType->getIntParameter(i: 3), EmitIR: true);
4070}
4071
4072static SPIRVTypeInst getSampledImageType(const TargetExtType *OpaqueType,
4073 MachineIRBuilder &MIRBuilder,
4074 SPIRVGlobalRegistry *GR) {
4075 SPIRVTypeInst OpaqueImageType = GR->getImageType(
4076 ExtensionType: OpaqueType, Qualifier: SPIRV::AccessQualifier::ReadOnly, MIRBuilder);
4077 // Create or get an existing type from GlobalRegistry.
4078 return GR->getOrCreateOpTypeSampledImage(ImageType: OpaqueImageType, MIRBuilder);
4079}
4080
4081static SPIRVTypeInst getInlineSpirvType(const TargetExtType *ExtensionType,
4082 MachineIRBuilder &MIRBuilder,
4083 SPIRVGlobalRegistry *GR) {
4084 assert(ExtensionType->getNumIntParameters() == 3 &&
4085 "Inline SPIR-V type builtin takes an opcode, size, and alignment "
4086 "parameter");
4087 auto Opcode = ExtensionType->getIntParameter(i: 0);
4088
4089 SmallVector<MCOperand> Operands;
4090 for (Type *Param : ExtensionType->type_params()) {
4091 if (const TargetExtType *ParamEType = dyn_cast<TargetExtType>(Val: Param)) {
4092 if (ParamEType->getName() == "spirv.IntegralConstant") {
4093 assert(ParamEType->getNumTypeParameters() == 1 &&
4094 "Inline SPIR-V integral constant builtin must have a type "
4095 "parameter");
4096 assert(ParamEType->getNumIntParameters() == 1 &&
4097 "Inline SPIR-V integral constant builtin must have a "
4098 "value parameter");
4099
4100 auto OperandValue = ParamEType->getIntParameter(i: 0);
4101 auto *OperandType = ParamEType->getTypeParameter(i: 0);
4102
4103 SPIRVTypeInst OperandSPIRVType = GR->getOrCreateSPIRVType(
4104 Type: OperandType, MIRBuilder, AQ: SPIRV::AccessQualifier::ReadWrite, EmitIR: true);
4105
4106 Operands.push_back(Elt: MCOperand::createReg(Reg: GR->buildConstantInt(
4107 Val: OperandValue, MIRBuilder, SpvType: OperandSPIRVType, EmitIR: true)));
4108 continue;
4109 } else if (ParamEType->getName() == "spirv.Literal") {
4110 assert(ParamEType->getNumTypeParameters() == 0 &&
4111 "Inline SPIR-V literal builtin does not take type "
4112 "parameters");
4113 assert(ParamEType->getNumIntParameters() == 1 &&
4114 "Inline SPIR-V literal builtin must have an integer "
4115 "parameter");
4116
4117 auto OperandValue = ParamEType->getIntParameter(i: 0);
4118
4119 Operands.push_back(Elt: MCOperand::createImm(Val: OperandValue));
4120 continue;
4121 }
4122 }
4123 SPIRVTypeInst TypeOperand = GR->getOrCreateSPIRVType(
4124 Type: Param, MIRBuilder, AQ: SPIRV::AccessQualifier::ReadWrite, EmitIR: true);
4125 Operands.push_back(Elt: MCOperand::createReg(Reg: GR->getSPIRVTypeID(SpirvType: TypeOperand)));
4126 }
4127
4128 return GR->getOrCreateUnknownType(Ty: ExtensionType, MIRBuilder, Opcode,
4129 Operands);
4130}
4131
4132static SPIRVTypeInst getVulkanBufferType(const TargetExtType *ExtensionType,
4133 MachineIRBuilder &MIRBuilder,
4134 SPIRVGlobalRegistry *GR) {
4135 assert(ExtensionType->getNumTypeParameters() == 1 &&
4136 "Vulkan buffers have exactly one type for the type of the buffer.");
4137 assert(ExtensionType->getNumIntParameters() == 2 &&
4138 "Vulkan buffer have 2 integer parameters: storage class and is "
4139 "writable.");
4140
4141 auto *T = ExtensionType->getTypeParameter(i: 0);
4142 auto SC = static_cast<SPIRV::StorageClass::StorageClass>(
4143 ExtensionType->getIntParameter(i: 0));
4144 bool IsWritable = ExtensionType->getIntParameter(i: 1);
4145 return GR->getOrCreateVulkanBufferType(MIRBuilder, ElemType: T, SC, IsWritable);
4146}
4147
4148static SPIRVTypeInst
4149getVulkanPushConstantType(const TargetExtType *ExtensionType,
4150 MachineIRBuilder &MIRBuilder,
4151 SPIRVGlobalRegistry *GR) {
4152 assert(ExtensionType->getNumTypeParameters() == 1 &&
4153 "Vulkan push constants have exactly one type as argument.");
4154 auto *T = ExtensionType->getTypeParameter(i: 0);
4155 return GR->getOrCreateVulkanPushConstantType(MIRBuilder, ElemType: T);
4156}
4157
4158static SPIRVTypeInst getLayoutType(const TargetExtType *ExtensionType,
4159 MachineIRBuilder &MIRBuilder,
4160 SPIRVGlobalRegistry *GR) {
4161 return GR->getOrCreateLayoutType(MIRBuilder, T: ExtensionType);
4162}
4163
4164namespace SPIRV {
4165TargetExtType *parseBuiltinTypeNameToTargetExtType(std::string TypeName,
4166 LLVMContext &Context) {
4167 StringRef NameWithParameters = TypeName;
4168
4169 // Pointers-to-opaque-structs representing OpenCL types are first translated
4170 // to equivalent SPIR-V types. OpenCL builtin type names should have the
4171 // following format: e.g. %opencl.event_t
4172 if (NameWithParameters.starts_with(Prefix: "opencl.")) {
4173 const SPIRV::OpenCLType *OCLTypeRecord =
4174 SPIRV::lookupOpenCLType(Name: NameWithParameters);
4175 if (!OCLTypeRecord)
4176 report_fatal_error(reason: "Missing TableGen record for OpenCL type: " +
4177 NameWithParameters);
4178 NameWithParameters =
4179 SPIRV::getOpenCLTypeStr(Offset: OCLTypeRecord->SpirvTypeLiteral);
4180 // Continue with the SPIR-V builtin type...
4181 }
4182
4183 // Names of the opaque structs representing a SPIR-V builtins without
4184 // parameters should have the following format: e.g. %spirv.Event
4185 assert(NameWithParameters.starts_with("spirv.") &&
4186 "Unknown builtin opaque type!");
4187
4188 // Parameterized SPIR-V builtins names follow this format:
4189 // e.g. %spirv.Image._void_1_0_0_0_0_0_0, %spirv.Pipe._0
4190 if (!NameWithParameters.contains(C: '_'))
4191 return TargetExtType::get(Context, Name: NameWithParameters);
4192
4193 SmallVector<StringRef> Parameters;
4194 unsigned BaseNameLength = NameWithParameters.find(C: '_') - 1;
4195 SplitString(Source: NameWithParameters.substr(Start: BaseNameLength + 1), OutFragments&: Parameters, Delimiters: "_");
4196
4197 SmallVector<Type *, 1> TypeParameters;
4198 bool HasTypeParameter = !isDigit(C: Parameters[0][0]);
4199 if (HasTypeParameter)
4200 TypeParameters.push_back(Elt: parseTypeString(Name: Parameters[0], Context));
4201 SmallVector<unsigned> IntParameters;
4202 for (unsigned i = HasTypeParameter ? 1 : 0; i < Parameters.size(); i++) {
4203 unsigned IntParameter = 0;
4204 bool ValidLiteral = !Parameters[i].getAsInteger(Radix: 10, Result&: IntParameter);
4205 (void)ValidLiteral;
4206 assert(ValidLiteral &&
4207 "Invalid format of SPIR-V builtin parameter literal!");
4208 IntParameters.push_back(Elt: IntParameter);
4209 }
4210 return TargetExtType::get(Context,
4211 Name: NameWithParameters.substr(Start: 0, N: BaseNameLength),
4212 Types: TypeParameters, Ints: IntParameters);
4213}
4214
4215SPIRVTypeInst
4216lowerBuiltinType(const Type *OpaqueType,
4217 SPIRV::AccessQualifier::AccessQualifier AccessQual,
4218 MachineIRBuilder &MIRBuilder, SPIRVGlobalRegistry *GR) {
4219 // In LLVM IR, SPIR-V and OpenCL builtin types are represented as either
4220 // target(...) target extension types or pointers-to-opaque-structs. The
4221 // approach relying on structs is deprecated and works only in the non-opaque
4222 // pointer mode (-opaque-pointers=0).
4223 // In order to maintain compatibility with LLVM IR generated by older versions
4224 // of Clang and LLVM/SPIR-V Translator, the pointers-to-opaque-structs are
4225 // "translated" to target extension types. This translation is temporary and
4226 // will be removed in the future release of LLVM.
4227 const TargetExtType *BuiltinType = dyn_cast<TargetExtType>(Val: OpaqueType);
4228 if (!BuiltinType)
4229 BuiltinType = parseBuiltinTypeNameToTargetExtType(
4230 TypeName: OpaqueType->getStructName().str(), Context&: MIRBuilder.getContext());
4231
4232 unsigned NumStartingVRegs = MIRBuilder.getMRI()->getNumVirtRegs();
4233
4234 StringRef Name = BuiltinType->getName();
4235 LLVM_DEBUG(dbgs() << "Lowering builtin type: " << Name << "\n");
4236
4237 SPIRVTypeInst TargetType = nullptr;
4238 if (Name == "spirv.Type") {
4239 TargetType = getInlineSpirvType(ExtensionType: BuiltinType, MIRBuilder, GR);
4240 } else if (Name == "spirv.VulkanBuffer") {
4241 TargetType = getVulkanBufferType(ExtensionType: BuiltinType, MIRBuilder, GR);
4242 } else if (Name == "spirv.Padding") {
4243 TargetType = GR->getOrCreatePaddingType(MIRBuilder);
4244 } else if (Name == "spirv.PushConstant") {
4245 TargetType = getVulkanPushConstantType(ExtensionType: BuiltinType, MIRBuilder, GR);
4246 } else if (Name == "spirv.Layout") {
4247 TargetType = getLayoutType(ExtensionType: BuiltinType, MIRBuilder, GR);
4248 } else {
4249 // Lookup the demangled builtin type in the TableGen records.
4250 const SPIRV::BuiltinType *TypeRecord = SPIRV::lookupBuiltinType(Name);
4251 if (!TypeRecord)
4252 report_fatal_error(reason: "Missing TableGen record for builtin type: " + Name);
4253
4254 // "Lower" the BuiltinType into TargetType. The following get<...>Type
4255 // methods use the implementation details from TableGen records or
4256 // TargetExtType parameters to either create a new OpType<...> machine
4257 // instruction or get an existing equivalent SPIRV type from
4258 // GlobalRegistry.
4259
4260 switch (TypeRecord->Opcode) {
4261 case SPIRV::OpTypeImage:
4262 TargetType = GR->getImageType(ExtensionType: BuiltinType, Qualifier: AccessQual, MIRBuilder);
4263 break;
4264 case SPIRV::OpTypePipe:
4265 TargetType = getPipeType(ExtensionType: BuiltinType, MIRBuilder, GR);
4266 break;
4267 case SPIRV::OpTypeDeviceEvent:
4268 TargetType = GR->getOrCreateOpTypeDeviceEvent(MIRBuilder);
4269 break;
4270 case SPIRV::OpTypeSampler:
4271 TargetType = getSamplerType(MIRBuilder, GR);
4272 break;
4273 case SPIRV::OpTypeSampledImage:
4274 TargetType = getSampledImageType(OpaqueType: BuiltinType, MIRBuilder, GR);
4275 break;
4276 case SPIRV::OpTypeCooperativeMatrixKHR:
4277 TargetType = getCoopMatrType(ExtensionType: BuiltinType, MIRBuilder, GR);
4278 break;
4279 default:
4280 TargetType =
4281 getNonParameterizedType(ExtensionType: BuiltinType, TypeRecord, MIRBuilder, GR);
4282 break;
4283 }
4284 }
4285
4286 // Emit OpName instruction if a new OpType<...> instruction was added
4287 // (equivalent type was not found in GlobalRegistry).
4288 if (NumStartingVRegs < MIRBuilder.getMRI()->getNumVirtRegs())
4289 buildOpName(Target: GR->getSPIRVTypeID(SpirvType: TargetType), Name, MIRBuilder);
4290
4291 return TargetType;
4292}
4293
4294bool isPipeOrAddressSpaceCastBuiltin(StringRef Name) {
4295 const DemangledBuiltin *Builtin = lookupBuiltin(Name, Set: OpenCL_std);
4296 if (!Builtin)
4297 return false;
4298 return Builtin->Group == Pipe || Builtin->Group == CastToPtr ||
4299 Builtin->Group == BlockingPipes;
4300}
4301} // namespace SPIRV
4302} // namespace llvm
4303