1//===--- SPIRVUtils.cpp ---- SPIR-V Utility 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 contains miscellaneous utility functions.
10//
11//===----------------------------------------------------------------------===//
12
13#include "SPIRVUtils.h"
14#include "MCTargetDesc/SPIRVBaseInfo.h"
15#include "SPIRV.h"
16#include "SPIRVBuiltins.h"
17#include "SPIRVGlobalRegistry.h"
18#include "SPIRVInstrInfo.h"
19#include "SPIRVSubtarget.h"
20#include "llvm/ADT/STLExtras.h"
21#include "llvm/ADT/StringRef.h"
22#include "llvm/CodeGen/GlobalISel/GenericMachineInstrs.h"
23#include "llvm/CodeGen/GlobalISel/MachineIRBuilder.h"
24#include "llvm/CodeGen/MachineInstr.h"
25#include "llvm/CodeGen/MachineInstrBuilder.h"
26#include "llvm/Demangle/Demangle.h"
27#include "llvm/IR/IntrinsicInst.h"
28#include "llvm/IR/IntrinsicsSPIRV.h"
29#include "llvm/Support/MathExtras.h"
30#include "llvm/TargetParser/AtomicScope.h"
31#include <queue>
32#include <vector>
33
34namespace llvm {
35namespace SPIRV {
36static MDNode *findNamedMDOperand(NamedMDNode *NMD, StringRef Name) {
37 auto It = find_if(Range: NMD->operands(), P: [Name](MDNode *N) {
38 if (auto *MDS = dyn_cast_or_null<MDString>(Val: N->getOperand(I: 0)))
39 return MDS->getString() == Name;
40 return false;
41 });
42 return It == NMD->op_end() ? nullptr : *It;
43}
44
45// This code restores function args/retvalue types for composite cases
46// because the final types should still be aggregate whereas they're i32
47// during the translation to cope with aggregate flattening etc.
48// TODO: should these just return nullptr when there's no metadata?
49static FunctionType *extractFunctionTypeFromMetadata(NamedMDNode *NMD,
50 FunctionType *FTy,
51 StringRef Name) {
52 if (!NMD)
53 return FTy;
54
55 MDNode *Match = findNamedMDOperand(NMD, Name);
56 if (!Match)
57 return FTy;
58
59 Type *RetTy = FTy->getReturnType();
60 SmallVector<Type *, 4> PTys(FTy->params());
61
62 for (unsigned I = 1; I != Match->getNumOperands(); ++I) {
63 MDNode *MD = dyn_cast<MDNode>(Val: Match->getOperand(I));
64 assert(MD && "MDNode operand is expected");
65
66 if (auto *Const = getMDOperandAsConstInt(N: MD, I: 0)) {
67 auto *CMeta = dyn_cast<ConstantAsMetadata>(Val: MD->getOperand(I: 1));
68 assert(CMeta && "ConstantAsMetadata operand is expected");
69 int64_t Idx = Const->getSExtValue();
70 // Currently -1 indicates return value, greater values mean
71 // argument numbers.
72 if (Idx == -1) {
73 RetTy = CMeta->getType();
74 continue;
75 }
76 if (Idx >= 0 && static_cast<uint64_t>(Idx) < PTys.size()) {
77 PTys[Idx] = CMeta->getType();
78 continue;
79 }
80 report_fatal_error(reason: "invalid argument index in function type metadata");
81 }
82 }
83
84 return FunctionType::get(Result: RetTy, Params: PTys, isVarArg: FTy->isVarArg());
85}
86
87static StringRef extractAsmConstraintsFromMetadata(NamedMDNode *NMD,
88 StringRef Constraints,
89 StringRef Name) {
90 if (!NMD)
91 return Constraints;
92
93 MDNode *Match = findNamedMDOperand(NMD, Name);
94 if (!Match)
95 return Constraints;
96
97 // By convention, the constraints string is stored in the final MD operand.
98 MDNode *MD = dyn_cast<MDNode>(Val: Match->getOperand(I: Match->getNumOperands() - 1));
99 assert(MD && "MDNode operand is expected");
100
101 if (auto *MDS = dyn_cast<MDString>(Val: MD->getOperand(I: 0)))
102 Constraints = MDS->getString();
103
104 return Constraints;
105}
106
107FunctionType *getOriginalFunctionType(const Function &F) {
108 return extractFunctionTypeFromMetadata(
109 NMD: F.getParent()->getNamedMetadata(Name: "spv.cloned_funcs"), FTy: F.getFunctionType(),
110 Name: F.getName());
111}
112
113// Keyed via instruction metadata, not a name.
114static std::optional<StringRef> getMutatedCallsiteKey(const CallBase &CB) {
115 if (MDNode *MD = CB.getMetadata(Kind: "spv.mutated_callsite"))
116 if (MD->getNumOperands() > 0)
117 if (auto *MDS = dyn_cast<MDString>(Val: MD->getOperand(I: 0)))
118 return MDS->getString();
119 return std::nullopt;
120}
121
122FunctionType *getOriginalFunctionType(const CallBase &CB) {
123 std::optional<StringRef> Key = getMutatedCallsiteKey(CB);
124 if (!Key)
125 return CB.getFunctionType();
126 return extractFunctionTypeFromMetadata(
127 NMD: CB.getModule()->getNamedMetadata(Name: "spv.mutated_callsites"),
128 FTy: CB.getFunctionType(), Name: *Key);
129}
130
131StringRef getOriginalAsmConstraints(const CallBase &CB) {
132 StringRef Constraints =
133 cast<InlineAsm>(Val: CB.getCalledOperand())->getConstraintString();
134 std::optional<StringRef> Key = getMutatedCallsiteKey(CB);
135 if (!Key)
136 return Constraints;
137 return extractAsmConstraintsFromMetadata(
138 NMD: CB.getModule()->getNamedMetadata(Name: "spv.mutated_callsites"), Constraints,
139 Name: *Key);
140}
141} // Namespace SPIRV
142
143// The following functions are used to add these string literals as a series of
144// 32-bit integer operands with the correct format, and unpack them if necessary
145// when making string comparisons in compiler passes.
146// SPIR-V requires null-terminated UTF-8 strings padded to 32-bit alignment.
147static uint32_t convertCharsToWord(StringRef Str, unsigned i) {
148 uint32_t Word = 0u; // Build up this 32-bit word from 4 8-bit chars.
149 for (unsigned WordIndex = 0; WordIndex < 4; ++WordIndex) {
150 unsigned StrIndex = i + WordIndex;
151 uint8_t CharToAdd = 0; // Initilize char as padding/null.
152 if (StrIndex < Str.size()) { // If it's within the string, get a real char.
153 CharToAdd = Str[StrIndex];
154 }
155 Word |= (CharToAdd << (WordIndex * 8));
156 }
157 return Word;
158}
159
160// Get length including padding and null terminator.
161static size_t getPaddedLen(StringRef Str) { return alignTo(Value: Str.size() + 1, Align: 4); }
162
163void addStringImm(StringRef Str, MCInst &Inst) {
164 const size_t PaddedLen = getPaddedLen(Str);
165 for (unsigned i = 0; i < PaddedLen; i += 4) {
166 // Add an operand for the 32-bits of chars or padding.
167 Inst.addOperand(Op: MCOperand::createImm(Val: convertCharsToWord(Str, i)));
168 }
169}
170
171void addStringImm(StringRef Str, MachineInstrBuilder &MIB) {
172 const size_t PaddedLen = getPaddedLen(Str);
173 for (unsigned i = 0; i < PaddedLen; i += 4) {
174 // Add an operand for the 32-bits of chars or padding.
175 MIB.addImm(Val: convertCharsToWord(Str, i));
176 }
177}
178
179std::string getStringImm(const MachineInstr &MI, unsigned StartIndex) {
180 return getSPIRVStringOperand(MI, StartIndex);
181}
182
183std::string getStringValueFromReg(Register Reg, MachineRegisterInfo &MRI) {
184 MachineInstr *Def = getVRegDef(MRI, Reg);
185 assert(Def && Def->getOpcode() == TargetOpcode::G_GLOBAL_VALUE &&
186 "Expected G_GLOBAL_VALUE");
187 const GlobalValue *GV = Def->getOperand(i: 1).getGlobal();
188 Value *V = GV->getOperand(i: 0);
189 const ConstantDataArray *CDA = cast<ConstantDataArray>(Val: V);
190 return CDA->getAsCString().str();
191}
192
193void addNumImm(const APInt &Imm, MachineInstrBuilder &MIB) {
194 const auto Bitwidth = Imm.getBitWidth();
195 if (Bitwidth == 1)
196 return; // Already handled
197 else if (Bitwidth <= 32) {
198 MIB.addImm(Val: Imm.getZExtValue());
199 // Asm Printer needs this info to print floating-type correctly
200 if (Bitwidth == 16)
201 MIB.getInstr()->setAsmPrinterFlag(SPIRV::ASM_PRINTER_WIDTH16);
202 return;
203 } else if (Bitwidth <= 64) {
204 uint64_t FullImm = Imm.getZExtValue();
205 MIB.addImm(Val: Lo_32(Value: FullImm)).addImm(Val: Hi_32(Value: FullImm));
206 // Asm Printer needs this info to print 64-bit operands correctly
207 MIB.getInstr()->setAsmPrinterFlag(SPIRV::ASM_PRINTER_WIDTH64);
208 return;
209 } else {
210 // Emit ceil(Bitwidth / 32) words to conform SPIR-V spec.
211 unsigned NumWords = divideCeil(Numerator: Bitwidth, Denominator: 32);
212 for (unsigned I = 0; I < NumWords; ++I) {
213 unsigned LimbIdx = I / 2;
214 unsigned LimbShift = (I % 2) * 32;
215 uint32_t Word = (Imm.getRawData()[LimbIdx] >> LimbShift) & 0xffffffff;
216 MIB.addImm(Val: Word);
217 }
218 return;
219 }
220}
221
222void buildOpName(Register Target, StringRef Name,
223 MachineIRBuilder &MIRBuilder) {
224 if (!Name.empty()) {
225 auto MIB = MIRBuilder.buildInstr(Opcode: SPIRV::OpName).addUse(RegNo: Target);
226 addStringImm(Str: Name, MIB);
227 }
228}
229
230void buildOpName(Register Target, StringRef Name, MachineInstr &I,
231 const SPIRVInstrInfo &TII) {
232 if (!Name.empty()) {
233 auto MIB =
234 BuildMI(BB&: *I.getParent(), I, MIMD: I.getDebugLoc(), MCID: TII.get(Opcode: SPIRV::OpName))
235 .addUse(RegNo: Target);
236 addStringImm(Str: Name, MIB);
237 }
238}
239
240static void finishBuildOpDecorate(MachineInstrBuilder &MIB,
241 ArrayRef<uint32_t> DecArgs,
242 StringRef StrImm) {
243 if (!StrImm.empty())
244 addStringImm(Str: StrImm, MIB);
245 for (const auto &DecArg : DecArgs)
246 MIB.addImm(Val: DecArg);
247}
248
249void buildOpDecorate(Register Reg, MachineIRBuilder &MIRBuilder,
250 SPIRV::Decoration::Decoration Dec,
251 ArrayRef<uint32_t> DecArgs, StringRef StrImm) {
252 auto MIB = MIRBuilder.buildInstr(Opcode: SPIRV::OpDecorate)
253 .addUse(RegNo: Reg)
254 .addImm(Val: static_cast<uint32_t>(Dec));
255 finishBuildOpDecorate(MIB, DecArgs, StrImm);
256}
257
258void buildOpDecorate(Register Reg, MachineInstr &I, const SPIRVInstrInfo &TII,
259 SPIRV::Decoration::Decoration Dec,
260 ArrayRef<uint32_t> DecArgs, StringRef StrImm) {
261 MachineBasicBlock &MBB = *I.getParent();
262 auto MIB = BuildMI(BB&: MBB, I, MIMD: I.getDebugLoc(), MCID: TII.get(Opcode: SPIRV::OpDecorate))
263 .addUse(RegNo: Reg)
264 .addImm(Val: static_cast<uint32_t>(Dec));
265 finishBuildOpDecorate(MIB, DecArgs, StrImm);
266}
267
268void buildOpMemberDecorate(Register Reg, MachineIRBuilder &MIRBuilder,
269 SPIRV::Decoration::Decoration Dec, uint32_t Member,
270 ArrayRef<uint32_t> DecArgs, StringRef StrImm) {
271 auto MIB = MIRBuilder.buildInstr(Opcode: SPIRV::OpMemberDecorate)
272 .addUse(RegNo: Reg)
273 .addImm(Val: Member)
274 .addImm(Val: static_cast<uint32_t>(Dec));
275 finishBuildOpDecorate(MIB, DecArgs, StrImm);
276}
277
278void buildOpSpirvDecorations(Register Reg, MachineIRBuilder &MIRBuilder,
279 const MDNode *GVarMD, const SPIRVSubtarget &ST) {
280 for (unsigned I = 0, E = GVarMD->getNumOperands(); I != E; ++I) {
281 auto *OpMD = dyn_cast<MDNode>(Val: GVarMD->getOperand(I));
282 if (!OpMD)
283 report_fatal_error(reason: "Invalid decoration");
284 if (OpMD->getNumOperands() == 0)
285 report_fatal_error(reason: "Expect operand(s) of the decoration");
286 ConstantInt *DecorationId =
287 mdconst::dyn_extract<ConstantInt>(MD: OpMD->getOperand(I: 0));
288 if (!DecorationId)
289 report_fatal_error(reason: "Expect SPIR-V <Decoration> operand to be the first "
290 "element of the decoration");
291
292 // The goal of `spirv.Decorations` metadata is to provide a way to
293 // represent SPIR-V entities that do not map to LLVM in an obvious way.
294 // FP flags do have obvious matches between LLVM IR and SPIR-V.
295 // Additionally, we have no guarantee at this point that the flags passed
296 // through the decoration are not violated already in the optimizer passes.
297 // Therefore, we simply ignore FP flags, including NoContraction, and
298 // FPFastMathMode.
299 if (DecorationId->getZExtValue() ==
300 static_cast<uint32_t>(SPIRV::Decoration::NoContraction) ||
301 DecorationId->getZExtValue() ==
302 static_cast<uint32_t>(SPIRV::Decoration::FPFastMathMode)) {
303 continue; // Ignored.
304 }
305 uint32_t Dec = static_cast<uint32_t>(DecorationId->getZExtValue());
306 if (Dec == static_cast<uint32_t>(SPIRV::Decoration::UniformId) ||
307 Dec == static_cast<uint32_t>(SPIRV::Decoration::AlignmentId) ||
308 Dec == static_cast<uint32_t>(SPIRV::Decoration::MaxByteOffsetId)) {
309 ConstantInt *IdV =
310 OpMD->getNumOperands() == 2
311 ? mdconst::dyn_extract<ConstantInt>(MD: OpMD->getOperand(I: 1))
312 : nullptr;
313 if (!IdV || !isUInt<32>(x: IdV->getZExtValue()))
314 report_fatal_error(reason: "Expect a single integer <id> operand of the "
315 "decoration");
316 SPIRVGlobalRegistry *GR = ST.getSPIRVGlobalRegistry();
317 SPIRVTypeInst SpvTypeInt32 =
318 GR->getOrCreateSPIRVIntegerType(BitWidth: 32, MIRBuilder);
319 Register IdReg = GR->buildConstantInt(Val: IdV->getZExtValue(), MIRBuilder,
320 SpvType: SpvTypeInt32, /*EmitIR=*/false);
321 MIRBuilder.buildInstr(Opcode: SPIRV::OpDecorateId)
322 .addUse(RegNo: Reg)
323 .addImm(Val: Dec)
324 .addUse(RegNo: IdReg);
325 continue;
326 }
327 auto MIB = MIRBuilder.buildInstr(Opcode: SPIRV::OpDecorate).addUse(RegNo: Reg).addImm(Val: Dec);
328 for (unsigned OpI = 1, OpE = OpMD->getNumOperands(); OpI != OpE; ++OpI) {
329 if (ConstantInt *OpV =
330 mdconst::dyn_extract<ConstantInt>(MD: OpMD->getOperand(I: OpI)))
331 MIB.addImm(Val: static_cast<uint32_t>(OpV->getZExtValue()));
332 else if (MDString *OpV = dyn_cast<MDString>(Val: OpMD->getOperand(I: OpI)))
333 addStringImm(Str: OpV->getString(), MIB);
334 else
335 report_fatal_error(reason: "Unexpected operand of the decoration");
336 }
337 }
338}
339
340MachineBasicBlock::iterator getOpVariableMBBIt(MachineFunction &MF) {
341 MachineBasicBlock &MBB = MF.front();
342 // Find the position to insert the OpVariable instruction.
343 // We will insert it after the last OpFunctionParameter, if any, or
344 // after OpFunction otherwise.
345 auto IsPreamble = [](const MachineInstr &MI) {
346 switch (MI.getOpcode()) {
347 case SPIRV::OpFunction:
348 case SPIRV::OpFunctionParameter:
349 case SPIRV::OpLabel:
350 case SPIRV::ASSIGN_TYPE:
351 return true;
352 default:
353 return false;
354 }
355 };
356 MachineBasicBlock::iterator VarPos = MBB.SkipPHIsAndLabels(I: MBB.begin());
357 while (VarPos != MBB.end() && VarPos->getOpcode() != SPIRV::OpFunction)
358 ++VarPos;
359 // Advance past the preamble.
360 while (VarPos != MBB.end() && IsPreamble(*VarPos))
361 ++VarPos;
362 return VarPos;
363}
364
365MachineBasicBlock::iterator getInsertPtValidEnd(MachineBasicBlock *MBB) {
366 MachineBasicBlock::iterator I = MBB->end();
367 if (I == MBB->begin())
368 return I;
369 --I;
370 while (I->isTerminator() || I->isDebugValue()) {
371 if (I == MBB->begin())
372 break;
373 --I;
374 }
375 return I;
376}
377
378SPIRV::StorageClass::StorageClass
379addressSpaceToStorageClass(unsigned AddrSpace, const SPIRVSubtarget &STI) {
380 switch (AddrSpace) {
381 case 0:
382 return SPIRV::StorageClass::Function;
383 case 1:
384 return SPIRV::StorageClass::CrossWorkgroup;
385 case 2:
386 return SPIRV::StorageClass::UniformConstant;
387 case 3:
388 return SPIRV::StorageClass::Workgroup;
389 case 4:
390 return SPIRV::StorageClass::Generic;
391 case 5:
392 return STI.canUseExtension(E: SPIRV::Extension::SPV_INTEL_usm_storage_classes)
393 ? SPIRV::StorageClass::DeviceOnlyINTEL
394 : SPIRV::StorageClass::CrossWorkgroup;
395 case 6:
396 return STI.canUseExtension(E: SPIRV::Extension::SPV_INTEL_usm_storage_classes)
397 ? SPIRV::StorageClass::HostOnlyINTEL
398 : SPIRV::StorageClass::CrossWorkgroup;
399 case 7:
400 return SPIRV::StorageClass::Input;
401 case 8:
402 return SPIRV::StorageClass::Output;
403 case 9:
404 return SPIRV::StorageClass::CodeSectionINTEL;
405 case 10:
406 return SPIRV::StorageClass::Private;
407 case 11:
408 return SPIRV::StorageClass::StorageBuffer;
409 case 12:
410 return SPIRV::StorageClass::Uniform;
411 case 13:
412 return SPIRV::StorageClass::PushConstant;
413 default:
414 report_fatal_error(reason: "Unknown address space");
415 }
416}
417
418SPIRV::MemorySemantics::MemorySemantics
419getMemSemanticsForStorageClass(SPIRV::StorageClass::StorageClass SC) {
420 switch (SC) {
421 case SPIRV::StorageClass::StorageBuffer:
422 case SPIRV::StorageClass::Uniform:
423 return SPIRV::MemorySemantics::UniformMemory;
424 case SPIRV::StorageClass::Workgroup:
425 return SPIRV::MemorySemantics::WorkgroupMemory;
426 case SPIRV::StorageClass::CrossWorkgroup:
427 return SPIRV::MemorySemantics::CrossWorkgroupMemory;
428 case SPIRV::StorageClass::Generic:
429 return SPIRV::MemorySemantics::MemorySemantics(
430 SPIRV::MemorySemantics::WorkgroupMemory |
431 SPIRV::MemorySemantics::CrossWorkgroupMemory);
432 case SPIRV::StorageClass::AtomicCounter:
433 return SPIRV::MemorySemantics::AtomicCounterMemory;
434 case SPIRV::StorageClass::Image:
435 return SPIRV::MemorySemantics::ImageMemory;
436 default:
437 return SPIRV::MemorySemantics::None;
438 }
439}
440
441SPIRV::MemorySemantics::MemorySemantics getMemSemantics(AtomicOrdering Ord) {
442 switch (Ord) {
443 case AtomicOrdering::Acquire:
444 return SPIRV::MemorySemantics::Acquire;
445 case AtomicOrdering::Release:
446 return SPIRV::MemorySemantics::Release;
447 case AtomicOrdering::AcquireRelease:
448 return SPIRV::MemorySemantics::AcquireRelease;
449 case AtomicOrdering::SequentiallyConsistent:
450 return SPIRV::MemorySemantics::SequentiallyConsistent;
451 case AtomicOrdering::Unordered:
452 case AtomicOrdering::Monotonic:
453 case AtomicOrdering::NotAtomic:
454 return SPIRV::MemorySemantics::None;
455 }
456 llvm_unreachable(nullptr);
457}
458
459uint32_t getMemSemanticsWithStorageClass(const Triple &TT, uint32_t OrderSem,
460 uint32_t StorageClassSem) {
461 bool DropStorageClass =
462 TT.isVulkanOS() &&
463 OrderSem == static_cast<uint32_t>(SPIRV::MemorySemantics::None);
464 return OrderSem | (DropStorageClass ? 0 : StorageClassSem);
465}
466
467SPIRV::Scope::Scope getMemScope(const Triple &TT, LLVMContext &Ctx,
468 SyncScope::ID Id) {
469 // Named by
470 // https://registry.khronos.org/SPIR-V/specs/unified1/SPIRV.html#_scope_id.
471 // We don't need aliases for Invocation and CrossDevice, as we already have
472 // them covered by "singlethread" and "" strings respectively (see
473 // implementation of LLVMContext::LLVMContext()).
474 auto ScopeID = [&](AtomicScope Scope) {
475 return Ctx.getOrInsertSyncScopeID(SSN: *getAtomicScopeIRString(T: TT, S: Scope));
476 };
477 static const llvm::SyncScope::ID SubGroup = ScopeID(AtomicScope::Wavefront);
478 static const llvm::SyncScope::ID WorkGroup = ScopeID(AtomicScope::Workgroup);
479 static const llvm::SyncScope::ID Device = ScopeID(AtomicScope::Device);
480
481 if (Id == llvm::SyncScope::SingleThread)
482 return SPIRV::Scope::Invocation;
483 else if (Id == llvm::SyncScope::System)
484 return SPIRV::Scope::CrossDevice;
485 else if (Id == SubGroup)
486 return SPIRV::Scope::Subgroup;
487 else if (Id == WorkGroup)
488 return SPIRV::Scope::Workgroup;
489 else if (Id == Device)
490 return SPIRV::Scope::Device;
491 return SPIRV::Scope::CrossDevice;
492}
493
494MachineInstr *getDefInstrMaybeConstant(Register &ConstReg,
495 const MachineRegisterInfo *MRI) {
496 MachineInstr *MI = MRI->getVRegDef(Reg: ConstReg);
497 MachineInstr *ConstInstr =
498 MI->getOpcode() == SPIRV::G_TRUNC || MI->getOpcode() == SPIRV::G_ZEXT
499 ? MRI->getVRegDef(Reg: MI->getOperand(i: 1).getReg())
500 : MI;
501 if (auto *GI = dyn_cast<GIntrinsic>(Val: ConstInstr)) {
502 if (GI->is(ID: Intrinsic::spv_track_constant)) {
503 ConstReg = ConstInstr->getOperand(i: 2).getReg();
504 return MRI->getVRegDef(Reg: ConstReg);
505 }
506 } else if (ConstInstr->getOpcode() == SPIRV::ASSIGN_TYPE) {
507 ConstReg = ConstInstr->getOperand(i: 1).getReg();
508 return MRI->getVRegDef(Reg: ConstReg);
509 } else if (ConstInstr->getOpcode() == TargetOpcode::G_CONSTANT ||
510 ConstInstr->getOpcode() == TargetOpcode::G_FCONSTANT) {
511 ConstReg = ConstInstr->getOperand(i: 0).getReg();
512 return ConstInstr;
513 }
514 return MRI->getVRegDef(Reg: ConstReg);
515}
516
517uint64_t getIConstVal(Register ConstReg, const MachineRegisterInfo *MRI) {
518 const MachineInstr *MI = getDefInstrMaybeConstant(ConstReg, MRI);
519 assert(MI && MI->getOpcode() == TargetOpcode::G_CONSTANT);
520 return MI->getOperand(i: 1).getCImm()->getValue().getZExtValue();
521}
522
523int64_t getIConstValSext(Register ConstReg, const MachineRegisterInfo *MRI) {
524 const MachineInstr *MI = getDefInstrMaybeConstant(ConstReg, MRI);
525 assert(MI && MI->getOpcode() == TargetOpcode::G_CONSTANT);
526 return MI->getOperand(i: 1).getCImm()->getSExtValue();
527}
528
529bool isSpvIntrinsic(const MachineInstr &MI, Intrinsic::ID IntrinsicID) {
530 if (const auto *GI = dyn_cast<GIntrinsic>(Val: &MI))
531 return GI->is(ID: IntrinsicID);
532 return false;
533}
534
535Type *getMDOperandAsType(const MDNode *N, unsigned I) {
536 Type *ElementTy = cast<ValueAsMetadata>(Val: N->getOperand(I))->getType();
537 return toTypedPointer(Ty: ElementTy);
538}
539
540ConstantInt *getMDOperandAsConstInt(const MDNode *N, unsigned I) {
541 if (N->getNumOperands() <= I)
542 return nullptr;
543 if (auto *CMeta = dyn_cast<ConstantAsMetadata>(Val: N->getOperand(I)))
544 return dyn_cast<ConstantInt>(Val: CMeta->getValue());
545 return nullptr;
546}
547
548static bool isEnqueueKernelBI(StringRef MangledName) {
549 return MangledName == "__enqueue_kernel_basic" ||
550 MangledName == "__enqueue_kernel_basic_events" ||
551 MangledName == "__enqueue_kernel_varargs" ||
552 MangledName == "__enqueue_kernel_events_varargs";
553}
554
555static bool isKernelQueryBI(StringRef MangledName) {
556 return MangledName == "__get_kernel_work_group_size_impl" ||
557 MangledName == "__get_kernel_sub_group_count_for_ndrange_impl" ||
558 MangledName == "__get_kernel_max_sub_group_size_for_ndrange_impl" ||
559 MangledName == "__get_kernel_preferred_work_group_size_multiple_impl";
560}
561
562static bool isNonMangledOCLBuiltin(StringRef Name) {
563 if (!Name.starts_with(Prefix: "__"))
564 return false;
565
566 return isEnqueueKernelBI(MangledName: Name) || isKernelQueryBI(MangledName: Name) ||
567 SPIRV::isPipeOrAddressSpaceCastBuiltin(Name) ||
568 Name == "__translate_sampler_initializer";
569}
570
571std::string getOclOrSpirvBuiltinDemangledName(StringRef Name) {
572 bool IsNonMangledOCL = isNonMangledOCLBuiltin(Name);
573 bool IsNonMangledSPIRV = Name.starts_with(Prefix: "__spirv_");
574 bool IsNonMangledHLSL = Name.starts_with(Prefix: "__hlsl_");
575 bool IsMangled = Name.starts_with(Prefix: "_Z");
576
577 // Otherwise use simple demangling to return the function name.
578 if (IsNonMangledOCL || IsNonMangledSPIRV || IsNonMangledHLSL || !IsMangled)
579 return Name.str();
580
581 // Try to use the itanium demangler.
582 if (char *DemangledName = itaniumDemangle(mangled_name: Name.data())) {
583 std::string Result = DemangledName;
584 free(ptr: DemangledName);
585 return Result;
586 }
587
588 // Autocheck C++, maybe need to do explicit check of the source language.
589 // OpenCL C++ built-ins are declared in cl namespace.
590 // TODO: consider using 'St' abbriviation for cl namespace mangling.
591 // Similar to ::std:: in C++.
592 size_t Start, Len = 0;
593 size_t DemangledNameLenStart = 2;
594 if (Name.starts_with(Prefix: "_ZN")) {
595 // Skip CV and ref qualifiers.
596 size_t NameSpaceStart = Name.find_first_not_of(Chars: "rVKRO", From: 3);
597 // All built-ins are in the ::cl:: namespace.
598 if (Name.substr(Start: NameSpaceStart, N: 11) != "2cl7__spirv")
599 return std::string();
600 DemangledNameLenStart = NameSpaceStart + 11;
601 }
602 Start = Name.find_first_not_of(Chars: "0123456789", From: DemangledNameLenStart);
603 bool Error = Name.substr(Start: DemangledNameLenStart, N: Start - DemangledNameLenStart)
604 .getAsInteger(Radix: 10, Result&: Len);
605 if (Error)
606 return std::string();
607 return Name.substr(Start, N: Len).str();
608}
609
610bool hasBuiltinTypePrefix(StringRef Name) {
611 if (Name.starts_with(Prefix: "opencl.") || Name.starts_with(Prefix: "ocl_") ||
612 Name.starts_with(Prefix: "spirv."))
613 return true;
614 return false;
615}
616
617bool isSpecialOpaqueType(const Type *Ty) {
618 if (const TargetExtType *ExtTy = dyn_cast<TargetExtType>(Val: Ty))
619 return isTypedPointerWrapper(ExtTy)
620 ? false
621 : hasBuiltinTypePrefix(Name: ExtTy->getName());
622
623 return false;
624}
625
626bool isEntryPoint(const Function &F) {
627 // OpenCL handling: any function with the SPIR_KERNEL
628 // calling convention will be a potential entry point.
629 if (F.getCallingConv() == CallingConv::SPIR_KERNEL)
630 return true;
631
632 // HLSL handling: special attribute are emitted from the
633 // front-end.
634 if (F.getFnAttribute(Kind: "hlsl.shader").isValid())
635 return true;
636
637 return false;
638}
639
640Type *parseBasicTypeName(StringRef &TypeName, LLVMContext &Ctx) {
641 TypeName.consume_front(Prefix: "atomic_");
642 if (TypeName.consume_front(Prefix: "void"))
643 return Type::getVoidTy(C&: Ctx);
644 else if (TypeName.consume_front(Prefix: "bool") || TypeName.consume_front(Prefix: "_Bool"))
645 return Type::getIntNTy(C&: Ctx, N: 1);
646 else if (TypeName.consume_front(Prefix: "char") ||
647 TypeName.consume_front(Prefix: "signed char") ||
648 TypeName.consume_front(Prefix: "unsigned char") ||
649 TypeName.consume_front(Prefix: "uchar"))
650 return Type::getInt8Ty(C&: Ctx);
651 else if (TypeName.consume_front(Prefix: "short") ||
652 TypeName.consume_front(Prefix: "signed short") ||
653 TypeName.consume_front(Prefix: "unsigned short") ||
654 TypeName.consume_front(Prefix: "ushort"))
655 return Type::getInt16Ty(C&: Ctx);
656 else if (TypeName.consume_front(Prefix: "int") ||
657 TypeName.consume_front(Prefix: "signed int") ||
658 TypeName.consume_front(Prefix: "unsigned int") ||
659 TypeName.consume_front(Prefix: "uint"))
660 return Type::getInt32Ty(C&: Ctx);
661 else if (TypeName.consume_front(Prefix: "long") ||
662 TypeName.consume_front(Prefix: "signed long") ||
663 TypeName.consume_front(Prefix: "unsigned long") ||
664 TypeName.consume_front(Prefix: "ulong"))
665 return Type::getInt64Ty(C&: Ctx);
666 else if (TypeName.consume_front(Prefix: "half") ||
667 TypeName.consume_front(Prefix: "_Float16") ||
668 TypeName.consume_front(Prefix: "__fp16"))
669 return Type::getHalfTy(C&: Ctx);
670 else if (TypeName.consume_front(Prefix: "float"))
671 return Type::getFloatTy(C&: Ctx);
672 else if (TypeName.consume_front(Prefix: "double"))
673 return Type::getDoubleTy(C&: Ctx);
674
675 // Unable to recognize SPIRV type name
676 return nullptr;
677}
678
679SmallPtrSet<BasicBlock *, 0>
680PartialOrderingVisitor::getReachableFrom(BasicBlock *Start) {
681 std::queue<BasicBlock *> ToVisit;
682 ToVisit.push(x: Start);
683
684 SmallPtrSet<BasicBlock *, 0> Output;
685 while (ToVisit.size() != 0) {
686 BasicBlock *BB = ToVisit.front();
687 ToVisit.pop();
688
689 if (Output.count(Ptr: BB) != 0)
690 continue;
691 Output.insert(Ptr: BB);
692
693 for (BasicBlock *Successor : successors(BB)) {
694 if (DT.dominates(A: Successor, B: BB))
695 continue;
696 ToVisit.push(x: Successor);
697 }
698 }
699
700 return Output;
701}
702
703bool PartialOrderingVisitor::CanBeVisited(BasicBlock *BB) const {
704 for (BasicBlock *P : predecessors(BB)) {
705 // Ignore back-edges.
706 if (DT.dominates(A: BB, B: P))
707 continue;
708
709 // One of the predecessor hasn't been visited. Not ready yet.
710 if (BlockToOrder.count(Val: P) == 0)
711 return false;
712
713 // If the block is a loop exit, the loop must be finished before
714 // we can continue.
715 Loop *L = LI.getLoopFor(BB: P);
716 if (L == nullptr || L->contains(BB))
717 continue;
718
719 // SPIR-V requires a single back-edge. And the backend first
720 // step transforms loops into the simplified format. If we have
721 // more than 1 back-edge, something is wrong.
722 assert(L->getNumBackEdges() <= 1);
723
724 // If the loop has no latch, loop's rank won't matter, so we can
725 // proceed.
726 BasicBlock *Latch = L->getLoopLatch();
727 assert(Latch);
728 if (Latch == nullptr)
729 continue;
730
731 // The latch is not ready yet, let's wait.
732 if (BlockToOrder.count(Val: Latch) == 0)
733 return false;
734 }
735
736 return true;
737}
738
739size_t PartialOrderingVisitor::GetNodeRank(BasicBlock *BB) const {
740 auto It = BlockToOrder.find(Val: BB);
741 if (It != BlockToOrder.end())
742 return It->second.Rank;
743
744 size_t result = 0;
745 for (BasicBlock *P : predecessors(BB)) {
746 // Ignore back-edges.
747 if (DT.dominates(A: BB, B: P))
748 continue;
749
750 auto Iterator = BlockToOrder.end();
751 Loop *L = LI.getLoopFor(BB: P);
752 BasicBlock *Latch = L ? L->getLoopLatch() : nullptr;
753
754 // If the predecessor is either outside a loop, or part of
755 // the same loop, simply take its rank + 1.
756 if (L == nullptr || L->contains(BB) || Latch == nullptr) {
757 Iterator = BlockToOrder.find(Val: P);
758 } else {
759 // Otherwise, take the loop's rank (highest rank in the loop) as base.
760 // Since loops have a single latch, highest rank is easy to find.
761 // If the loop has no latch, then it doesn't matter.
762 Iterator = BlockToOrder.find(Val: Latch);
763 }
764
765 assert(Iterator != BlockToOrder.end());
766 result = std::max(a: result, b: Iterator->second.Rank + 1);
767 }
768
769 return result;
770}
771
772size_t PartialOrderingVisitor::visit(BasicBlock *BB, size_t Unused) {
773 ToVisit.push(x: BB);
774 Queued.insert(Ptr: BB);
775
776 size_t QueueIndex = 0;
777 while (ToVisit.size() != 0) {
778 BasicBlock *BB = ToVisit.front();
779 ToVisit.pop();
780
781 if (!CanBeVisited(BB)) {
782 ToVisit.push(x: BB);
783 if (QueueIndex >= ToVisit.size())
784 llvm::report_fatal_error(
785 reason: "No valid candidate in the queue. Is the graph reducible?");
786 QueueIndex++;
787 continue;
788 }
789
790 QueueIndex = 0;
791 size_t Rank = GetNodeRank(BB);
792 OrderInfo Info = {.Rank: Rank, .TraversalIndex: BlockToOrder.size()};
793 BlockToOrder.try_emplace(Key: BB, Args&: Info);
794
795 for (BasicBlock *S : successors(BB)) {
796 if (Queued.count(Ptr: S) != 0)
797 continue;
798 ToVisit.push(x: S);
799 Queued.insert(Ptr: S);
800 }
801 }
802
803 return 0;
804}
805
806PartialOrderingVisitor::PartialOrderingVisitor(Function &F) {
807 DT.recalculate(Func&: F);
808 LI = LoopInfo(DT);
809
810 visit(BB: &*F.begin(), Unused: 0);
811
812 Order.reserve(n: F.size());
813 for (auto &[BB, Info] : BlockToOrder)
814 Order.emplace_back(args&: BB);
815
816 llvm::sort(C&: Order, Comp: [&](const auto &LHS, const auto &RHS) {
817 return compare(LHS, RHS);
818 });
819}
820
821bool PartialOrderingVisitor::compare(const BasicBlock *LHS,
822 const BasicBlock *RHS) const {
823 const OrderInfo &InfoLHS = BlockToOrder.at(Val: const_cast<BasicBlock *>(LHS));
824 const OrderInfo &InfoRHS = BlockToOrder.at(Val: const_cast<BasicBlock *>(RHS));
825 if (InfoLHS.Rank != InfoRHS.Rank)
826 return InfoLHS.Rank < InfoRHS.Rank;
827 return InfoLHS.TraversalIndex < InfoRHS.TraversalIndex;
828}
829
830void PartialOrderingVisitor::partialOrderVisit(
831 BasicBlock &Start, std::function<bool(BasicBlock *)> Op) {
832 SmallPtrSet<BasicBlock *, 0> Reachable = getReachableFrom(Start: &Start);
833 assert(BlockToOrder.count(&Start) != 0);
834
835 // Skipping blocks with a rank inferior to |Start|'s rank.
836 auto It = Order.begin();
837 while (It != Order.end() && *It != &Start)
838 ++It;
839
840 // This is unexpected. Worst case |Start| is the last block,
841 // so It should point to the last block, not past-end.
842 assert(It != Order.end());
843
844 // By default, there is no rank limit. Setting it to the maximum value.
845 std::optional<size_t> EndRank = std::nullopt;
846 for (; It != Order.end(); ++It) {
847 if (EndRank.has_value() && BlockToOrder[*It].Rank > *EndRank)
848 break;
849
850 if (Reachable.count(Ptr: *It) == 0) {
851 continue;
852 }
853
854 if (!Op(*It)) {
855 EndRank = BlockToOrder[*It].Rank;
856 }
857 }
858}
859
860bool sortBlocks(Function &F) {
861 if (F.size() == 0)
862 return false;
863
864 bool Modified = false;
865 std::vector<BasicBlock *> Order;
866 Order.reserve(n: F.size());
867
868 ReversePostOrderTraversal<Function *> RPOT(&F);
869 llvm::append_range(C&: Order, R&: RPOT);
870
871 assert(&*F.begin() == Order[0]);
872 BasicBlock *LastBlock = &*F.begin();
873 for (BasicBlock *BB : Order) {
874 if (BB != LastBlock && &*LastBlock->getNextNode() != BB) {
875 Modified = true;
876 BB->moveAfter(MovePos: LastBlock);
877 }
878 LastBlock = BB;
879 }
880
881 return Modified;
882}
883
884AllocaInst *createVariable(Function &F, Type *Type) {
885 const DataLayout &DL = F.getDataLayout();
886 return new AllocaInst(Type, DL.getAllocaAddrSpace(), nullptr, "reg",
887 F.begin()->getFirstInsertionPt());
888}
889
890Value *
891createExitVariable(BasicBlock *BB,
892 const DenseMap<BasicBlock *, ConstantInt *> &TargetToValue) {
893 auto *T = BB->getTerminator();
894 if (isa<ReturnInst>(Val: T))
895 return nullptr;
896 if (auto *BI = dyn_cast<UncondBrInst>(Val: T))
897 return TargetToValue.lookup(Val: BI->getSuccessor());
898
899 IRBuilder<> Builder(BB);
900 Builder.SetInsertPoint(T);
901
902 if (auto *BI = dyn_cast<CondBrInst>(Val: T)) {
903 Value *LHS = TargetToValue.lookup(Val: BI->getSuccessor(i: 0));
904 Value *RHS = TargetToValue.lookup(Val: BI->getSuccessor(i: 1));
905
906 if (LHS == nullptr || RHS == nullptr)
907 return LHS == nullptr ? RHS : LHS;
908 return Builder.CreateSelect(C: BI->getCondition(), True: LHS, False: RHS);
909 }
910
911 if (auto *SI = dyn_cast<SwitchInst>(Val: T)) {
912 Value *Condition = SI->getCondition();
913 // The default destination acts as the fallback value of the select chain.
914 Value *Result = TargetToValue.lookup(Val: SI->getDefaultDest());
915 for (const auto &Case : SI->cases()) {
916 Value *CaseValue = TargetToValue.lookup(Val: Case.getCaseSuccessor());
917 // Successors that are internal to the region have no exit value.
918 if (CaseValue == nullptr)
919 continue;
920 // The first known exit value becomes the base of the select chain.
921 if (Result == nullptr) {
922 Result = CaseValue;
923 continue;
924 }
925 Value *Cmp = Builder.CreateICmpEQ(LHS: Condition, RHS: Case.getCaseValue());
926 Result = Builder.CreateSelect(C: Cmp, True: CaseValue, False: Result);
927 }
928 return Result;
929 }
930
931 llvm_unreachable("Unhandled terminator type.");
932}
933
934MachineInstr *getVRegDef(MachineRegisterInfo &MRI, Register Reg) {
935 MachineInstr *MaybeDef = MRI.getVRegDef(Reg);
936 if (MaybeDef && MaybeDef->getOpcode() == SPIRV::ASSIGN_TYPE)
937 MaybeDef = MRI.getVRegDef(Reg: MaybeDef->getOperand(i: 1).getReg());
938 return MaybeDef;
939}
940
941static bool getVacantFunctionName(Module &M, std::string &Name) {
942 // It's a bit of paranoia, but still we don't want to have even a chance that
943 // the loop will work for too long.
944 constexpr unsigned MaxIters = 1024;
945 for (unsigned I = 0; I < MaxIters; ++I) {
946 std::string OrdName = Name + Twine(I).str();
947 if (!M.getFunction(Name: OrdName)) {
948 Name = std::move(OrdName);
949 return true;
950 }
951 }
952 return false;
953}
954
955// Assign SPIR-V type to the register. If the register has no valid assigned
956// class, set register LLT type and class according to the SPIR-V type.
957void setRegClassType(Register Reg, SPIRVTypeInst SpvType,
958 SPIRVGlobalRegistry *GR, MachineRegisterInfo *MRI,
959 const MachineFunction &MF, bool Force) {
960 GR->assignSPIRVTypeToVReg(Type: SpvType, VReg: Reg, MF);
961 if (!MRI->getRegClassOrNull(Reg) || Force) {
962 MRI->setRegClass(Reg, RC: GR->getRegClass(SpvType));
963 LLT RegType = GR->getRegType(SpvType);
964 if (Force || !MRI->getType(Reg).isValid())
965 MRI->setType(VReg: Reg, Ty: RegType);
966 }
967}
968
969// Create a SPIR-V type, assign SPIR-V type to the register. If the register has
970// no valid assigned class, set register LLT type and class according to the
971// SPIR-V type.
972void setRegClassType(Register Reg, const Type *Ty, SPIRVGlobalRegistry *GR,
973 MachineIRBuilder &MIRBuilder,
974 SPIRV::AccessQualifier::AccessQualifier AccessQual,
975 bool EmitIR, bool Force) {
976 setRegClassType(Reg,
977 SpvType: GR->getOrCreateSPIRVType(Type: Ty, MIRBuilder, AQ: AccessQual, EmitIR),
978 GR, MRI: MIRBuilder.getMRI(), MF: MIRBuilder.getMF(), Force);
979}
980
981// Create a virtual register and assign SPIR-V type to the register. Set
982// register LLT type and class according to the SPIR-V type.
983Register createVirtualRegister(SPIRVTypeInst SpvType, SPIRVGlobalRegistry *GR,
984 MachineRegisterInfo *MRI,
985 const MachineFunction &MF) {
986 Register Reg = MRI->createVirtualRegister(RegClass: GR->getRegClass(SpvType));
987 MRI->setType(VReg: Reg, Ty: GR->getRegType(SpvType));
988 GR->assignSPIRVTypeToVReg(Type: SpvType, VReg: Reg, MF);
989 return Reg;
990}
991
992// Create a virtual register and assign SPIR-V type to the register. Set
993// register LLT type and class according to the SPIR-V type.
994Register createVirtualRegister(SPIRVTypeInst SpvType, SPIRVGlobalRegistry *GR,
995 MachineIRBuilder &MIRBuilder) {
996 return createVirtualRegister(SpvType, GR, MRI: MIRBuilder.getMRI(),
997 MF: MIRBuilder.getMF());
998}
999
1000// Create a SPIR-V type, virtual register and assign SPIR-V type to the
1001// register. Set register LLT type and class according to the SPIR-V type.
1002Register createVirtualRegister(
1003 const Type *Ty, SPIRVGlobalRegistry *GR, MachineIRBuilder &MIRBuilder,
1004 SPIRV::AccessQualifier::AccessQualifier AccessQual, bool EmitIR) {
1005 return createVirtualRegister(
1006 SpvType: GR->getOrCreateSPIRVType(Type: Ty, MIRBuilder, AQ: AccessQual, EmitIR), GR,
1007 MIRBuilder);
1008}
1009
1010bool isVectorType(SPIRVTypeInst SPVTy) {
1011 return SPVTy->getOpcode() == SPIRV::OpTypeVector ||
1012 SPVTy->getOpcode() == SPIRV::OpTypeVectorIdEXT;
1013}
1014
1015CallInst *buildIntrWithMD(Intrinsic::ID IntrID, ArrayRef<Type *> Types,
1016 Value *Arg, Value *Arg2, ArrayRef<Constant *> Imms,
1017 IRBuilder<> &B) {
1018 SmallVector<Value *, 4> Args;
1019 Args.push_back(Elt: Arg2);
1020 Args.push_back(Elt: buildMD(Arg));
1021 llvm::append_range(C&: Args, R&: Imms);
1022 return B.CreateIntrinsicWithoutFolding(ID: IntrID, OverloadTypes: {Types}, Args);
1023}
1024
1025// Return true if there is an opaque pointer type nested in the argument.
1026bool isNestedPointer(const Type *Ty) {
1027 if (Ty->isPtrOrPtrVectorTy())
1028 return true;
1029 if (const FunctionType *RefTy = dyn_cast<FunctionType>(Val: Ty)) {
1030 if (isNestedPointer(Ty: RefTy->getReturnType()))
1031 return true;
1032 for (const Type *ArgTy : RefTy->params())
1033 if (isNestedPointer(Ty: ArgTy))
1034 return true;
1035 return false;
1036 }
1037 if (const ArrayType *RefTy = dyn_cast<ArrayType>(Val: Ty))
1038 return isNestedPointer(Ty: RefTy->getElementType());
1039 return false;
1040}
1041
1042bool isSpvIntrinsic(const Value *Arg) {
1043 if (const auto *II = dyn_cast<IntrinsicInst>(Val: Arg))
1044 if (Function *F = II->getCalledFunction())
1045 if (F->getName().starts_with(Prefix: "llvm.spv."))
1046 return true;
1047 return false;
1048}
1049
1050// Function to create continued instructions for SPV_INTEL_long_composites
1051// extension
1052SmallVector<MachineInstr *, 4>
1053createContinuedInstructions(MachineIRBuilder &MIRBuilder, unsigned Opcode,
1054 unsigned MinWC, unsigned ContinuedOpcode,
1055 ArrayRef<Register> Args, Register ReturnRegister,
1056 Register TypeID) {
1057
1058 SmallVector<MachineInstr *, 4> Instructions;
1059 constexpr unsigned MaxWordCount = UINT16_MAX;
1060 const size_t NumElements = Args.size();
1061 size_t MaxNumElements = MaxWordCount - MinWC;
1062 size_t SPIRVStructNumElements = NumElements;
1063
1064 if (NumElements > MaxNumElements) {
1065 // Do adjustments for continued instructions which always had only one
1066 // minumum word count.
1067 SPIRVStructNumElements = MaxNumElements;
1068 MaxNumElements = MaxWordCount - 1;
1069 }
1070
1071 auto MIB =
1072 MIRBuilder.buildInstr(Opcode).addDef(RegNo: ReturnRegister).addUse(RegNo: TypeID);
1073
1074 for (size_t I = 0; I < SPIRVStructNumElements; ++I)
1075 MIB.addUse(RegNo: Args[I]);
1076
1077 Instructions.push_back(Elt: MIB.getInstr());
1078
1079 for (size_t I = SPIRVStructNumElements; I < NumElements;
1080 I += MaxNumElements) {
1081 auto MIB = MIRBuilder.buildInstr(Opcode: ContinuedOpcode);
1082 for (size_t J = I; J < std::min(a: I + MaxNumElements, b: NumElements); ++J)
1083 MIB.addUse(RegNo: Args[J]);
1084 Instructions.push_back(Elt: MIB.getInstr());
1085 }
1086 return Instructions;
1087}
1088
1089SmallVector<unsigned, 1>
1090getSpirvLoopControlOperandsFromLoopMetadata(MDNode *LoopMD) {
1091 unsigned LC = SPIRV::LoopControl::None;
1092 // Currently used only to store PartialCount value. Later when other
1093 // LoopControls are added - this map should be sorted before making
1094 // them loop_merge operands to satisfy 3.23. Loop Control requirements.
1095 std::vector<std::pair<unsigned, unsigned>> MaskToValueMap;
1096 if (findOptionMDForLoopID(LoopID: LoopMD, Name: "llvm.loop.unroll.disable")) {
1097 LC |= SPIRV::LoopControl::DontUnroll;
1098 } else {
1099 if (findOptionMDForLoopID(LoopID: LoopMD, Name: "llvm.loop.unroll.enable") ||
1100 findOptionMDForLoopID(LoopID: LoopMD, Name: "llvm.loop.unroll.full")) {
1101 LC |= SPIRV::LoopControl::Unroll;
1102 }
1103 if (MDNode *CountMD =
1104 findOptionMDForLoopID(LoopID: LoopMD, Name: "llvm.loop.unroll.count")) {
1105 if (auto *CI =
1106 mdconst::extract_or_null<ConstantInt>(MD: CountMD->getOperand(I: 1))) {
1107 unsigned Count = CI->getZExtValue();
1108 if (Count != 1) {
1109 LC |= SPIRV::LoopControl::PartialCount;
1110 MaskToValueMap.emplace_back(
1111 args: std::make_pair(x: SPIRV::LoopControl::PartialCount, y&: Count));
1112 }
1113 }
1114 }
1115 }
1116 SmallVector<unsigned, 1> Result = {LC};
1117 for (auto &[Mask, Val] : MaskToValueMap)
1118 Result.push_back(Elt: Val);
1119 return Result;
1120}
1121
1122SmallVector<unsigned, 1> getSpirvLoopControlOperandsFromLoopMetadata(Loop *L) {
1123 return getSpirvLoopControlOperandsFromLoopMetadata(LoopMD: L->getLoopID());
1124}
1125
1126const std::set<unsigned> &getTypeFoldingSupportedOpcodes() {
1127 // clang-format off
1128 static const std::set<unsigned> TypeFoldingSupportingOpcs = {
1129 TargetOpcode::G_ADD,
1130 TargetOpcode::G_FADD,
1131 TargetOpcode::G_STRICT_FADD,
1132 TargetOpcode::G_SUB,
1133 TargetOpcode::G_FSUB,
1134 TargetOpcode::G_STRICT_FSUB,
1135 TargetOpcode::G_MUL,
1136 TargetOpcode::G_FMUL,
1137 TargetOpcode::G_STRICT_FMUL,
1138 TargetOpcode::G_SDIV,
1139 TargetOpcode::G_UDIV,
1140 TargetOpcode::G_FDIV,
1141 TargetOpcode::G_STRICT_FDIV,
1142 TargetOpcode::G_SREM,
1143 TargetOpcode::G_UREM,
1144 TargetOpcode::G_FREM,
1145 TargetOpcode::G_STRICT_FREM,
1146 TargetOpcode::G_FNEG,
1147 TargetOpcode::G_CONSTANT,
1148 TargetOpcode::G_FCONSTANT,
1149 TargetOpcode::G_AND,
1150 TargetOpcode::G_OR,
1151 TargetOpcode::G_XOR,
1152 TargetOpcode::G_SHL,
1153 TargetOpcode::G_ASHR,
1154 TargetOpcode::G_LSHR,
1155 TargetOpcode::G_SELECT,
1156 TargetOpcode::G_EXTRACT_VECTOR_ELT,
1157 };
1158 // clang-format on
1159 return TypeFoldingSupportingOpcs;
1160}
1161
1162bool isTypeFoldingSupported(unsigned Opcode) {
1163 return getTypeFoldingSupportedOpcodes().count(x: Opcode) > 0;
1164}
1165
1166// Traversing [g]MIR accounting for pseudo-instructions.
1167MachineInstr *passCopy(MachineInstr *Def, const MachineRegisterInfo *MRI) {
1168 return (Def->getOpcode() == SPIRV::ASSIGN_TYPE ||
1169 Def->getOpcode() == TargetOpcode::COPY)
1170 ? MRI->getVRegDef(Reg: Def->getOperand(i: 1).getReg())
1171 : Def;
1172}
1173
1174MachineInstr *getDef(const MachineOperand &MO, const MachineRegisterInfo *MRI) {
1175 if (MachineInstr *Def = MRI->getVRegDef(Reg: MO.getReg()))
1176 return passCopy(Def, MRI);
1177 return nullptr;
1178}
1179
1180MachineInstr *getImm(const MachineOperand &MO, const MachineRegisterInfo *MRI) {
1181 if (MachineInstr *Def = getDef(MO, MRI)) {
1182 if (Def->getOpcode() == TargetOpcode::G_CONSTANT ||
1183 Def->getOpcode() == SPIRV::OpConstantI)
1184 return Def;
1185 }
1186 return nullptr;
1187}
1188
1189int64_t foldImm(const MachineOperand &MO, const MachineRegisterInfo *MRI) {
1190 if (MachineInstr *Def = getImm(MO, MRI)) {
1191 if (Def->getOpcode() == SPIRV::OpConstantI)
1192 return Def->getOperand(i: 2).getImm();
1193 if (Def->getOpcode() == TargetOpcode::G_CONSTANT)
1194 return Def->getOperand(i: 1).getCImm()->getZExtValue();
1195 }
1196 llvm_unreachable("Unexpected integer constant pattern");
1197}
1198
1199unsigned getArrayComponentCount(const MachineRegisterInfo *MRI,
1200 const MachineInstr *ResType) {
1201 return foldImm(MO: ResType->getOperand(i: 2), MRI);
1202}
1203
1204bool matchPeeledArrayPattern(const StructType *Ty, Type *&OriginalElementType,
1205 uint64_t &TotalSize) {
1206 // An array of N padded structs is represented as {[N-1 x <{T, pad}>], T}.
1207 if (Ty->getStructNumElements() != 2)
1208 return false;
1209
1210 Type *FirstElement = Ty->getStructElementType(N: 0);
1211 Type *SecondElement = Ty->getStructElementType(N: 1);
1212
1213 if (!FirstElement->isArrayTy())
1214 return false;
1215
1216 Type *ArrayElementType = FirstElement->getArrayElementType();
1217 if (!ArrayElementType->isStructTy() ||
1218 ArrayElementType->getStructNumElements() != 2)
1219 return false;
1220
1221 Type *T_in_struct = ArrayElementType->getStructElementType(N: 0);
1222 if (T_in_struct != SecondElement)
1223 return false;
1224
1225 auto *Padding_in_struct =
1226 dyn_cast<TargetExtType>(Val: ArrayElementType->getStructElementType(N: 1));
1227 if (!Padding_in_struct || Padding_in_struct->getName() != "spirv.Padding")
1228 return false;
1229
1230 const uint64_t ArraySize = FirstElement->getArrayNumElements();
1231 TotalSize = ArraySize + 1;
1232 OriginalElementType = ArrayElementType;
1233 return true;
1234}
1235
1236Type *reconstitutePeeledArrayType(Type *Ty) {
1237 if (!Ty->isStructTy())
1238 return Ty;
1239
1240 auto *STy = cast<StructType>(Val: Ty);
1241 Type *OriginalElementType = nullptr;
1242 uint64_t TotalSize = 0;
1243 if (matchPeeledArrayPattern(Ty: STy, OriginalElementType, TotalSize)) {
1244 Type *ResultTy = ArrayType::get(
1245 ElementType: reconstitutePeeledArrayType(Ty: OriginalElementType), NumElements: TotalSize);
1246 return ResultTy;
1247 }
1248
1249 SmallVector<Type *, 4> NewElementTypes;
1250 bool Changed = false;
1251 for (Type *ElementTy : STy->elements()) {
1252 Type *NewElementTy = reconstitutePeeledArrayType(Ty: ElementTy);
1253 if (NewElementTy != ElementTy)
1254 Changed = true;
1255 NewElementTypes.push_back(Elt: NewElementTy);
1256 }
1257
1258 if (!Changed)
1259 return Ty;
1260
1261 Type *ResultTy;
1262 if (STy->isLiteral()) {
1263 ResultTy =
1264 StructType::get(Context&: STy->getContext(), Elements: NewElementTypes, isPacked: STy->isPacked());
1265 } else {
1266 ResultTy = StructType::create(Context&: STy->getContext(), Elements: NewElementTypes,
1267 Name: STy->getName(), isPacked: STy->isPacked());
1268 }
1269 return ResultTy;
1270}
1271
1272std::optional<SPIRV::LinkageType::LinkageType>
1273getSpirvLinkageTypeFor(const SPIRVSubtarget &ST, const GlobalValue &GV) {
1274 if (GV.hasLocalLinkage())
1275 return std::nullopt;
1276
1277 if (GV.isDeclarationForLinker()) {
1278 if (const auto *GVar = dyn_cast<GlobalVariable>(Val: &GV)) {
1279 auto SC = addressSpaceToStorageClass(AddrSpace: GVar->getAddressSpace(), STI: ST);
1280 // Interface variables must not get Import linkage.
1281 if (SC == SPIRV::StorageClass::Input ||
1282 SC == SPIRV::StorageClass::Output ||
1283 SC == SPIRV::StorageClass::PushConstant)
1284 return std::nullopt;
1285 // Shaders have no linker, so module-internal storage
1286 // (e.g. HLSL groupshared) can't be imported
1287 if (ST.isShader() && (SC == SPIRV::StorageClass::Workgroup ||
1288 SC == SPIRV::StorageClass::Private))
1289 return std::nullopt;
1290 }
1291 return SPIRV::LinkageType::Import;
1292 }
1293
1294 if (GV.hasHiddenVisibility())
1295 return std::nullopt;
1296
1297 if (GV.hasLinkOnceODRLinkage() &&
1298 ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_linkonce_odr))
1299 return SPIRV::LinkageType::LinkOnceODR;
1300
1301 if (GV.hasWeakLinkage() &&
1302 ST.canUseExtension(E: SPIRV::Extension::SPV_AMD_weak_linkage))
1303 return SPIRV::LinkageType::WeakAMD;
1304
1305 return SPIRV::LinkageType::Export;
1306}
1307
1308Function *getOrCreateBackendServiceFunction(Module &M) {
1309 std::string ServiceFunName = SPIRV_BACKEND_SERVICE_FUN_NAME;
1310 if (!getVacantFunctionName(M, Name&: ServiceFunName))
1311 report_fatal_error(
1312 reason: "cannot allocate a name for the internal service function");
1313 if (Function *SF = M.getFunction(Name: ServiceFunName)) {
1314 if (SF->getInstructionCount() > 0)
1315 report_fatal_error(
1316 reason: "Unexpected combination of global variables and function pointers");
1317 return SF;
1318 }
1319 Function *SF = Function::Create(
1320 Ty: FunctionType::get(Result: Type::getVoidTy(C&: M.getContext()), Params: {}, isVarArg: false),
1321 Linkage: GlobalValue::PrivateLinkage, N: ServiceFunName, M);
1322 SF->addFnAttr(SPIRV_BACKEND_SERVICE_FUN_NAME, Val: "");
1323 return SF;
1324}
1325
1326} // namespace llvm
1327