1//===- SPIRVModuleAnalysis.cpp - analysis of global instrs & regs - C++ -*-===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// The analysis collects instructions that should be output at the module level
10// and performs the global register numbering.
11//
12// The results of this analysis are used in AsmPrinter to rename registers
13// globally and to output required instructions at the module level.
14//
15//===----------------------------------------------------------------------===//
16
17// TODO: Per LLVM best practices, the report_fatal_error (deprecated) /
18// ReportFatalUsageError calls in this file should be replaced with the
19// Diagnostic infrastructure (e.g. the reportUnsupported function below).
20
21#include "SPIRVModuleAnalysis.h"
22#include "MCTargetDesc/SPIRVBaseInfo.h"
23#include "MCTargetDesc/SPIRVMCTargetDesc.h"
24#include "SPIRV.h"
25#include "SPIRVSubtarget.h"
26#include "SPIRVTargetMachine.h"
27#include "SPIRVUtils.h"
28#include "llvm/ADT/STLExtras.h"
29#include "llvm/CodeGen/MachineFunctionAnalysis.h"
30#include "llvm/CodeGen/MachineModuleInfo.h"
31#include "llvm/CodeGen/TargetPassConfig.h"
32
33using namespace llvm;
34
35#define DEBUG_TYPE "spirv-module-analysis"
36
37static cl::opt<bool>
38 SPVDumpDeps("spv-dump-deps",
39 cl::desc("Dump MIR with SPIR-V dependencies info"),
40 cl::init(Val: false));
41
42static cl::list<SPIRV::Capability::Capability>
43 AvoidCapabilities("avoid-spirv-capabilities",
44 cl::desc("SPIR-V capabilities to avoid if there are "
45 "other options enabling a feature"),
46 cl::Hidden,
47 cl::values(clEnumValN(SPIRV::Capability::Shader, "Shader",
48 "SPIR-V Shader capability")));
49// Use sets instead of cl::list to check "if contains" condition
50struct AvoidCapabilitiesSet {
51 SmallSet<SPIRV::Capability::Capability, 4> S;
52 AvoidCapabilitiesSet() { S.insert_range(R&: AvoidCapabilities); }
53};
54
55char llvm::SPIRVModuleAnalysisWrapperPass::ID = 0;
56
57INITIALIZE_PASS(SPIRVModuleAnalysisWrapperPass, DEBUG_TYPE,
58 "SPIRV module analysis", true, true)
59
60static void reportUnsupported(const MachineInstr &MI, const char *Msg) {
61 const Function &Func = MI.getMF()->getFunction();
62 Func.getContext().diagnose(
63 DI: DiagnosticInfoUnsupported(Func, Msg, MI.getDebugLoc()));
64}
65
66// Retrieve an unsigned from an MDNode with a list of them as operands.
67static unsigned getMetadataUInt(MDNode *MdNode, unsigned OpIndex,
68 unsigned DefaultVal = 0) {
69 if (MdNode && OpIndex < MdNode->getNumOperands()) {
70 const auto &Op = MdNode->getOperand(I: OpIndex);
71 return mdconst::extract<ConstantInt>(MD: Op)->getZExtValue();
72 }
73 return DefaultVal;
74}
75
76static SPIRV::Requirements
77getSymbolicOperandRequirements(SPIRV::OperandCategory::OperandCategory Category,
78 unsigned i, const SPIRVSubtarget &ST,
79 SPIRV::RequirementHandler &Reqs) {
80 // A set of capabilities to avoid if there is another option.
81 AvoidCapabilitiesSet AvoidCaps;
82 if (!ST.isShader())
83 AvoidCaps.S.insert(V: SPIRV::Capability::Shader);
84 else
85 AvoidCaps.S.insert(V: SPIRV::Capability::Kernel);
86
87 VersionTuple ReqMinVer = getSymbolicOperandMinVersion(Category, Value: i);
88 VersionTuple ReqMaxVer = getSymbolicOperandMaxVersion(Category, Value: i);
89 VersionTuple SPIRVVersion = ST.getSPIRVVersion();
90 bool MinVerOK = SPIRVVersion.empty() || SPIRVVersion >= ReqMinVer;
91 bool MaxVerOK =
92 ReqMaxVer.empty() || SPIRVVersion.empty() || SPIRVVersion <= ReqMaxVer;
93 CapabilityList ReqCaps = getSymbolicOperandCapabilities(Category, Value: i);
94 ExtensionList ReqExts = getSymbolicOperandExtensions(Category, Value: i);
95 if (ReqCaps.empty()) {
96 if (ReqExts.empty()) {
97 if (MinVerOK && MaxVerOK)
98 return {true, {}, {}, ReqMinVer, ReqMaxVer};
99 return {false, {}, {}, VersionTuple(), VersionTuple()};
100 }
101 } else if (MinVerOK && MaxVerOK) {
102 if (ReqCaps.size() == 1) {
103 auto Cap = ReqCaps[0];
104 if (Reqs.isCapabilityAvailable(Cap)) {
105 ReqExts.append(RHS: getSymbolicOperandExtensions(
106 Category: SPIRV::OperandCategory::CapabilityOperand, Value: Cap));
107 return {true, {Cap}, std::move(ReqExts), ReqMinVer, ReqMaxVer};
108 }
109 } else {
110 // By SPIR-V specification: "If an instruction, enumerant, or other
111 // feature specifies multiple enabling capabilities, only one such
112 // capability needs to be declared to use the feature." However, one
113 // capability may be preferred over another. We use command line
114 // argument(s) and AvoidCapabilities to avoid selection of certain
115 // capabilities if there are other options.
116 CapabilityList UseCaps;
117 for (auto Cap : ReqCaps)
118 if (Reqs.isCapabilityAvailable(Cap))
119 UseCaps.push_back(Elt: Cap);
120 for (size_t i = 0, Sz = UseCaps.size(); i < Sz; ++i) {
121 auto Cap = UseCaps[i];
122 if (i == Sz - 1 || !AvoidCaps.S.contains(V: Cap)) {
123 ReqExts.append(RHS: getSymbolicOperandExtensions(
124 Category: SPIRV::OperandCategory::CapabilityOperand, Value: Cap));
125 return {true, {Cap}, std::move(ReqExts), ReqMinVer, ReqMaxVer};
126 }
127 }
128 }
129 }
130 // If there are no capabilities, or we can't satisfy the version or
131 // capability requirements, use the list of extensions (if the subtarget
132 // can handle them all).
133 if (llvm::all_of(Range&: ReqExts, P: [&ST](const SPIRV::Extension::Extension &Ext) {
134 return ST.canUseExtension(E: Ext);
135 })) {
136 return {true,
137 {},
138 std::move(ReqExts),
139 VersionTuple(),
140 VersionTuple()}; // TODO: add versions to extensions.
141 }
142 return {false, {}, {}, VersionTuple(), VersionTuple()};
143}
144
145void SPIRVModuleAnalysisImpl::setBaseInfo(const Module &M) {
146 MAI.MaxID = 0;
147 for (int i = 0; i < SPIRV::NUM_MODULE_SECTIONS; i++)
148 MAI.MS[i].clear();
149 MAI.RegisterAliasTable.clear();
150 MAI.InstrsToDelete.clear();
151 MAI.GlobalObjMap.clear();
152 MAI.GlobalVarList.clear();
153 MAI.ExtInstSetMap.clear();
154 MAI.Reqs.clear();
155 MAI.Reqs.initAvailableCapabilities(ST: *ST);
156
157 // TODO: determine memory model and source language from the configuratoin.
158 if (auto MemModel = M.getNamedMetadata(Name: "spirv.MemoryModel")) {
159 auto MemMD = MemModel->getOperand(i: 0);
160 MAI.Addr = static_cast<SPIRV::AddressingModel::AddressingModel>(
161 getMetadataUInt(MdNode: MemMD, OpIndex: 0));
162 MAI.Mem =
163 static_cast<SPIRV::MemoryModel::MemoryModel>(getMetadataUInt(MdNode: MemMD, OpIndex: 1));
164 } else {
165 // TODO: Add support for VulkanMemoryModel.
166 MAI.Mem = ST->isShader() ? SPIRV::MemoryModel::GLSL450
167 : SPIRV::MemoryModel::OpenCL;
168 if (MAI.Mem == SPIRV::MemoryModel::OpenCL) {
169 unsigned PtrSize = ST->getPointerSize();
170 MAI.Addr = PtrSize == 32 ? SPIRV::AddressingModel::Physical32
171 : PtrSize == 64 ? SPIRV::AddressingModel::Physical64
172 : SPIRV::AddressingModel::Logical;
173 } else {
174 // TODO: Add support for PhysicalStorageBufferAddress.
175 MAI.Addr = SPIRV::AddressingModel::Logical;
176 }
177 }
178 // Get the OpenCL version number from metadata.
179 // TODO: support other source languages.
180 if (auto VerNode = M.getNamedMetadata(Name: "opencl.ocl.version")) {
181 MAI.SrcLang = SPIRV::SourceLanguage::OpenCL_C;
182 // Construct version literal in accordance with SPIRV-LLVM-Translator.
183 // TODO: support multiple OCL version metadata.
184 assert(VerNode->getNumOperands() > 0 && "Invalid SPIR");
185 auto VersionMD = VerNode->getOperand(i: 0);
186 unsigned MajorNum = getMetadataUInt(MdNode: VersionMD, OpIndex: 0, DefaultVal: 2);
187 unsigned MinorNum = getMetadataUInt(MdNode: VersionMD, OpIndex: 1);
188 unsigned RevNum = getMetadataUInt(MdNode: VersionMD, OpIndex: 2);
189 // Prevent Major part of OpenCL version to be 0
190 MAI.SrcLangVersion =
191 (std::max(a: 1U, b: MajorNum) * 100 + MinorNum) * 1000 + RevNum;
192 // When opencl.cxx.version is also present, validate compatibility
193 // and use C++ for OpenCL as source language with the C++ version.
194 if (auto *CxxVerNode = M.getNamedMetadata(Name: "opencl.cxx.version")) {
195 assert(CxxVerNode->getNumOperands() > 0 && "Invalid SPIR");
196 auto *CxxMD = CxxVerNode->getOperand(i: 0);
197 unsigned CxxVer =
198 (getMetadataUInt(MdNode: CxxMD, OpIndex: 0) * 100 + getMetadataUInt(MdNode: CxxMD, OpIndex: 1)) * 1000 +
199 getMetadataUInt(MdNode: CxxMD, OpIndex: 2);
200 if ((MAI.SrcLangVersion == 200000 && CxxVer == 100000) ||
201 (MAI.SrcLangVersion == 300000 && CxxVer == 202100000)) {
202 MAI.SrcLang = SPIRV::SourceLanguage::CPP_for_OpenCL;
203 MAI.SrcLangVersion = CxxVer;
204 } else {
205 report_fatal_error(
206 reason: "opencl cxx version is not compatible with opencl c version!");
207 }
208 }
209 } else {
210 // If there is no information about OpenCL version we are forced to generate
211 // OpenCL 1.0 by default for the OpenCL environment to avoid puzzling
212 // run-times with Unknown/0.0 version output. For a reference, LLVM-SPIRV
213 // Translator avoids potential issues with run-times in a similar manner.
214 if (!ST->isShader()) {
215 MAI.SrcLang = SPIRV::SourceLanguage::OpenCL_CPP;
216 MAI.SrcLangVersion = 100000;
217 } else {
218 MAI.SrcLang = SPIRV::SourceLanguage::Unknown;
219 MAI.SrcLangVersion = 0;
220 }
221 }
222
223 if (auto ExtNode = M.getNamedMetadata(Name: "opencl.used.extensions")) {
224 for (unsigned I = 0, E = ExtNode->getNumOperands(); I != E; ++I) {
225 MDNode *MD = ExtNode->getOperand(i: I);
226 if (!MD || MD->getNumOperands() == 0)
227 continue;
228 for (unsigned J = 0, N = MD->getNumOperands(); J != N; ++J)
229 MAI.SrcExt.insert(key: cast<MDString>(Val: MD->getOperand(I: J))->getString());
230 }
231 }
232
233 // Update required capabilities for this memory model, addressing model and
234 // source language.
235 MAI.Reqs.getAndAddRequirements(Category: SPIRV::OperandCategory::MemoryModelOperand,
236 i: MAI.Mem, ST: *ST);
237 MAI.Reqs.getAndAddRequirements(Category: SPIRV::OperandCategory::SourceLanguageOperand,
238 i: MAI.SrcLang, ST: *ST);
239 MAI.Reqs.getAndAddRequirements(Category: SPIRV::OperandCategory::AddressingModelOperand,
240 i: MAI.Addr, ST: *ST);
241
242 if (MAI.Mem == SPIRV::MemoryModel::VulkanKHR)
243 MAI.Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_KHR_vulkan_memory_model);
244
245 if (!ST->isShader()) {
246 // TODO: check if it's required by default.
247 MAI.ExtInstSetMap[static_cast<unsigned>(
248 SPIRV::InstructionSet::OpenCL_std)] = MAI.getNextIDRegister();
249 }
250}
251
252// Appends the signature of the decoration instructions that decorate R to
253// Signature.
254static void appendDecorationsForReg(const MachineRegisterInfo &MRI, Register R,
255 InstrSignature &Signature) {
256 for (MachineInstr &UseMI : MRI.use_instructions(Reg: R)) {
257 // We don't handle OpDecorateId because getting the register alias for the
258 // ID can cause problems, and we do not need it for now.
259 if (UseMI.getOpcode() != SPIRV::OpDecorate &&
260 UseMI.getOpcode() != SPIRV::OpMemberDecorate)
261 continue;
262
263 for (unsigned I = 0; I < UseMI.getNumOperands(); ++I) {
264 const MachineOperand &MO = UseMI.getOperand(i: I);
265 if (MO.isReg())
266 continue;
267 Signature.push_back(Elt: hash_value(MO));
268 }
269 }
270}
271
272// Returns a representation of an instruction as a vector of MachineOperand
273// hash values, see llvm::hash_value(const MachineOperand &MO) for details.
274// This creates a signature of the instruction with the same content
275// that MachineOperand::isIdenticalTo uses for comparison.
276static InstrSignature instrToSignature(const MachineInstr &MI,
277 SPIRV::ModuleAnalysisInfo &MAI,
278 bool UseDefReg) {
279 Register DefReg;
280 InstrSignature Signature{MI.getOpcode()};
281 for (unsigned i = 0; i < MI.getNumOperands(); ++i) {
282 // The only decorations that can be applied more than once to a given <id>
283 // or structure member are FuncParamAttr (38), UserSemantic (5635),
284 // CacheControlLoadINTEL (6442), and CacheControlStoreINTEL (6443). For all
285 // the rest of decorations, we will only add to the signature the Opcode,
286 // the id to which it applies, and the decoration id, disregarding any
287 // decoration flags. This will ensure that any subsequent decoration with
288 // the same id will be deemed as a duplicate. Then, at the call site, we
289 // will be able to handle duplicates in the best way.
290 unsigned Opcode = MI.getOpcode();
291 if ((Opcode == SPIRV::OpDecorate) && i >= 2) {
292 unsigned DecorationID = MI.getOperand(i: 1).getImm();
293 if (DecorationID != SPIRV::Decoration::FuncParamAttr &&
294 DecorationID != SPIRV::Decoration::UserSemantic &&
295 DecorationID != SPIRV::Decoration::CacheControlLoadINTEL &&
296 DecorationID != SPIRV::Decoration::CacheControlStoreINTEL)
297 continue;
298 }
299 const MachineOperand &MO = MI.getOperand(i);
300 size_t h;
301 if (MO.isReg()) {
302 if (!UseDefReg && MO.isDef()) {
303 assert(!DefReg.isValid() && "Multiple def registers.");
304 DefReg = MO.getReg();
305 continue;
306 }
307 Register RegAlias = MAI.getRegisterAlias(MF: MI.getMF(), Reg: MO.getReg());
308 if (!RegAlias.isValid()) {
309 LLVM_DEBUG({
310 dbgs() << "Unexpectedly, no global id found for the operand ";
311 MO.print(dbgs());
312 dbgs() << "\nInstruction: ";
313 MI.print(dbgs());
314 dbgs() << "\n";
315 });
316 report_fatal_error(reason: "All v-regs must have been mapped to global id's");
317 }
318 // mimic llvm::hash_value(const MachineOperand &MO)
319 h = hash_combine(args: MO.getType(), args: (unsigned)RegAlias, args: MO.getSubReg(),
320 args: MO.isDef());
321 } else {
322 h = hash_value(MO);
323 }
324 Signature.push_back(Elt: h);
325 }
326
327 if (DefReg.isValid()) {
328 // Decorations change the semantics of the current instruction. So two
329 // identical instruction with different decorations cannot be merged. That
330 // is why we add the decorations to the signature.
331 appendDecorationsForReg(MRI: MI.getMF()->getRegInfo(), R: DefReg, Signature);
332 }
333 return Signature;
334}
335
336// Operand index of Invoke in device enqueue instructions, 0 if none.
337static unsigned getInvokeOperandIdx(unsigned Opcode) {
338 switch (Opcode) {
339 case SPIRV::OpEnqueueKernel:
340 return 8;
341 case SPIRV::OpGetKernelNDrangeSubGroupCount:
342 case SPIRV::OpGetKernelNDrangeMaxSubGroupSize:
343 return 3;
344 case SPIRV::OpGetKernelWorkGroupSize:
345 case SPIRV::OpGetKernelPreferredWorkGroupSizeMultiple:
346 return 2;
347 default:
348 return 0;
349 }
350}
351
352bool SPIRVModuleAnalysisImpl::isDeclSection(const MachineRegisterInfo &MRI,
353 const MachineInstr &MI) {
354 unsigned Opcode = MI.getOpcode();
355 switch (Opcode) {
356 case SPIRV::OpTypeForwardPointer:
357 // omit now, collect later
358 return false;
359 case SPIRV::OpVariable:
360 case SPIRV::OpUntypedVariableKHR:
361 return static_cast<SPIRV::StorageClass::StorageClass>(
362 MI.getOperand(i: 2).getImm()) != SPIRV::StorageClass::Function;
363 case SPIRV::OpFunction:
364 case SPIRV::OpFunctionParameter:
365 return true;
366 }
367 if (GR->hasConstFunPtr() && Opcode == SPIRV::OpUndef) {
368 // The OpUndef may be a placeholder for a function reference recorded by
369 // selectGlobalValue. Skip emitting it if any user consumes it as a
370 // function-pointer-like operand (OpConstantFunctionPointerINTEL operand 2,
371 // or the Invoke operand of a device enqueue instruction). The rewrite
372 // happens in visitFunPtrUse, which aliases the OpUndef's vreg to the
373 // function's global <id>.
374 Register DefReg = MI.getOperand(i: 0).getReg();
375 if (GR->getFunctionDefinitionByUse(Use: &MI.getOperand(i: 0))) {
376 for (MachineInstr &UseMI : MRI.use_instructions(Reg: DefReg)) {
377 unsigned UseOp = UseMI.getOpcode();
378 if (UseOp == SPIRV::OpConstantFunctionPointerINTEL ||
379 getInvokeOperandIdx(Opcode: UseOp)) {
380 MAI.setSkipEmission(&MI);
381 return false;
382 }
383 }
384 }
385 for (MachineInstr &UseMI : MRI.use_instructions(Reg: DefReg)) {
386 if (UseMI.getOpcode() != SPIRV::OpConstantFunctionPointerINTEL)
387 continue;
388 // it's a dummy definition, FP constant refers to a function,
389 // and this is resolved in another way; let's skip this definition
390 assert(UseMI.getOperand(2).isReg() &&
391 UseMI.getOperand(2).getReg() == DefReg);
392 MAI.setSkipEmission(&MI);
393 return false;
394 }
395 }
396 return TII->isTypeDeclInstr(MI) || TII->isConstantInstr(MI) ||
397 TII->isInlineAsmDefInstr(MI);
398}
399
400// This is a special case of a function pointer referring to a possibly
401// forward function declaration. The operand is a dummy OpUndef that
402// requires a special treatment.
403// FunPtrOp is the MachineOperand previously recorded via
404// SPIRVGlobalRegistry::recordFunctionPointer, identifying which Function
405// this placeholder refers to.
406void SPIRVModuleAnalysisImpl::visitFunPtrUse(
407 Register OpReg, const MachineOperand *FunPtrOp,
408 InstrGRegsMap &SignatureToGReg,
409 std::map<const Value *, unsigned> &GlobalToGReg,
410 const MachineFunction *MF) {
411 const MachineOperand *OpFunDef = GR->getFunctionDefinitionByUse(Use: FunPtrOp);
412 assert(OpFunDef && OpFunDef->isReg());
413 // find the actual function definition and number it globally in advance
414 const MachineInstr *OpDefMI = OpFunDef->getParent();
415 assert(OpDefMI && OpDefMI->getOpcode() == SPIRV::OpFunction);
416 const MachineFunction *FunDefMF = OpDefMI->getParent()->getParent();
417 const MachineRegisterInfo &FunDefMRI = FunDefMF->getRegInfo();
418 do {
419 visitDecl(MRI: FunDefMRI, SignatureToGReg, GlobalToGReg, MF: FunDefMF, MI: *OpDefMI);
420 OpDefMI = OpDefMI->getNextNode();
421 } while (OpDefMI && (OpDefMI->getOpcode() == SPIRV::OpFunction ||
422 OpDefMI->getOpcode() == SPIRV::OpFunctionParameter));
423 // associate the function pointer with the newly assigned global number
424 MCRegister GlobalFunDefReg =
425 MAI.getRegisterAlias(MF: FunDefMF, Reg: OpFunDef->getReg());
426 assert(GlobalFunDefReg.isValid() &&
427 "Function definition must refer to a global register");
428 MAI.setRegisterAlias(MF, Reg: OpReg, AliasReg: GlobalFunDefReg);
429}
430
431// Depth first recursive traversal of dependencies. Repeated visits are guarded
432// by MAI.hasRegisterAlias().
433void SPIRVModuleAnalysisImpl::visitDecl(
434 const MachineRegisterInfo &MRI, InstrGRegsMap &SignatureToGReg,
435 std::map<const Value *, unsigned> &GlobalToGReg, const MachineFunction *MF,
436 const MachineInstr &MI) {
437 unsigned Opcode = MI.getOpcode();
438
439 // Process each operand of the instruction to resolve dependencies
440 for (const MachineOperand &MO : MI.operands()) {
441 if (!MO.isReg() || MO.isDef())
442 continue;
443 Register OpReg = MO.getReg();
444 // Handle function pointers special case
445 if (Opcode == SPIRV::OpConstantFunctionPointerINTEL &&
446 MRI.getRegClass(Reg: OpReg) == &SPIRV::pIDRegClass) {
447 visitFunPtrUse(OpReg, FunPtrOp: &MI.getOperand(i: 2), SignatureToGReg, GlobalToGReg,
448 MF);
449 continue;
450 }
451 // Skip already processed instructions
452 if (MAI.hasRegisterAlias(MF, Reg: MO.getReg()))
453 continue;
454 // Recursively visit dependencies
455 if (const MachineInstr *OpDefMI = MRI.getUniqueVRegDef(Reg: OpReg)) {
456 if (isDeclSection(MRI, MI: *OpDefMI))
457 visitDecl(MRI, SignatureToGReg, GlobalToGReg, MF, MI: *OpDefMI);
458 continue;
459 }
460 // Handle the unexpected case of no unique definition for the SPIR-V
461 // instruction
462 LLVM_DEBUG({
463 dbgs() << "Unexpectedly, no unique definition for the operand ";
464 MO.print(dbgs());
465 dbgs() << "\nInstruction: ";
466 MI.print(dbgs());
467 dbgs() << "\n";
468 });
469 report_fatal_error(
470 reason: "No unique definition is found for the virtual register");
471 }
472
473 MCRegister GReg;
474 bool IsFunDef = false;
475 if (TII->isSpecConstantInstr(MI)) {
476 GReg = MAI.getNextIDRegister();
477 MAI.MS[SPIRV::MB_TypeConstVars].push_back(Elt: &MI);
478 } else if (Opcode == SPIRV::OpFunction ||
479 Opcode == SPIRV::OpFunctionParameter) {
480 GReg = handleFunctionOrParameter(MF, MI, GlobalToGReg, IsFunDef);
481 } else if (Opcode == SPIRV::OpTypeStruct ||
482 Opcode == SPIRV::OpConstantComposite) {
483 GReg = handleTypeDeclOrConstant(MI, SignatureToGReg);
484 const MachineInstr *NextInstr = MI.getNextNode();
485 while (NextInstr &&
486 ((Opcode == SPIRV::OpTypeStruct &&
487 NextInstr->getOpcode() == SPIRV::OpTypeStructContinuedINTEL) ||
488 (Opcode == SPIRV::OpConstantComposite &&
489 NextInstr->getOpcode() ==
490 SPIRV::OpConstantCompositeContinuedINTEL))) {
491 MCRegister Tmp = handleTypeDeclOrConstant(MI: *NextInstr, SignatureToGReg);
492 MAI.setRegisterAlias(MF, Reg: NextInstr->getOperand(i: 0).getReg(), AliasReg: Tmp);
493 MAI.setSkipEmission(NextInstr);
494 NextInstr = NextInstr->getNextNode();
495 }
496 } else if (TII->isTypeDeclInstr(MI) || TII->isConstantInstr(MI) ||
497 TII->isInlineAsmDefInstr(MI)) {
498 GReg = handleTypeDeclOrConstant(MI, SignatureToGReg);
499 } else if (Opcode == SPIRV::OpVariable ||
500 Opcode == SPIRV::OpUntypedVariableKHR) {
501 GReg = handleVariable(MF, MI, GlobalToGReg);
502 } else {
503 LLVM_DEBUG({
504 dbgs() << "\nInstruction: ";
505 MI.print(dbgs());
506 dbgs() << "\n";
507 });
508 llvm_unreachable("Unexpected instruction is visited");
509 }
510 MAI.setRegisterAlias(MF, Reg: MI.getOperand(i: 0).getReg(), AliasReg: GReg);
511 if (!IsFunDef)
512 MAI.setSkipEmission(&MI);
513}
514
515MCRegister SPIRVModuleAnalysisImpl::handleFunctionOrParameter(
516 const MachineFunction *MF, const MachineInstr &MI,
517 std::map<const Value *, unsigned> &GlobalToGReg, bool &IsFunDef) {
518 const Value *GObj = GR->getGlobalObject(MF, R: MI.getOperand(i: 0).getReg());
519 assert(GObj && "Unregistered global definition");
520 const Function *F = dyn_cast<Function>(Val: GObj);
521 if (!F)
522 F = dyn_cast<Argument>(Val: GObj)->getParent();
523 assert(F && "Expected a reference to a function or an argument");
524 IsFunDef = !F->isDeclaration();
525 auto [It, Inserted] = GlobalToGReg.try_emplace(k: GObj);
526 if (!Inserted)
527 return It->second;
528 MCRegister GReg = MAI.getNextIDRegister();
529 It->second = GReg;
530 if (!IsFunDef)
531 MAI.MS[SPIRV::MB_ExtFuncDecls].push_back(Elt: &MI);
532 return GReg;
533}
534
535MCRegister SPIRVModuleAnalysisImpl::handleTypeDeclOrConstant(
536 const MachineInstr &MI, InstrGRegsMap &SignatureToGReg) {
537 InstrSignature MISign = instrToSignature(MI, MAI, UseDefReg: false);
538 auto [It, Inserted] = SignatureToGReg.try_emplace(k: MISign);
539 if (!Inserted)
540 return It->second;
541 MCRegister GReg = MAI.getNextIDRegister();
542 It->second = GReg;
543 MAI.MS[SPIRV::MB_TypeConstVars].push_back(Elt: &MI);
544 return GReg;
545}
546
547MCRegister SPIRVModuleAnalysisImpl::handleVariable(
548 const MachineFunction *MF, const MachineInstr &MI,
549 std::map<const Value *, unsigned> &GlobalToGReg) {
550 MAI.GlobalVarList.push_back(Elt: &MI);
551 const Value *GObj = GR->getGlobalObject(MF, R: MI.getOperand(i: 0).getReg());
552 assert(GObj && "Unregistered global definition");
553 auto [It, Inserted] = GlobalToGReg.try_emplace(k: GObj);
554 if (!Inserted)
555 return It->second;
556 MCRegister GReg = MAI.getNextIDRegister();
557 It->second = GReg;
558 MAI.MS[SPIRV::MB_TypeConstVars].push_back(Elt: &MI);
559 if (const auto *GV = dyn_cast<GlobalVariable>(Val: GObj))
560 MAI.GlobalObjMap[GV] = GReg;
561 return GReg;
562}
563
564void SPIRVModuleAnalysisImpl::collectDeclarations(const Module &M) {
565 InstrGRegsMap SignatureToGReg;
566 std::map<const Value *, unsigned> GlobalToGReg;
567 for (const Function &F : M) {
568 MachineFunction *MF = GetMF(F);
569 if (!MF)
570 continue;
571 const MachineRegisterInfo &MRI = MF->getRegInfo();
572 unsigned PastHeader = 0;
573 for (MachineBasicBlock &MBB : *MF) {
574 for (MachineInstr &MI : MBB) {
575 if (MI.getNumOperands() == 0)
576 continue;
577 unsigned Opcode = MI.getOpcode();
578 if (Opcode == SPIRV::OpFunction) {
579 if (PastHeader == 0) {
580 PastHeader = 1;
581 continue;
582 }
583 } else if (Opcode == SPIRV::OpFunctionParameter) {
584 if (PastHeader < 2)
585 continue;
586 } else if (PastHeader > 0) {
587 PastHeader = 2;
588 }
589
590 const MachineOperand &DefMO = MI.getOperand(i: 0);
591 switch (Opcode) {
592 case SPIRV::OpExtension:
593 MAI.Reqs.addExtension(ToAdd: SPIRV::Extension::Extension(DefMO.getImm()));
594 MAI.setSkipEmission(&MI);
595 break;
596 case SPIRV::OpCapability:
597 MAI.Reqs.addCapability(ToAdd: SPIRV::Capability::Capability(DefMO.getImm()));
598 MAI.setSkipEmission(&MI);
599 if (PastHeader > 0)
600 PastHeader = 2;
601 break;
602 default:
603 if (DefMO.isReg() && isDeclSection(MRI, MI) &&
604 !MAI.hasRegisterAlias(MF, Reg: DefMO.getReg()))
605 visitDecl(MRI, SignatureToGReg, GlobalToGReg, MF, MI);
606 // Device enqueue instructions are not decls, but their Invoke
607 // operand may be a function-pointer placeholder OpUndef. Resolve it
608 // to the OpFunction's global <id> via visitFunPtrUse.
609 if (unsigned InvokeIdx = getInvokeOperandIdx(Opcode)) {
610 const MachineOperand &InvokeMO = MI.getOperand(i: InvokeIdx);
611 if (InvokeMO.isReg()) {
612 Register InvokeReg = InvokeMO.getReg();
613 if (!MAI.hasRegisterAlias(MF, Reg: InvokeReg)) {
614 if (const MachineInstr *DefMI =
615 MRI.getUniqueVRegDef(Reg: InvokeReg)) {
616 if (DefMI->getOpcode() == SPIRV::OpUndef) {
617 const MachineOperand *FunPtrOp = &DefMI->getOperand(i: 0);
618 if (GR->getFunctionDefinitionByUse(Use: FunPtrOp))
619 visitFunPtrUse(OpReg: InvokeReg, FunPtrOp, SignatureToGReg,
620 GlobalToGReg, MF);
621 }
622 }
623 }
624 }
625 }
626 }
627 }
628 }
629 }
630}
631
632// Look for IDs declared with Import linkage, and map the corresponding function
633// to the register defining that variable (which will usually be the result of
634// an OpFunction). This lets us call externally imported functions using
635// the correct ID registers.
636void SPIRVModuleAnalysisImpl::collectFuncNames(MachineInstr &MI,
637 const Function *F) {
638 if (MI.getOpcode() == SPIRV::OpDecorate) {
639 // If it's got Import linkage.
640 auto Dec = MI.getOperand(i: 1).getImm();
641 if (Dec == SPIRV::Decoration::LinkageAttributes) {
642 auto Lnk = MI.getOperand(i: MI.getNumOperands() - 1).getImm();
643 if (Lnk == SPIRV::LinkageType::Import) {
644 // Map imported function name to function ID register.
645 const Function *ImportedFunc =
646 F->getParent()->getFunction(Name: getStringImm(MI, StartIndex: 2));
647 Register Target = MI.getOperand(i: 0).getReg();
648 MAI.GlobalObjMap[ImportedFunc] =
649 MAI.getRegisterAlias(MF: MI.getMF(), Reg: Target);
650 }
651 }
652 } else if (MI.getOpcode() == SPIRV::OpFunction) {
653 // Record all internal OpFunction declarations.
654 Register Reg = MI.defs().begin()->getReg();
655 MCRegister GlobalReg = MAI.getRegisterAlias(MF: MI.getMF(), Reg);
656 assert(GlobalReg.isValid());
657 MAI.GlobalObjMap[F] = GlobalReg;
658 }
659}
660
661// Collect the given instruction in the specified MS. We assume global register
662// numbering has already occurred by this point. We can directly compare reg
663// arguments when detecting duplicates.
664static void collectOtherInstr(MachineInstr &MI, SPIRV::ModuleAnalysisInfo &MAI,
665 SPIRV::ModuleSectionType MSType, InstrTraces &IS,
666 bool Append = true) {
667 MAI.setSkipEmission(&MI);
668 InstrSignature MISign = instrToSignature(MI, MAI, UseDefReg: true);
669 auto FoundMI = IS.insert(x: std::move(MISign));
670 if (!FoundMI.second) {
671 if (MI.getOpcode() == SPIRV::OpDecorate) {
672 assert(MI.getNumOperands() >= 2 &&
673 "Decoration instructions must have at least 2 operands");
674 assert(MSType == SPIRV::MB_Annotations &&
675 "Only OpDecorate instructions can be duplicates");
676 // For FPFastMathMode decoration, we need to merge the flags of the
677 // duplicate decoration with the original one, so we need to find the
678 // original instruction that has the same signature. For the rest of
679 // instructions, we will simply skip the duplicate.
680 if (MI.getOperand(i: 1).getImm() != SPIRV::Decoration::FPFastMathMode)
681 return; // Skip duplicates of other decorations.
682
683 const SPIRV::InstrList &Decorations = MAI.MS[MSType];
684 for (const MachineInstr *OrigMI : Decorations) {
685 if (instrToSignature(MI: *OrigMI, MAI, UseDefReg: true) == MISign) {
686 assert(OrigMI->getNumOperands() == MI.getNumOperands() &&
687 "Original instruction must have the same number of operands");
688 assert(
689 OrigMI->getNumOperands() == 3 &&
690 "FPFastMathMode decoration must have 3 operands for OpDecorate");
691 unsigned OrigFlags = OrigMI->getOperand(i: 2).getImm();
692 unsigned NewFlags = MI.getOperand(i: 2).getImm();
693 if (OrigFlags == NewFlags)
694 return; // No need to merge, the flags are the same.
695
696 // Emit warning about possible conflict between flags.
697 unsigned FinalFlags = OrigFlags | NewFlags;
698 llvm::errs()
699 << "Warning: Conflicting FPFastMathMode decoration flags "
700 "in instruction: "
701 << *OrigMI << "Original flags: " << OrigFlags
702 << ", new flags: " << NewFlags
703 << ". They will be merged on a best effort basis, but not "
704 "validated. Final flags: "
705 << FinalFlags << "\n";
706 MachineInstr *OrigMINonConst = const_cast<MachineInstr *>(OrigMI);
707 MachineOperand &OrigFlagsOp = OrigMINonConst->getOperand(i: 2);
708 OrigFlagsOp = MachineOperand::CreateImm(Val: FinalFlags);
709 return; // Merge done, so we found a duplicate; don't add it to MAI.MS
710 }
711 }
712 assert(false && "No original instruction found for the duplicate "
713 "OpDecorate, but we found one in IS.");
714 }
715 return; // insert failed, so we found a duplicate; don't add it to MAI.MS
716 }
717 // No duplicates, so add it.
718 if (Append)
719 MAI.MS[MSType].push_back(Elt: &MI);
720 else
721 MAI.MS[MSType].insert(I: MAI.MS[MSType].begin(), Elt: &MI);
722}
723
724// Some global instructions make reference to function-local ID regs, so cannot
725// be correctly collected until these registers are globally numbered.
726void SPIRVModuleAnalysisImpl::processOtherInstrs(const Module &M) {
727 InstrTraces IS;
728 for (const Function &F : M) {
729 if (F.isDeclaration())
730 continue;
731 MachineFunction *MF = GetMF(F);
732 assert(MF);
733
734 for (MachineBasicBlock &MBB : *MF)
735 for (MachineInstr &MI : MBB) {
736 if (MAI.getSkipEmission(MI: &MI))
737 continue;
738 const unsigned OpCode = MI.getOpcode();
739 if (OpCode == SPIRV::OpString) {
740 collectOtherInstr(MI, MAI, MSType: SPIRV::MB_DebugStrings, IS);
741 } else if (OpCode == SPIRV::OpExtInst && MI.getOperand(i: 2).isImm() &&
742 MI.getOperand(i: 2).getImm() ==
743 SPIRV::InstructionSet::
744 NonSemantic_Shader_DebugInfo_100) {
745 // TODO: This branch is dead. SPIRVNonSemanticDebugHandler emits NSDI
746 // instructions directly as MCInsts at print time; no
747 // MachineInstructions with the NSDI ext set are created anymore.
748 // Remove this block and
749 // MB_NonSemanticGlobalDI once per-function NSDI emission is confirmed
750 // not to need MIR routing.
751 MachineOperand Ins = MI.getOperand(i: 3);
752 namespace NS = SPIRV::NonSemanticExtInst;
753 static constexpr int64_t GlobalNonSemanticDITy[] = {
754 NS::DebugSource, NS::DebugCompilationUnit, NS::DebugInfoNone,
755 NS::DebugTypeBasic, NS::DebugTypePointer};
756 bool IsGlobalDI = false;
757 for (unsigned Idx = 0; Idx < std::size(GlobalNonSemanticDITy); ++Idx)
758 IsGlobalDI |= Ins.getImm() == GlobalNonSemanticDITy[Idx];
759 if (IsGlobalDI)
760 collectOtherInstr(MI, MAI, MSType: SPIRV::MB_NonSemanticGlobalDI, IS);
761 } else if (OpCode == SPIRV::OpName || OpCode == SPIRV::OpMemberName) {
762 collectOtherInstr(MI, MAI, MSType: SPIRV::MB_DebugNames, IS);
763 } else if (OpCode == SPIRV::OpEntryPoint) {
764 collectOtherInstr(MI, MAI, MSType: SPIRV::MB_EntryPoints, IS);
765 } else if (TII->isAliasingInstr(MI)) {
766 collectOtherInstr(MI, MAI, MSType: SPIRV::MB_AliasingInsts, IS);
767 } else if (TII->isDecorationInstr(MI)) {
768 collectOtherInstr(MI, MAI, MSType: SPIRV::MB_Annotations, IS);
769 collectFuncNames(MI, F: &F);
770 } else if (TII->isConstantInstr(MI)) {
771 // Now OpSpecConstant*s are not in DT,
772 // but they need to be collected anyway.
773 collectOtherInstr(MI, MAI, MSType: SPIRV::MB_TypeConstVars, IS);
774 } else if (OpCode == SPIRV::OpFunction) {
775 collectFuncNames(MI, F: &F);
776 } else if (OpCode == SPIRV::OpTypeForwardPointer) {
777 collectOtherInstr(MI, MAI, MSType: SPIRV::MB_TypeConstVars, IS, Append: false);
778 }
779 }
780 }
781 // Selection order can place a scope/list ahead of a domain/scope it
782 // references. The dependency meanwhile is domain -> scope -> list, so sort
783 // the def before its uses.
784 auto AliasingTier = [](const MachineInstr *MI) {
785 switch (MI->getOpcode()) {
786 case SPIRV::OpAliasDomainDeclINTEL:
787 return 0;
788 case SPIRV::OpAliasScopeDeclINTEL:
789 return 1;
790 case SPIRV::OpAliasScopeListDeclINTEL:
791 return 2;
792 default:
793 llvm_unreachable("unexpected aliasing instruction");
794 }
795 };
796 stable_sort(Range&: MAI.MS[SPIRV::MB_AliasingInsts],
797 C: [&](const MachineInstr *LHS, const MachineInstr *RHS) {
798 return AliasingTier(LHS) < AliasingTier(RHS);
799 });
800}
801
802// Number registers in all functions globally from 0 onwards and store
803// the result in global register alias table. Some registers are already
804// numbered.
805void SPIRVModuleAnalysisImpl::numberRegistersGlobally(const Module &M) {
806 for (const Function &F : M) {
807 if (F.isDeclaration())
808 continue;
809 MachineFunction *MF = GetMF(F);
810 assert(MF);
811 for (MachineBasicBlock &MBB : *MF) {
812 for (MachineInstr &MI : MBB) {
813 for (MachineOperand &Op : MI.operands()) {
814 if (!Op.isReg())
815 continue;
816 Register Reg = Op.getReg();
817 if (MAI.hasRegisterAlias(MF, Reg))
818 continue;
819 MCRegister NewReg = MAI.getNextIDRegister();
820 MAI.setRegisterAlias(MF, Reg, AliasReg: NewReg);
821 }
822 if (MI.getOpcode() != SPIRV::OpExtInst)
823 continue;
824 auto Set = MI.getOperand(i: 2).getImm();
825 auto [It, Inserted] = MAI.ExtInstSetMap.try_emplace(Key: Set);
826 if (Inserted)
827 It->second = MAI.getNextIDRegister();
828 }
829 }
830 }
831}
832
833// RequirementHandler implementations.
834void SPIRV::RequirementHandler::getAndAddRequirements(
835 SPIRV::OperandCategory::OperandCategory Category, uint32_t i,
836 const SPIRVSubtarget &ST) {
837 addRequirements(Req: getSymbolicOperandRequirements(Category, i, ST, Reqs&: *this));
838}
839
840void SPIRV::RequirementHandler::recursiveAddCapabilities(
841 const CapabilityList &ToPrune) {
842 for (const auto &Cap : ToPrune) {
843 AllCaps.insert(V: Cap);
844 CapabilityList ImplicitDecls =
845 getSymbolicOperandCapabilities(Category: OperandCategory::CapabilityOperand, Value: Cap);
846 recursiveAddCapabilities(ToPrune: ImplicitDecls);
847 }
848}
849
850void SPIRV::RequirementHandler::addCapabilities(const CapabilityList &ToAdd) {
851 for (const auto &Cap : ToAdd) {
852 bool IsNewlyInserted = AllCaps.insert(V: Cap).second;
853 if (!IsNewlyInserted) // Don't re-add if it's already been declared.
854 continue;
855 CapabilityList ImplicitDecls =
856 getSymbolicOperandCapabilities(Category: OperandCategory::CapabilityOperand, Value: Cap);
857 recursiveAddCapabilities(ToPrune: ImplicitDecls);
858 MinimalCaps.push_back(Elt: Cap);
859 }
860}
861
862void SPIRV::RequirementHandler::addRequirements(
863 const SPIRV::Requirements &Req) {
864 if (!Req.IsSatisfiable)
865 report_fatal_error(reason: "Adding SPIR-V requirements this target can't satisfy.");
866
867 if (Req.Cap.has_value())
868 addCapabilities(ToAdd: {Req.Cap.value()});
869
870 addExtensions(ToAdd: Req.Exts);
871
872 if (!Req.MinVer.empty()) {
873 if (!MaxVersion.empty() && Req.MinVer > MaxVersion) {
874 LLVM_DEBUG(dbgs() << "Conflicting version requirements: >= " << Req.MinVer
875 << " and <= " << MaxVersion << "\n");
876 report_fatal_error(reason: "Adding SPIR-V requirements that can't be satisfied.");
877 }
878
879 if (MinVersion.empty() || Req.MinVer > MinVersion)
880 MinVersion = Req.MinVer;
881 }
882
883 if (!Req.MaxVer.empty()) {
884 if (!MinVersion.empty() && Req.MaxVer < MinVersion) {
885 LLVM_DEBUG(dbgs() << "Conflicting version requirements: <= " << Req.MaxVer
886 << " and >= " << MinVersion << "\n");
887 report_fatal_error(reason: "Adding SPIR-V requirements that can't be satisfied.");
888 }
889
890 if (MaxVersion.empty() || Req.MaxVer < MaxVersion)
891 MaxVersion = Req.MaxVer;
892 }
893}
894
895void SPIRV::RequirementHandler::checkSatisfiable(
896 const SPIRVSubtarget &ST) const {
897 // Report as many errors as possible before aborting the compilation.
898 bool IsSatisfiable = true;
899 auto TargetVer = ST.getSPIRVVersion();
900
901 if (!MaxVersion.empty() && !TargetVer.empty() && MaxVersion < TargetVer) {
902 LLVM_DEBUG(
903 dbgs() << "Target SPIR-V version too high for required features\n"
904 << "Required max version: " << MaxVersion << " target version "
905 << TargetVer << "\n");
906 IsSatisfiable = false;
907 }
908
909 if (!MinVersion.empty() && !TargetVer.empty() && MinVersion > TargetVer) {
910 LLVM_DEBUG(dbgs() << "Target SPIR-V version too low for required features\n"
911 << "Required min version: " << MinVersion
912 << " target version " << TargetVer << "\n");
913 IsSatisfiable = false;
914 }
915
916 if (!MinVersion.empty() && !MaxVersion.empty() && MinVersion > MaxVersion) {
917 LLVM_DEBUG(
918 dbgs()
919 << "Version is too low for some features and too high for others.\n"
920 << "Required SPIR-V min version: " << MinVersion
921 << " required SPIR-V max version " << MaxVersion << "\n");
922 IsSatisfiable = false;
923 }
924
925 AvoidCapabilitiesSet AvoidCaps;
926 if (!ST.isShader())
927 AvoidCaps.S.insert(V: SPIRV::Capability::Shader);
928 else
929 AvoidCaps.S.insert(V: SPIRV::Capability::Kernel);
930
931 for (auto Cap : MinimalCaps) {
932 if (AvailableCaps.contains(V: Cap) && !AvoidCaps.S.contains(V: Cap))
933 continue;
934 LLVM_DEBUG(dbgs() << "Capability not supported: "
935 << getSymbolicOperandMnemonic(
936 OperandCategory::CapabilityOperand, Cap)
937 << "\n");
938 IsSatisfiable = false;
939 }
940
941 for (auto Ext : AllExtensions) {
942 if (ST.canUseExtension(E: Ext))
943 continue;
944 LLVM_DEBUG(dbgs() << "Extension not supported: "
945 << getSymbolicOperandMnemonic(
946 OperandCategory::ExtensionOperand, Ext)
947 << "\n");
948 IsSatisfiable = false;
949 }
950
951 if (!IsSatisfiable)
952 report_fatal_error(reason: "Unable to meet SPIR-V requirements for this target.");
953}
954
955// Add the given capabilities and all their implicitly defined capabilities too.
956void SPIRV::RequirementHandler::addAvailableCaps(const CapabilityList &ToAdd) {
957 for (const auto Cap : ToAdd)
958 if (AvailableCaps.insert(V: Cap).second)
959 addAvailableCaps(ToAdd: getSymbolicOperandCapabilities(
960 Category: SPIRV::OperandCategory::CapabilityOperand, Value: Cap));
961}
962
963void SPIRV::RequirementHandler::removeCapabilityIf(
964 const Capability::Capability ToRemove,
965 const Capability::Capability IfPresent) {
966 if (AllCaps.contains(V: IfPresent)) {
967 AllCaps.erase(V: ToRemove);
968 llvm::erase(C&: MinimalCaps, V: ToRemove);
969 }
970}
971
972namespace llvm {
973namespace SPIRV {
974void RequirementHandler::initAvailableCapabilities(const SPIRVSubtarget &ST) {
975 // Provided by both all supported Vulkan versions and OpenCl.
976 addAvailableCaps(ToAdd: {Capability::Shader, Capability::Linkage, Capability::Int8,
977 Capability::Int16});
978
979 if (ST.isAtLeastSPIRVVer(VerToCompareTo: VersionTuple(1, 3)))
980 addAvailableCaps(ToAdd: {Capability::GroupNonUniform,
981 Capability::GroupNonUniformVote,
982 Capability::GroupNonUniformArithmetic,
983 Capability::GroupNonUniformBallot,
984 Capability::GroupNonUniformClustered,
985 Capability::GroupNonUniformShuffle,
986 Capability::GroupNonUniformShuffleRelative,
987 Capability::GroupNonUniformQuad});
988
989 if (ST.isAtLeastSPIRVVer(VerToCompareTo: VersionTuple(1, 6)))
990 addAvailableCaps(ToAdd: {Capability::DotProduct, Capability::DotProductInputAll,
991 Capability::DotProductInput4x8Bit,
992 Capability::DotProductInput4x8BitPacked,
993 Capability::DemoteToHelperInvocation});
994
995 // Add capabilities enabled by extensions.
996 for (auto Extension : ST.getAllAvailableExtensions()) {
997 CapabilityList EnabledCapabilities =
998 getCapabilitiesEnabledByExtension(Extension);
999 addAvailableCaps(ToAdd: EnabledCapabilities);
1000 }
1001
1002 if (!ST.isShader()) {
1003 initAvailableCapabilitiesForOpenCL(ST);
1004 return;
1005 }
1006
1007 if (ST.isShader()) {
1008 initAvailableCapabilitiesForVulkan(ST);
1009 return;
1010 }
1011
1012 report_fatal_error(reason: "Unimplemented environment for SPIR-V generation.");
1013}
1014
1015void RequirementHandler::initAvailableCapabilitiesForOpenCL(
1016 const SPIRVSubtarget &ST) {
1017 // Add the min requirements for different OpenCL and SPIR-V versions.
1018 addAvailableCaps(ToAdd: {Capability::Addresses, Capability::Float16Buffer,
1019 Capability::Kernel, Capability::Vector16,
1020 Capability::Groups, Capability::GenericPointer,
1021 Capability::StorageImageWriteWithoutFormat,
1022 Capability::StorageImageReadWithoutFormat});
1023 if (ST.hasOpenCLFullProfile())
1024 addAvailableCaps(ToAdd: {Capability::Int64, Capability::Int64Atomics});
1025 if (ST.hasOpenCLImageSupport()) {
1026 addAvailableCaps(ToAdd: {Capability::ImageBasic, Capability::LiteralSampler,
1027 Capability::Image1D, Capability::SampledBuffer,
1028 Capability::ImageBuffer});
1029 if (ST.isAtLeastOpenCLVer(VerToCompareTo: VersionTuple(2, 0)))
1030 addAvailableCaps(ToAdd: {Capability::ImageReadWrite});
1031 }
1032 if (ST.isAtLeastSPIRVVer(VerToCompareTo: VersionTuple(1, 1)) &&
1033 ST.isAtLeastOpenCLVer(VerToCompareTo: VersionTuple(2, 2)))
1034 addAvailableCaps(ToAdd: {Capability::SubgroupDispatch, Capability::PipeStorage});
1035 if (ST.isAtLeastSPIRVVer(VerToCompareTo: VersionTuple(1, 4)))
1036 addAvailableCaps(ToAdd: {Capability::DenormPreserve, Capability::DenormFlushToZero,
1037 Capability::SignedZeroInfNanPreserve,
1038 Capability::RoundingModeRTE,
1039 Capability::RoundingModeRTZ});
1040 // TODO: verify if this needs some checks.
1041 addAvailableCaps(ToAdd: {Capability::Float16, Capability::Float64});
1042
1043 // TODO: add OpenCL extensions.
1044}
1045
1046void RequirementHandler::initAvailableCapabilitiesForVulkan(
1047 const SPIRVSubtarget &ST) {
1048
1049 // Core in Vulkan 1.1 and earlier.
1050 addAvailableCaps(ToAdd: {Capability::Int64,
1051 Capability::Float16,
1052 Capability::Float64,
1053 Capability::GroupNonUniform,
1054 Capability::Image1D,
1055 Capability::SampledBuffer,
1056 Capability::ImageBuffer,
1057 Capability::UniformBufferArrayDynamicIndexing,
1058 Capability::SampledImageArrayDynamicIndexing,
1059 Capability::StorageBufferArrayDynamicIndexing,
1060 Capability::StorageImageArrayDynamicIndexing,
1061 Capability::DerivativeControl,
1062 Capability::MinLod,
1063 Capability::ImageQuery,
1064 Capability::ImageGatherExtended,
1065 Capability::Addresses,
1066 Capability::VulkanMemoryModelKHR,
1067 Capability::StorageImageExtendedFormats,
1068 Capability::StorageImageMultisample,
1069 Capability::ImageMSArray});
1070
1071 if (ST.isAtLeastSPIRVVer(VerToCompareTo: VersionTuple(1, 3)) ||
1072 ST.canUseExtension(E: Extension::SPV_KHR_variable_pointers))
1073 addAvailableCaps(ToAdd: {Capability::VariablePointersStorageBuffer,
1074 Capability::VariablePointers});
1075
1076 // Became core in Vulkan 1.2
1077 if (ST.isAtLeastSPIRVVer(VerToCompareTo: VersionTuple(1, 5))) {
1078 addAvailableCaps(
1079 ToAdd: {Capability::Int64Atomics, Capability::ShaderNonUniformEXT,
1080 Capability::RuntimeDescriptorArrayEXT,
1081 Capability::InputAttachmentArrayDynamicIndexingEXT,
1082 Capability::UniformTexelBufferArrayDynamicIndexingEXT,
1083 Capability::StorageTexelBufferArrayDynamicIndexingEXT,
1084 Capability::UniformBufferArrayNonUniformIndexingEXT,
1085 Capability::SampledImageArrayNonUniformIndexingEXT,
1086 Capability::StorageBufferArrayNonUniformIndexingEXT,
1087 Capability::StorageImageArrayNonUniformIndexingEXT,
1088 Capability::InputAttachmentArrayNonUniformIndexingEXT,
1089 Capability::UniformTexelBufferArrayNonUniformIndexingEXT,
1090 Capability::StorageTexelBufferArrayNonUniformIndexingEXT});
1091 }
1092
1093 // Became core in Vulkan 1.3
1094 if (ST.isAtLeastSPIRVVer(VerToCompareTo: VersionTuple(1, 6)))
1095 addAvailableCaps(ToAdd: {Capability::StorageImageWriteWithoutFormat,
1096 Capability::StorageImageReadWithoutFormat});
1097}
1098
1099} // namespace SPIRV
1100} // namespace llvm
1101
1102// Add the required capabilities from a decoration instruction (including
1103// BuiltIns).
1104static void addOpDecorateReqs(const MachineInstr &MI, unsigned DecIndex,
1105 SPIRV::RequirementHandler &Reqs,
1106 const SPIRVSubtarget &ST) {
1107 int64_t DecOp = MI.getOperand(i: DecIndex).getImm();
1108 auto Dec = static_cast<SPIRV::Decoration::Decoration>(DecOp);
1109 Reqs.addRequirements(Req: getSymbolicOperandRequirements(
1110 Category: SPIRV::OperandCategory::DecorationOperand, i: Dec, ST, Reqs));
1111
1112 if (Dec == SPIRV::Decoration::BuiltIn) {
1113 int64_t BuiltInOp = MI.getOperand(i: DecIndex + 1).getImm();
1114 auto BuiltIn = static_cast<SPIRV::BuiltIn::BuiltIn>(BuiltInOp);
1115 Reqs.addRequirements(Req: getSymbolicOperandRequirements(
1116 Category: SPIRV::OperandCategory::BuiltInOperand, i: BuiltIn, ST, Reqs));
1117 } else if (Dec == SPIRV::Decoration::LinkageAttributes) {
1118 int64_t LinkageOp = MI.getOperand(i: MI.getNumOperands() - 1).getImm();
1119 SPIRV::LinkageType::LinkageType LnkType =
1120 static_cast<SPIRV::LinkageType::LinkageType>(LinkageOp);
1121 if (LnkType == SPIRV::LinkageType::LinkOnceODR)
1122 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_KHR_linkonce_odr);
1123 else if (LnkType == SPIRV::LinkageType::WeakAMD) {
1124 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_AMD_weak_linkage);
1125 Reqs.addCapability(ToAdd: SPIRV::Capability::WeakLinkageAMD);
1126 }
1127 } else if (Dec == SPIRV::Decoration::CacheControlLoadINTEL ||
1128 Dec == SPIRV::Decoration::CacheControlStoreINTEL) {
1129 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_cache_controls);
1130 } else if (Dec == SPIRV::Decoration::HostAccessINTEL) {
1131 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_global_variable_host_access);
1132 } else if (Dec == SPIRV::Decoration::InitModeINTEL ||
1133 Dec == SPIRV::Decoration::ImplementInRegisterMapINTEL) {
1134 Reqs.addExtension(
1135 ToAdd: SPIRV::Extension::SPV_INTEL_global_variable_fpga_decorations);
1136 } else if (Dec == SPIRV::Decoration::NonUniformEXT) {
1137 Reqs.addRequirements(Req: SPIRV::Capability::ShaderNonUniformEXT);
1138 } else if (Dec == SPIRV::Decoration::FPMaxErrorDecorationINTEL) {
1139 Reqs.addRequirements(Req: SPIRV::Capability::FPMaxErrorINTEL);
1140 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_fp_max_error);
1141 } else if (Dec == SPIRV::Decoration::FPFastMathMode) {
1142 if (ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_float_controls2)) {
1143 Reqs.addRequirements(Req: SPIRV::Capability::FloatControls2);
1144 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_KHR_float_controls2);
1145 }
1146 }
1147}
1148
1149// Add requirements for image handling.
1150static void addOpTypeImageReqs(const MachineInstr &MI,
1151 SPIRV::RequirementHandler &Reqs,
1152 const SPIRVSubtarget &ST) {
1153 assert(MI.getNumOperands() >= 8 && "Insufficient operands for OpTypeImage");
1154 // The operand indices used here are based on the OpTypeImage layout, which
1155 // the MachineInstr follows as well.
1156 int64_t ImgFormatOp = MI.getOperand(i: 7).getImm();
1157 auto ImgFormat = static_cast<SPIRV::ImageFormat::ImageFormat>(ImgFormatOp);
1158 Reqs.getAndAddRequirements(Category: SPIRV::OperandCategory::ImageFormatOperand,
1159 i: ImgFormat, ST);
1160
1161 bool IsArrayed = MI.getOperand(i: 4).getImm() == 1;
1162 bool IsMultisampled = MI.getOperand(i: 5).getImm() == 1;
1163 bool NoSampler = MI.getOperand(i: 6).getImm() == 2;
1164 // Add dimension requirements.
1165 assert(MI.getOperand(2).isImm());
1166 switch (MI.getOperand(i: 2).getImm()) {
1167 case SPIRV::Dim::DIM_1D:
1168 Reqs.addRequirements(Req: NoSampler ? SPIRV::Capability::Image1D
1169 : SPIRV::Capability::Sampled1D);
1170 break;
1171 case SPIRV::Dim::DIM_2D:
1172 if (IsMultisampled && NoSampler)
1173 Reqs.addRequirements(Req: SPIRV::Capability::StorageImageMultisample);
1174 if (IsMultisampled && IsArrayed)
1175 Reqs.addRequirements(Req: SPIRV::Capability::ImageMSArray);
1176 break;
1177 case SPIRV::Dim::DIM_3D:
1178 break;
1179 case SPIRV::Dim::DIM_Cube:
1180 Reqs.addRequirements(Req: SPIRV::Capability::Shader);
1181 if (IsArrayed)
1182 Reqs.addRequirements(Req: NoSampler ? SPIRV::Capability::ImageCubeArray
1183 : SPIRV::Capability::SampledCubeArray);
1184 break;
1185 case SPIRV::Dim::DIM_Rect:
1186 Reqs.addRequirements(Req: NoSampler ? SPIRV::Capability::ImageRect
1187 : SPIRV::Capability::SampledRect);
1188 break;
1189 case SPIRV::Dim::DIM_Buffer:
1190 Reqs.addRequirements(Req: NoSampler ? SPIRV::Capability::ImageBuffer
1191 : SPIRV::Capability::SampledBuffer);
1192 break;
1193 case SPIRV::Dim::DIM_SubpassData:
1194 Reqs.addRequirements(Req: SPIRV::Capability::InputAttachment);
1195 break;
1196 }
1197
1198 // Has optional access qualifier.
1199 if (!ST.isShader()) {
1200 if (MI.getNumOperands() > 8 &&
1201 MI.getOperand(i: 8).getImm() == SPIRV::AccessQualifier::ReadWrite)
1202 Reqs.addRequirements(Req: SPIRV::Capability::ImageReadWrite);
1203 else
1204 Reqs.addRequirements(Req: SPIRV::Capability::ImageBasic);
1205 }
1206}
1207
1208static bool isBFloat16Type(SPIRVTypeInst TypeDef) {
1209 return TypeDef && TypeDef->getNumOperands() == 3 &&
1210 TypeDef->getOpcode() == SPIRV::OpTypeFloat &&
1211 TypeDef->getOperand(i: 1).getImm() == 16 &&
1212 TypeDef->getOperand(i: 2).getImm() == SPIRV::FPEncoding::BFloat16KHR;
1213}
1214
1215// Add requirements for handling atomic float instructions
1216#define ATOM_FLT_REQ_EXT_MSG(ExtName) \
1217 "The atomic float instruction requires the following SPIR-V " \
1218 "extension: SPV_EXT_shader_atomic_float" ExtName
1219static void AddAtomicVectorFloatRequirements(const MachineInstr &MI,
1220 SPIRV::RequirementHandler &Reqs,
1221 const SPIRVSubtarget &ST) {
1222 SPIRVTypeInst VecTypeDef =
1223 MI.getMF()->getRegInfo().getVRegDef(Reg: MI.getOperand(i: 1).getReg());
1224
1225 const unsigned Rank = VecTypeDef->getOperand(i: 2).getImm();
1226 if (Rank != 2 && Rank != 4)
1227 reportFatalUsageError(reason: "Result type of an atomic vector float instruction "
1228 "must be a 2-component or 4 component vector");
1229
1230 SPIRVTypeInst EltTypeDef =
1231 MI.getMF()->getRegInfo().getVRegDef(Reg: VecTypeDef->getOperand(i: 1).getReg());
1232
1233 if (EltTypeDef->getOpcode() != SPIRV::OpTypeFloat ||
1234 EltTypeDef->getOperand(i: 1).getImm() != 16)
1235 reportFatalUsageError(
1236 reason: "The element type for the result type of an atomic vector float "
1237 "instruction must be a 16-bit floating-point scalar");
1238
1239 // The extension is defined for fp16, but the AMD target lets a bf16 vector
1240 // use the same instruction so it can lower to a packed bf16 atomic.
1241 if (isBFloat16Type(TypeDef: EltTypeDef) &&
1242 ST.getTargetTriple().getVendor() != Triple::AMD)
1243 reportFatalUsageError(
1244 reason: "The element type for the result type of an atomic vector float "
1245 "instruction cannot be a bfloat16 scalar");
1246 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_NV_shader_atomic_fp16_vector))
1247 reportFatalUsageError(
1248 reason: "The atomic float16 vector instruction requires the following SPIR-V "
1249 "extension: SPV_NV_shader_atomic_fp16_vector");
1250
1251 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_NV_shader_atomic_fp16_vector);
1252 Reqs.addCapability(ToAdd: SPIRV::Capability::AtomicFloat16VectorNV);
1253}
1254
1255static void AddAtomicFloatRequirements(const MachineInstr &MI,
1256 SPIRV::RequirementHandler &Reqs,
1257 const SPIRVSubtarget &ST) {
1258 assert(MI.getOperand(1).isReg() &&
1259 "Expect register operand in atomic float instruction");
1260 Register TypeReg = MI.getOperand(i: 1).getReg();
1261 SPIRVTypeInst TypeDef = MI.getMF()->getRegInfo().getVRegDef(Reg: TypeReg);
1262
1263 if (isVectorType(SPVTy: TypeDef))
1264 return AddAtomicVectorFloatRequirements(MI, Reqs, ST);
1265
1266 if (TypeDef->getOpcode() != SPIRV::OpTypeFloat)
1267 report_fatal_error(reason: "Result type of an atomic float instruction must be a "
1268 "floating-point type scalar");
1269
1270 unsigned BitWidth = TypeDef->getOperand(i: 1).getImm();
1271 unsigned Op = MI.getOpcode();
1272 if (Op == SPIRV::OpAtomicFAddEXT) {
1273 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_EXT_shader_atomic_float_add))
1274 report_fatal_error(ATOM_FLT_REQ_EXT_MSG("_add"), gen_crash_diag: false);
1275 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_EXT_shader_atomic_float_add);
1276 switch (BitWidth) {
1277 case 16:
1278 if (isBFloat16Type(TypeDef)) {
1279 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_16bit_atomics))
1280 report_fatal_error(
1281 reason: "The atomic bfloat16 instruction requires the following SPIR-V "
1282 "extension: SPV_INTEL_16bit_atomics",
1283 gen_crash_diag: false);
1284 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_16bit_atomics);
1285 Reqs.addCapability(ToAdd: SPIRV::Capability::AtomicBFloat16AddINTEL);
1286 } else {
1287 if (!ST.canUseExtension(
1288 E: SPIRV::Extension::SPV_EXT_shader_atomic_float16_add))
1289 report_fatal_error(ATOM_FLT_REQ_EXT_MSG("16_add"), gen_crash_diag: false);
1290 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_EXT_shader_atomic_float16_add);
1291 Reqs.addCapability(ToAdd: SPIRV::Capability::AtomicFloat16AddEXT);
1292 }
1293 break;
1294 case 32:
1295 Reqs.addCapability(ToAdd: SPIRV::Capability::AtomicFloat32AddEXT);
1296 break;
1297 case 64:
1298 Reqs.addCapability(ToAdd: SPIRV::Capability::AtomicFloat64AddEXT);
1299 break;
1300 default:
1301 report_fatal_error(
1302 reason: "Unexpected floating-point type width in atomic float instruction");
1303 }
1304 } else {
1305 if (!ST.canUseExtension(
1306 E: SPIRV::Extension::SPV_EXT_shader_atomic_float_min_max))
1307 report_fatal_error(ATOM_FLT_REQ_EXT_MSG("_min_max"), gen_crash_diag: false);
1308 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_EXT_shader_atomic_float_min_max);
1309 switch (BitWidth) {
1310 case 16:
1311 if (isBFloat16Type(TypeDef)) {
1312 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_16bit_atomics))
1313 report_fatal_error(
1314 reason: "The atomic bfloat16 instruction requires the following SPIR-V "
1315 "extension: SPV_INTEL_16bit_atomics",
1316 gen_crash_diag: false);
1317 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_16bit_atomics);
1318 Reqs.addCapability(ToAdd: SPIRV::Capability::AtomicBFloat16MinMaxINTEL);
1319 } else {
1320 Reqs.addCapability(ToAdd: SPIRV::Capability::AtomicFloat16MinMaxEXT);
1321 }
1322 break;
1323 case 32:
1324 Reqs.addCapability(ToAdd: SPIRV::Capability::AtomicFloat32MinMaxEXT);
1325 break;
1326 case 64:
1327 Reqs.addCapability(ToAdd: SPIRV::Capability::AtomicFloat64MinMaxEXT);
1328 break;
1329 default:
1330 report_fatal_error(
1331 reason: "Unexpected floating-point type width in atomic float instruction");
1332 }
1333 }
1334}
1335
1336bool isUniformTexelBuffer(MachineInstr *ImageInst) {
1337 if (ImageInst->getOpcode() != SPIRV::OpTypeImage)
1338 return false;
1339 uint32_t Dim = ImageInst->getOperand(i: 2).getImm();
1340 uint32_t Sampled = ImageInst->getOperand(i: 6).getImm();
1341 return Dim == SPIRV::Dim::DIM_Buffer && Sampled == 1;
1342}
1343
1344bool isStorageTexelBuffer(MachineInstr *ImageInst) {
1345 if (ImageInst->getOpcode() != SPIRV::OpTypeImage)
1346 return false;
1347 uint32_t Dim = ImageInst->getOperand(i: 2).getImm();
1348 uint32_t Sampled = ImageInst->getOperand(i: 6).getImm();
1349 return Dim == SPIRV::Dim::DIM_Buffer && Sampled == 2;
1350}
1351
1352bool isSampledImage(MachineInstr *ImageInst) {
1353 if (ImageInst->getOpcode() != SPIRV::OpTypeImage)
1354 return false;
1355 uint32_t Dim = ImageInst->getOperand(i: 2).getImm();
1356 uint32_t Sampled = ImageInst->getOperand(i: 6).getImm();
1357 return Dim != SPIRV::Dim::DIM_Buffer && Sampled == 1;
1358}
1359
1360bool isInputAttachment(MachineInstr *ImageInst) {
1361 if (ImageInst->getOpcode() != SPIRV::OpTypeImage)
1362 return false;
1363 uint32_t Dim = ImageInst->getOperand(i: 2).getImm();
1364 uint32_t Sampled = ImageInst->getOperand(i: 6).getImm();
1365 return Dim == SPIRV::Dim::DIM_SubpassData && Sampled == 2;
1366}
1367
1368bool isStorageImage(MachineInstr *ImageInst) {
1369 if (ImageInst->getOpcode() != SPIRV::OpTypeImage)
1370 return false;
1371 uint32_t Dim = ImageInst->getOperand(i: 2).getImm();
1372 uint32_t Sampled = ImageInst->getOperand(i: 6).getImm();
1373 return Dim != SPIRV::Dim::DIM_Buffer && Sampled == 2;
1374}
1375
1376bool isCombinedImageSampler(MachineInstr *SampledImageInst) {
1377 if (SampledImageInst->getOpcode() != SPIRV::OpTypeSampledImage)
1378 return false;
1379
1380 const MachineRegisterInfo &MRI = SampledImageInst->getMF()->getRegInfo();
1381 Register ImageReg = SampledImageInst->getOperand(i: 1).getReg();
1382 auto *ImageInst = MRI.getUniqueVRegDef(Reg: ImageReg);
1383 return isSampledImage(ImageInst);
1384}
1385
1386bool hasNonUniformDecoration(Register Reg, const MachineRegisterInfo &MRI) {
1387 for (const auto &MI : MRI.reg_instructions(Reg)) {
1388 if (MI.getOpcode() != SPIRV::OpDecorate)
1389 continue;
1390
1391 uint32_t Dec = MI.getOperand(i: 1).getImm();
1392 if (Dec == SPIRV::Decoration::NonUniformEXT)
1393 return true;
1394 }
1395 return false;
1396}
1397
1398void addOpAccessChainReqs(const MachineInstr &Instr,
1399 SPIRV::RequirementHandler &Handler,
1400 const SPIRVSubtarget &Subtarget) {
1401 const MachineRegisterInfo &MRI = Instr.getMF()->getRegInfo();
1402 // Get the result type. If it is an image type, then the shader uses
1403 // descriptor indexing. The appropriate capabilities will be added based
1404 // on the specifics of the image.
1405 Register ResTypeReg = Instr.getOperand(i: 1).getReg();
1406 MachineInstr *ResTypeInst = MRI.getUniqueVRegDef(Reg: ResTypeReg);
1407
1408 assert(ResTypeInst->getOpcode() == SPIRV::OpTypePointer);
1409 uint32_t StorageClass = ResTypeInst->getOperand(i: 1).getImm();
1410 if (StorageClass != SPIRV::StorageClass::StorageClass::UniformConstant &&
1411 StorageClass != SPIRV::StorageClass::StorageClass::Uniform &&
1412 StorageClass != SPIRV::StorageClass::StorageClass::StorageBuffer) {
1413 return;
1414 }
1415
1416 bool IsNonUniform =
1417 hasNonUniformDecoration(Reg: Instr.getOperand(i: 0).getReg(), MRI);
1418
1419 auto FirstIndexReg = Instr.getOperand(i: 3).getReg();
1420 bool FirstIndexIsConstant =
1421 Subtarget.getInstrInfo()->isConstantInstr(MI: *MRI.getVRegDef(Reg: FirstIndexReg));
1422
1423 if (StorageClass == SPIRV::StorageClass::StorageClass::StorageBuffer) {
1424 if (IsNonUniform)
1425 Handler.addRequirements(
1426 Req: SPIRV::Capability::StorageBufferArrayNonUniformIndexingEXT);
1427 else if (!FirstIndexIsConstant)
1428 Handler.addRequirements(
1429 Req: SPIRV::Capability::StorageBufferArrayDynamicIndexing);
1430 return;
1431 }
1432
1433 Register PointeeTypeReg = ResTypeInst->getOperand(i: 2).getReg();
1434 MachineInstr *PointeeType = MRI.getUniqueVRegDef(Reg: PointeeTypeReg);
1435 if (PointeeType->getOpcode() != SPIRV::OpTypeImage &&
1436 PointeeType->getOpcode() != SPIRV::OpTypeSampledImage &&
1437 PointeeType->getOpcode() != SPIRV::OpTypeSampler) {
1438 return;
1439 }
1440
1441 if (isUniformTexelBuffer(ImageInst: PointeeType)) {
1442 if (IsNonUniform)
1443 Handler.addRequirements(
1444 Req: SPIRV::Capability::UniformTexelBufferArrayNonUniformIndexingEXT);
1445 else if (!FirstIndexIsConstant)
1446 Handler.addRequirements(
1447 Req: SPIRV::Capability::UniformTexelBufferArrayDynamicIndexingEXT);
1448 } else if (isInputAttachment(ImageInst: PointeeType)) {
1449 if (IsNonUniform)
1450 Handler.addRequirements(
1451 Req: SPIRV::Capability::InputAttachmentArrayNonUniformIndexingEXT);
1452 else if (!FirstIndexIsConstant)
1453 Handler.addRequirements(
1454 Req: SPIRV::Capability::InputAttachmentArrayDynamicIndexingEXT);
1455 } else if (isStorageTexelBuffer(ImageInst: PointeeType)) {
1456 if (IsNonUniform)
1457 Handler.addRequirements(
1458 Req: SPIRV::Capability::StorageTexelBufferArrayNonUniformIndexingEXT);
1459 else if (!FirstIndexIsConstant)
1460 Handler.addRequirements(
1461 Req: SPIRV::Capability::StorageTexelBufferArrayDynamicIndexingEXT);
1462 } else if (isSampledImage(ImageInst: PointeeType) ||
1463 isCombinedImageSampler(SampledImageInst: PointeeType) ||
1464 PointeeType->getOpcode() == SPIRV::OpTypeSampler) {
1465 if (IsNonUniform)
1466 Handler.addRequirements(
1467 Req: SPIRV::Capability::SampledImageArrayNonUniformIndexingEXT);
1468 else if (!FirstIndexIsConstant)
1469 Handler.addRequirements(
1470 Req: SPIRV::Capability::SampledImageArrayDynamicIndexing);
1471 } else if (isStorageImage(ImageInst: PointeeType)) {
1472 if (IsNonUniform)
1473 Handler.addRequirements(
1474 Req: SPIRV::Capability::StorageImageArrayNonUniformIndexingEXT);
1475 else if (!FirstIndexIsConstant)
1476 Handler.addRequirements(
1477 Req: SPIRV::Capability::StorageImageArrayDynamicIndexing);
1478 }
1479}
1480
1481static bool isImageTypeWithUnknownFormat(SPIRVTypeInst TypeInst) {
1482 if (TypeInst->getOpcode() != SPIRV::OpTypeImage)
1483 return false;
1484 assert(TypeInst->getOperand(7).isImm() && "The image format must be an imm.");
1485 return TypeInst->getOperand(i: 7).getImm() == 0;
1486}
1487
1488static void AddDotProductRequirements(const MachineInstr &MI,
1489 SPIRV::RequirementHandler &Reqs,
1490 const SPIRVSubtarget &ST) {
1491 if (ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_integer_dot_product))
1492 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_KHR_integer_dot_product);
1493 Reqs.addCapability(ToAdd: SPIRV::Capability::DotProduct);
1494
1495 const MachineRegisterInfo &MRI = MI.getMF()->getRegInfo();
1496 assert(MI.getOperand(2).isReg() && "Unexpected operand in dot");
1497 // We do not consider what the previous instruction is. This is just used
1498 // to get the input register and to check the type.
1499 const MachineInstr *Input = MRI.getVRegDef(Reg: MI.getOperand(i: 2).getReg());
1500 assert(Input->getOperand(1).isReg() && "Unexpected operand in dot input");
1501 Register InputReg = Input->getOperand(i: 1).getReg();
1502
1503 SPIRVTypeInst TypeDef = MRI.getVRegDef(Reg: InputReg);
1504 if (TypeDef->getOpcode() == SPIRV::OpTypeInt) {
1505 assert(TypeDef->getOperand(1).getImm() == 32);
1506 Reqs.addCapability(ToAdd: SPIRV::Capability::DotProductInput4x8BitPacked);
1507 } else if (isVectorType(SPVTy: TypeDef)) {
1508 SPIRVTypeInst ScalarTypeDef =
1509 MRI.getVRegDef(Reg: TypeDef->getOperand(i: 1).getReg());
1510 assert(ScalarTypeDef->getOpcode() == SPIRV::OpTypeInt);
1511 if (ScalarTypeDef->getOperand(i: 1).getImm() == 8) {
1512 assert(TypeDef->getOperand(2).getImm() == 4 &&
1513 "Dot operand of 8-bit integer type requires 4 components");
1514 Reqs.addCapability(ToAdd: SPIRV::Capability::DotProductInput4x8Bit);
1515 } else {
1516 Reqs.addCapability(ToAdd: SPIRV::Capability::DotProductInputAll);
1517 }
1518 }
1519}
1520
1521void addPrintfRequirements(const MachineInstr &MI,
1522 SPIRV::RequirementHandler &Reqs,
1523 const SPIRVSubtarget &ST) {
1524 SPIRVGlobalRegistry *GR = ST.getSPIRVGlobalRegistry();
1525 SPIRVTypeInst PtrType =
1526 GR->getSPIRVTypeForVReg(VReg: MI.getOperand(i: 4).getReg(), MF: MI.getMF());
1527 if (PtrType) {
1528 MachineOperand ASOp = PtrType->getOperand(i: 1);
1529 if (ASOp.isImm()) {
1530 unsigned AddrSpace = ASOp.getImm();
1531 if (AddrSpace != SPIRV::StorageClass::UniformConstant) {
1532 if (!ST.canUseExtension(
1533 E: SPIRV::Extension::
1534 SPV_EXT_relaxed_printf_string_address_space)) {
1535 report_fatal_error(reason: "SPV_EXT_relaxed_printf_string_address_space is "
1536 "required because printf uses a format string not "
1537 "in constant address space.",
1538 gen_crash_diag: false);
1539 }
1540 Reqs.addExtension(
1541 ToAdd: SPIRV::Extension::SPV_EXT_relaxed_printf_string_address_space);
1542 }
1543 }
1544 }
1545}
1546
1547static void addImageOperandReqs(const MachineInstr &MI,
1548 SPIRV::RequirementHandler &Reqs,
1549 const SPIRVSubtarget &ST, unsigned OpIdx) {
1550 if (MI.getNumOperands() <= OpIdx)
1551 return;
1552 uint32_t Mask = MI.getOperand(i: OpIdx).getImm();
1553 for (uint32_t I = 0; I < 32; ++I)
1554 if (Mask & (1U << I))
1555 Reqs.getAndAddRequirements(Category: SPIRV::OperandCategory::ImageOperandOperand,
1556 i: 1U << I, ST);
1557}
1558
1559static inline void maybeAddScatterGatherReq(const MachineInstr &MI,
1560 SPIRV::RequirementHandler &Reqs,
1561 const SPIRVSubtarget &ST) {
1562 assert(MI.getOperand(1).isReg());
1563 const MachineRegisterInfo &MRI = MI.getMF()->getRegInfo();
1564 SPIRVTypeInst ElemTypeDef = MRI.getVRegDef(Reg: MI.getOperand(i: 1).getReg());
1565 if (ElemTypeDef->getOpcode() == SPIRV::OpTypePointer &&
1566 ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_masked_gather_scatter)) {
1567 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_masked_gather_scatter);
1568 Reqs.addCapability(ToAdd: SPIRV::Capability::MaskedGatherScatterINTEL);
1569 }
1570}
1571
1572void addInstrRequirements(const MachineInstr &MI,
1573 SPIRV::ModuleAnalysisInfo &MAI,
1574 const SPIRVSubtarget &ST) {
1575 SPIRV::RequirementHandler &Reqs = MAI.Reqs;
1576 unsigned Op = MI.getOpcode();
1577 switch (Op) {
1578 case SPIRV::OpMemoryModel: {
1579 int64_t Addr = MI.getOperand(i: 0).getImm();
1580 Reqs.getAndAddRequirements(Category: SPIRV::OperandCategory::AddressingModelOperand,
1581 i: Addr, ST);
1582 int64_t Mem = MI.getOperand(i: 1).getImm();
1583 Reqs.getAndAddRequirements(Category: SPIRV::OperandCategory::MemoryModelOperand, i: Mem,
1584 ST);
1585 break;
1586 }
1587 case SPIRV::OpEntryPoint: {
1588 int64_t Exe = MI.getOperand(i: 0).getImm();
1589 Reqs.getAndAddRequirements(Category: SPIRV::OperandCategory::ExecutionModelOperand,
1590 i: Exe, ST);
1591 break;
1592 }
1593 case SPIRV::OpExecutionMode:
1594 case SPIRV::OpExecutionModeId: {
1595 int64_t Exe = MI.getOperand(i: 1).getImm();
1596 Reqs.getAndAddRequirements(Category: SPIRV::OperandCategory::ExecutionModeOperand,
1597 i: Exe, ST);
1598 break;
1599 }
1600 case SPIRV::OpTypeMatrix:
1601 Reqs.addCapability(ToAdd: SPIRV::Capability::Matrix);
1602 break;
1603 case SPIRV::OpTypeInt: {
1604 unsigned BitWidth = MI.getOperand(i: 1).getImm();
1605 if (BitWidth == 64)
1606 Reqs.addCapability(ToAdd: SPIRV::Capability::Int64);
1607 else if (BitWidth == 16)
1608 Reqs.addCapability(ToAdd: SPIRV::Capability::Int16);
1609 else if (BitWidth == 8)
1610 Reqs.addCapability(ToAdd: SPIRV::Capability::Int8);
1611 else if (BitWidth == 4 &&
1612 ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_int4)) {
1613 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_int4);
1614 Reqs.addCapability(ToAdd: SPIRV::Capability::Int4TypeINTEL);
1615 } else if (BitWidth != 32) {
1616 if (!ST.canUseExtension(
1617 E: SPIRV::Extension::SPV_ALTERA_arbitrary_precision_integers))
1618 reportFatalUsageError(
1619 reason: "OpTypeInt type with a width other than 8, 16, 32 or 64 bits "
1620 "requires the following SPIR-V extension: "
1621 "SPV_ALTERA_arbitrary_precision_integers");
1622 Reqs.addExtension(
1623 ToAdd: SPIRV::Extension::SPV_ALTERA_arbitrary_precision_integers);
1624 Reqs.addCapability(ToAdd: SPIRV::Capability::ArbitraryPrecisionIntegersALTERA);
1625 }
1626 break;
1627 }
1628 case SPIRV::OpDot: {
1629 const MachineRegisterInfo &MRI = MI.getMF()->getRegInfo();
1630 SPIRVTypeInst TypeDef = MRI.getVRegDef(Reg: MI.getOperand(i: 1).getReg());
1631 if (isBFloat16Type(TypeDef))
1632 Reqs.addCapability(ToAdd: SPIRV::Capability::BFloat16DotProductKHR);
1633 break;
1634 }
1635 case SPIRV::OpTypeFloat: {
1636 unsigned BitWidth = MI.getOperand(i: 1).getImm();
1637 if (BitWidth == 64)
1638 Reqs.addCapability(ToAdd: SPIRV::Capability::Float64);
1639 else if (BitWidth == 16) {
1640 if (isBFloat16Type(TypeDef: &MI)) {
1641 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_bfloat16))
1642 report_fatal_error(reason: "OpTypeFloat type with bfloat requires the "
1643 "following SPIR-V extension: SPV_KHR_bfloat16",
1644 gen_crash_diag: false);
1645 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_KHR_bfloat16);
1646 Reqs.addCapability(ToAdd: SPIRV::Capability::BFloat16TypeKHR);
1647 } else {
1648 Reqs.addCapability(ToAdd: SPIRV::Capability::Float16);
1649 }
1650 }
1651 break;
1652 }
1653 case SPIRV::OpTypeVector: {
1654 unsigned NumComponents = MI.getOperand(i: 2).getImm();
1655 if (NumComponents == 8 || NumComponents == 16)
1656 Reqs.addCapability(ToAdd: SPIRV::Capability::Vector16);
1657 else if (requiresLongVectorEXT(NumComponents))
1658 // Such widths are only expressible as OpTypeVectorIdEXT.
1659 reportFatalUsageError(
1660 reason: "OpTypeVector with " + Twine(NumComponents) +
1661 " components requires the following SPIR-V extension: "
1662 "SPV_EXT_long_vector");
1663
1664 maybeAddScatterGatherReq(MI, Reqs, ST);
1665 break;
1666 }
1667 case SPIRV::OpTypeVectorIdEXT: {
1668 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_EXT_long_vector))
1669 reportFatalUsageError(reason: "OpTypeVectorIdEXT requires the following SPIR-V "
1670 "extension: SPV_EXT_long_vector extension");
1671 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_EXT_long_vector);
1672 Reqs.addCapability(ToAdd: SPIRV::Capability::LongVectorEXT);
1673 maybeAddScatterGatherReq(MI, Reqs, ST);
1674 break;
1675 }
1676 case SPIRV::OpTypePointer: {
1677 auto SC = MI.getOperand(i: 1).getImm();
1678 Reqs.getAndAddRequirements(Category: SPIRV::OperandCategory::StorageClassOperand, i: SC,
1679 ST);
1680 // If it's a type of pointer to float16 targeting OpenCL, add Float16Buffer
1681 // capability.
1682 if (ST.isShader())
1683 break;
1684 assert(MI.getOperand(2).isReg());
1685 const MachineRegisterInfo &MRI = MI.getMF()->getRegInfo();
1686 SPIRVTypeInst TypeDef = MRI.getVRegDef(Reg: MI.getOperand(i: 2).getReg());
1687 if ((TypeDef->getNumOperands() == 2) &&
1688 (TypeDef->getOpcode() == SPIRV::OpTypeFloat) &&
1689 (TypeDef->getOperand(i: 1).getImm() == 16))
1690 Reqs.addCapability(ToAdd: SPIRV::Capability::Float16Buffer);
1691 break;
1692 }
1693 case SPIRV::OpExtInst: {
1694 if (MI.getOperand(i: 2).getImm() ==
1695 static_cast<int64_t>(
1696 SPIRV::InstructionSet::NonSemantic_Shader_DebugInfo_100)) {
1697 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_KHR_non_semantic_info);
1698 break;
1699 }
1700 if (MI.getOperand(i: 3).getImm() ==
1701 static_cast<int64_t>(SPIRV::OpenCLExtInst::printf)) {
1702 addPrintfRequirements(MI, Reqs, ST);
1703 break;
1704 }
1705 if (MI.getOperand(i: 2).getImm() ==
1706 static_cast<int64_t>(SPIRV::InstructionSet::OpenCL_std)) {
1707 const MachineFunction *MF = MI.getMF();
1708 const MachineRegisterInfo &MRI = MF->getRegInfo();
1709 SPIRVGlobalRegistry *GR = ST.getSPIRVGlobalRegistry();
1710
1711 auto IsBFloat16 = [&](SPIRVTypeInst TypeDef) {
1712 if (TypeDef && TypeDef->getOpcode() == SPIRV::OpTypeVector)
1713 TypeDef = MRI.getVRegDef(Reg: TypeDef->getOperand(i: 1).getReg());
1714 return isBFloat16Type(TypeDef);
1715 };
1716
1717 // Result type is operand 1; arguments start at operand 4.
1718 bool UsesBFloat16 = IsBFloat16(MRI.getVRegDef(Reg: MI.getOperand(i: 1).getReg()));
1719 for (unsigned I = 4, E = MI.getNumOperands(); I < E && !UsesBFloat16;
1720 ++I) {
1721 const MachineOperand &MO = MI.getOperand(i: I);
1722 if (MO.isReg())
1723 UsesBFloat16 = IsBFloat16(GR->getResultType(
1724 VReg: MO.getReg(), MF: const_cast<MachineFunction *>(MF)));
1725 }
1726
1727 if (UsesBFloat16) {
1728 if (!ST.canUseExtension(
1729 E: SPIRV::Extension::SPV_INTEL_bfloat16_arithmetic)) {
1730 reportUnsupported(
1731 MI, Msg: "OpenCL Extended instructions with bfloat16 require the "
1732 "following SPIR-V extension: SPV_INTEL_bfloat16_arithmetic");
1733 break;
1734 }
1735 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_bfloat16_arithmetic);
1736 Reqs.addCapability(ToAdd: SPIRV::Capability::BFloat16ArithmeticINTEL);
1737 }
1738 }
1739 break;
1740 }
1741 case SPIRV::OpAliasDomainDeclINTEL:
1742 case SPIRV::OpAliasScopeDeclINTEL:
1743 case SPIRV::OpAliasScopeListDeclINTEL: {
1744 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_memory_access_aliasing);
1745 Reqs.addCapability(ToAdd: SPIRV::Capability::MemoryAccessAliasingINTEL);
1746 break;
1747 }
1748 case SPIRV::OpBitReverse:
1749 case SPIRV::OpBitFieldInsert:
1750 case SPIRV::OpBitFieldSExtract:
1751 case SPIRV::OpBitFieldUExtract:
1752 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_bit_instructions)) {
1753 Reqs.addCapability(ToAdd: SPIRV::Capability::Shader);
1754 break;
1755 }
1756 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_KHR_bit_instructions);
1757 Reqs.addCapability(ToAdd: SPIRV::Capability::BitInstructions);
1758 break;
1759 case SPIRV::OpTypeRuntimeArray:
1760 Reqs.addCapability(ToAdd: SPIRV::Capability::Shader);
1761 break;
1762 case SPIRV::OpTypeOpaque:
1763 case SPIRV::OpTypeEvent:
1764 Reqs.addCapability(ToAdd: SPIRV::Capability::Kernel);
1765 break;
1766 case SPIRV::OpTypePipe:
1767 case SPIRV::OpTypeReserveId:
1768 Reqs.addCapability(ToAdd: SPIRV::Capability::Pipes);
1769 break;
1770 case SPIRV::OpTypeDeviceEvent:
1771 case SPIRV::OpTypeQueue:
1772 case SPIRV::OpBuildNDRange:
1773 case SPIRV::OpEnqueueKernel:
1774 case SPIRV::OpGetKernelNDrangeSubGroupCount:
1775 case SPIRV::OpGetKernelNDrangeMaxSubGroupSize:
1776 case SPIRV::OpGetKernelWorkGroupSize:
1777 case SPIRV::OpGetKernelPreferredWorkGroupSizeMultiple:
1778 Reqs.addCapability(ToAdd: SPIRV::Capability::DeviceEnqueue);
1779 break;
1780 case SPIRV::OpDecorate:
1781 case SPIRV::OpDecorateId:
1782 case SPIRV::OpDecorateString:
1783 addOpDecorateReqs(MI, DecIndex: 1, Reqs, ST);
1784 break;
1785 case SPIRV::OpMemberDecorate:
1786 case SPIRV::OpMemberDecorateString:
1787 addOpDecorateReqs(MI, DecIndex: 2, Reqs, ST);
1788 break;
1789 case SPIRV::OpInBoundsPtrAccessChain:
1790 Reqs.addCapability(ToAdd: SPIRV::Capability::Addresses);
1791 break;
1792 case SPIRV::OpConstantSampler:
1793 Reqs.addCapability(ToAdd: SPIRV::Capability::LiteralSampler);
1794 break;
1795 case SPIRV::OpInBoundsAccessChain:
1796 case SPIRV::OpAccessChain:
1797 addOpAccessChainReqs(Instr: MI, Handler&: Reqs, Subtarget: ST);
1798 break;
1799 case SPIRV::OpTypeImage:
1800 addOpTypeImageReqs(MI, Reqs, ST);
1801 break;
1802 case SPIRV::OpTypeSampler:
1803 if (!ST.isShader()) {
1804 Reqs.addCapability(ToAdd: SPIRV::Capability::ImageBasic);
1805 }
1806 break;
1807 case SPIRV::OpTypeForwardPointer:
1808 // TODO: check if it's OpenCL's kernel.
1809 Reqs.addCapability(ToAdd: SPIRV::Capability::Addresses);
1810 break;
1811 case SPIRV::OpAtomicFlagTestAndSet:
1812 case SPIRV::OpAtomicLoad:
1813 case SPIRV::OpAtomicStore:
1814 case SPIRV::OpAtomicExchange:
1815 case SPIRV::OpAtomicCompareExchange:
1816 case SPIRV::OpAtomicCompareExchangeWeak:
1817 case SPIRV::OpAtomicIIncrement:
1818 case SPIRV::OpAtomicIDecrement:
1819 case SPIRV::OpAtomicIAdd:
1820 case SPIRV::OpAtomicISub:
1821 case SPIRV::OpAtomicUMin:
1822 case SPIRV::OpAtomicUMax:
1823 case SPIRV::OpAtomicSMin:
1824 case SPIRV::OpAtomicSMax:
1825 case SPIRV::OpAtomicAnd:
1826 case SPIRV::OpAtomicOr:
1827 case SPIRV::OpAtomicXor: {
1828 const MachineRegisterInfo &MRI = MI.getMF()->getRegInfo();
1829 const MachineInstr *InstrPtr = &MI;
1830 if (Op == SPIRV::OpAtomicStore) {
1831 assert(MI.getOperand(3).isReg());
1832 InstrPtr = MRI.getVRegDef(Reg: MI.getOperand(i: 3).getReg());
1833 assert(InstrPtr && "Unexpected type instruction for OpAtomicStore");
1834 }
1835 assert(InstrPtr->getOperand(1).isReg() && "Unexpected operand in atomic");
1836 Register TypeReg = InstrPtr->getOperand(i: 1).getReg();
1837 SPIRVTypeInst TypeDef = MRI.getVRegDef(Reg: TypeReg);
1838
1839 if (TypeDef->getOpcode() == SPIRV::OpTypeInt) {
1840 unsigned BitWidth = TypeDef->getOperand(i: 1).getImm();
1841 if (BitWidth == 64)
1842 Reqs.addCapability(ToAdd: SPIRV::Capability::Int64Atomics);
1843 else if (BitWidth == 16) {
1844 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_16bit_atomics))
1845 report_fatal_error(
1846 reason: "16-bit integer atomic operations require the following SPIR-V "
1847 "extension: SPV_INTEL_16bit_atomics",
1848 gen_crash_diag: false);
1849 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_16bit_atomics);
1850 switch (Op) {
1851 case SPIRV::OpAtomicLoad:
1852 case SPIRV::OpAtomicStore:
1853 case SPIRV::OpAtomicExchange:
1854 case SPIRV::OpAtomicCompareExchange:
1855 case SPIRV::OpAtomicCompareExchangeWeak:
1856 Reqs.addCapability(
1857 ToAdd: SPIRV::Capability::AtomicInt16CompareExchangeINTEL);
1858 break;
1859 default:
1860 Reqs.addCapability(ToAdd: SPIRV::Capability::Int16AtomicsINTEL);
1861 break;
1862 }
1863 }
1864 } else if (isBFloat16Type(TypeDef)) {
1865 if (is_contained(Set: {SPIRV::OpAtomicLoad, SPIRV::OpAtomicStore,
1866 SPIRV::OpAtomicExchange},
1867 Element: Op)) {
1868 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_16bit_atomics))
1869 report_fatal_error(
1870 reason: "The atomic bfloat16 instruction requires the following SPIR-V "
1871 "extension: SPV_INTEL_16bit_atomics",
1872 gen_crash_diag: false);
1873 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_16bit_atomics);
1874 Reqs.addCapability(ToAdd: SPIRV::Capability::AtomicBFloat16LoadStoreINTEL);
1875 }
1876 }
1877 break;
1878 }
1879 case SPIRV::OpGroupNonUniformIAdd:
1880 case SPIRV::OpGroupNonUniformFAdd:
1881 case SPIRV::OpGroupNonUniformIMul:
1882 case SPIRV::OpGroupNonUniformFMul:
1883 case SPIRV::OpGroupNonUniformSMin:
1884 case SPIRV::OpGroupNonUniformUMin:
1885 case SPIRV::OpGroupNonUniformFMin:
1886 case SPIRV::OpGroupNonUniformSMax:
1887 case SPIRV::OpGroupNonUniformUMax:
1888 case SPIRV::OpGroupNonUniformFMax:
1889 case SPIRV::OpGroupNonUniformBitwiseAnd:
1890 case SPIRV::OpGroupNonUniformBitwiseOr:
1891 case SPIRV::OpGroupNonUniformBitwiseXor:
1892 case SPIRV::OpGroupNonUniformLogicalAnd:
1893 case SPIRV::OpGroupNonUniformLogicalOr:
1894 case SPIRV::OpGroupNonUniformLogicalXor: {
1895 assert(MI.getOperand(3).isImm());
1896 int64_t GroupOp = MI.getOperand(i: 3).getImm();
1897 switch (GroupOp) {
1898 case SPIRV::GroupOperation::Reduce:
1899 case SPIRV::GroupOperation::InclusiveScan:
1900 case SPIRV::GroupOperation::ExclusiveScan:
1901 Reqs.addCapability(ToAdd: SPIRV::Capability::GroupNonUniformArithmetic);
1902 break;
1903 case SPIRV::GroupOperation::ClusteredReduce:
1904 Reqs.addCapability(ToAdd: SPIRV::Capability::GroupNonUniformClustered);
1905 break;
1906 case SPIRV::GroupOperation::PartitionedReduceNV:
1907 case SPIRV::GroupOperation::PartitionedInclusiveScanNV:
1908 case SPIRV::GroupOperation::PartitionedExclusiveScanNV:
1909 Reqs.addCapability(ToAdd: SPIRV::Capability::GroupNonUniformPartitionedNV);
1910 break;
1911 }
1912 break;
1913 }
1914 case SPIRV::OpGroupNonUniformQuadSwap:
1915 Reqs.addCapability(ToAdd: SPIRV::Capability::GroupNonUniformQuad);
1916 break;
1917 case SPIRV::OpImageQueryLod:
1918 Reqs.addCapability(ToAdd: SPIRV::Capability::ImageQuery);
1919 break;
1920 case SPIRV::OpImageQuerySize:
1921 case SPIRV::OpImageQuerySizeLod:
1922 case SPIRV::OpImageQueryLevels:
1923 case SPIRV::OpImageQuerySamples:
1924 if (ST.isShader())
1925 Reqs.addCapability(ToAdd: SPIRV::Capability::ImageQuery);
1926 break;
1927 case SPIRV::OpImageQueryFormat: {
1928 Register ResultReg = MI.getOperand(i: 0).getReg();
1929 const MachineRegisterInfo &MRI = MI.getMF()->getRegInfo();
1930 static const unsigned CompareOps[] = {
1931 SPIRV::OpIEqual, SPIRV::OpINotEqual,
1932 SPIRV::OpUGreaterThan, SPIRV::OpUGreaterThanEqual,
1933 SPIRV::OpULessThan, SPIRV::OpULessThanEqual,
1934 SPIRV::OpSGreaterThan, SPIRV::OpSGreaterThanEqual,
1935 SPIRV::OpSLessThan, SPIRV::OpSLessThanEqual};
1936
1937 auto CheckAndAddExtension = [&](int64_t ImmVal) {
1938 if (ImmVal == 4323 || ImmVal == 4324) {
1939 if (ST.canUseExtension(E: SPIRV::Extension::SPV_EXT_image_raw10_raw12))
1940 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_EXT_image_raw10_raw12);
1941 else
1942 report_fatal_error(reason: "This requires the "
1943 "SPV_EXT_image_raw10_raw12 extension");
1944 }
1945 };
1946
1947 for (MachineInstr &UseInst : MRI.use_instructions(Reg: ResultReg)) {
1948 unsigned Opc = UseInst.getOpcode();
1949
1950 if (Opc == SPIRV::OpSwitch) {
1951 for (const MachineOperand &Op : UseInst.operands())
1952 if (Op.isImm())
1953 CheckAndAddExtension(Op.getImm());
1954 } else if (llvm::is_contained(Range: CompareOps, Element: Opc)) {
1955 for (unsigned i = 1; i < UseInst.getNumOperands(); ++i) {
1956 Register UseReg = UseInst.getOperand(i).getReg();
1957 MachineInstr *ConstInst = MRI.getVRegDef(Reg: UseReg);
1958 if (ConstInst && ConstInst->getOpcode() == SPIRV::OpConstantI) {
1959 int64_t ImmVal = ConstInst->getOperand(i: 2).getImm();
1960 if (ImmVal)
1961 CheckAndAddExtension(ImmVal);
1962 }
1963 }
1964 }
1965 }
1966 break;
1967 }
1968
1969 case SPIRV::OpGroupNonUniformShuffle:
1970 case SPIRV::OpGroupNonUniformShuffleXor:
1971 Reqs.addCapability(ToAdd: SPIRV::Capability::GroupNonUniformShuffle);
1972 break;
1973 case SPIRV::OpGroupNonUniformShuffleUp:
1974 case SPIRV::OpGroupNonUniformShuffleDown:
1975 Reqs.addCapability(ToAdd: SPIRV::Capability::GroupNonUniformShuffleRelative);
1976 break;
1977 case SPIRV::OpGroupAll:
1978 case SPIRV::OpGroupAny:
1979 case SPIRV::OpGroupBroadcast:
1980 case SPIRV::OpGroupIAdd:
1981 case SPIRV::OpGroupFAdd:
1982 case SPIRV::OpGroupFMin:
1983 case SPIRV::OpGroupUMin:
1984 case SPIRV::OpGroupSMin:
1985 case SPIRV::OpGroupFMax:
1986 case SPIRV::OpGroupUMax:
1987 case SPIRV::OpGroupSMax:
1988 Reqs.addCapability(ToAdd: SPIRV::Capability::Groups);
1989 break;
1990 case SPIRV::OpGroupNonUniformElect:
1991 Reqs.addCapability(ToAdd: SPIRV::Capability::GroupNonUniform);
1992 break;
1993 case SPIRV::OpGroupNonUniformAll:
1994 case SPIRV::OpGroupNonUniformAny:
1995 case SPIRV::OpGroupNonUniformAllEqual:
1996 Reqs.addCapability(ToAdd: SPIRV::Capability::GroupNonUniformVote);
1997 break;
1998 case SPIRV::OpGroupNonUniformBroadcast:
1999 case SPIRV::OpGroupNonUniformBroadcastFirst:
2000 case SPIRV::OpGroupNonUniformBallot:
2001 case SPIRV::OpGroupNonUniformInverseBallot:
2002 case SPIRV::OpGroupNonUniformBallotBitExtract:
2003 case SPIRV::OpGroupNonUniformBallotBitCount:
2004 case SPIRV::OpGroupNonUniformBallotFindLSB:
2005 case SPIRV::OpGroupNonUniformBallotFindMSB:
2006 Reqs.addCapability(ToAdd: SPIRV::Capability::GroupNonUniformBallot);
2007 break;
2008 case SPIRV::OpSubgroupShuffleINTEL:
2009 case SPIRV::OpSubgroupShuffleDownINTEL:
2010 case SPIRV::OpSubgroupShuffleUpINTEL:
2011 case SPIRV::OpSubgroupShuffleXorINTEL:
2012 if (ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_subgroups)) {
2013 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_subgroups);
2014 Reqs.addCapability(ToAdd: SPIRV::Capability::SubgroupShuffleINTEL);
2015 }
2016 break;
2017 case SPIRV::OpSubgroupBlockReadINTEL:
2018 case SPIRV::OpSubgroupBlockWriteINTEL:
2019 if (ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_subgroups)) {
2020 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_subgroups);
2021 Reqs.addCapability(ToAdd: SPIRV::Capability::SubgroupBufferBlockIOINTEL);
2022 }
2023 break;
2024 case SPIRV::OpSubgroupImageBlockReadINTEL:
2025 case SPIRV::OpSubgroupImageBlockWriteINTEL:
2026 if (ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_subgroups)) {
2027 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_subgroups);
2028 Reqs.addCapability(ToAdd: SPIRV::Capability::SubgroupImageBlockIOINTEL);
2029 }
2030 break;
2031 case SPIRV::OpSubgroupImageMediaBlockReadINTEL:
2032 case SPIRV::OpSubgroupImageMediaBlockWriteINTEL:
2033 if (ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_media_block_io)) {
2034 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_media_block_io);
2035 Reqs.addCapability(ToAdd: SPIRV::Capability::SubgroupImageMediaBlockIOINTEL);
2036 }
2037 break;
2038 case SPIRV::OpAssumeTrueKHR:
2039 case SPIRV::OpExpectKHR:
2040 if (ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_expect_assume)) {
2041 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_KHR_expect_assume);
2042 Reqs.addCapability(ToAdd: SPIRV::Capability::ExpectAssumeKHR);
2043 }
2044 break;
2045 case SPIRV::OpFmaKHR:
2046 if (ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_fma)) {
2047 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_KHR_fma);
2048 Reqs.addCapability(ToAdd: SPIRV::Capability::FmaKHR);
2049 }
2050 break;
2051 case SPIRV::OpPtrCastToCrossWorkgroupINTEL:
2052 case SPIRV::OpCrossWorkgroupCastToPtrINTEL:
2053 if (ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_usm_storage_classes)) {
2054 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_usm_storage_classes);
2055 Reqs.addCapability(ToAdd: SPIRV::Capability::USMStorageClassesINTEL);
2056 }
2057 break;
2058 case SPIRV::OpConstantFunctionPointerINTEL:
2059 case SPIRV::OpFunctionPointerCallINTEL:
2060 if (ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_function_pointers)) {
2061 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_function_pointers);
2062 Reqs.addCapability(ToAdd: SPIRV::Capability::FunctionPointersINTEL);
2063 }
2064 break;
2065 case SPIRV::OpGroupNonUniformRotateKHR:
2066 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_subgroup_rotate))
2067 report_fatal_error(reason: "OpGroupNonUniformRotateKHR instruction requires the "
2068 "following SPIR-V extension: SPV_KHR_subgroup_rotate",
2069 gen_crash_diag: false);
2070 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_KHR_subgroup_rotate);
2071 Reqs.addCapability(ToAdd: SPIRV::Capability::GroupNonUniformRotateKHR);
2072 Reqs.addCapability(ToAdd: SPIRV::Capability::GroupNonUniform);
2073 break;
2074 case SPIRV::OpFixedCosALTERA:
2075 case SPIRV::OpFixedSinALTERA:
2076 case SPIRV::OpFixedCosPiALTERA:
2077 case SPIRV::OpFixedSinPiALTERA:
2078 case SPIRV::OpFixedExpALTERA:
2079 case SPIRV::OpFixedLogALTERA:
2080 case SPIRV::OpFixedRecipALTERA:
2081 case SPIRV::OpFixedSqrtALTERA:
2082 case SPIRV::OpFixedSinCosALTERA:
2083 case SPIRV::OpFixedSinCosPiALTERA:
2084 case SPIRV::OpFixedRsqrtALTERA:
2085 if (!ST.canUseExtension(
2086 E: SPIRV::Extension::SPV_ALTERA_arbitrary_precision_fixed_point))
2087 report_fatal_error(reason: "This instruction requires the "
2088 "following SPIR-V extension: "
2089 "SPV_ALTERA_arbitrary_precision_fixed_point",
2090 gen_crash_diag: false);
2091 Reqs.addExtension(
2092 ToAdd: SPIRV::Extension::SPV_ALTERA_arbitrary_precision_fixed_point);
2093 Reqs.addCapability(ToAdd: SPIRV::Capability::ArbitraryPrecisionFixedPointALTERA);
2094 break;
2095 case SPIRV::OpGroupIMulKHR:
2096 case SPIRV::OpGroupFMulKHR:
2097 case SPIRV::OpGroupBitwiseAndKHR:
2098 case SPIRV::OpGroupBitwiseOrKHR:
2099 case SPIRV::OpGroupBitwiseXorKHR:
2100 case SPIRV::OpGroupLogicalAndKHR:
2101 case SPIRV::OpGroupLogicalOrKHR:
2102 case SPIRV::OpGroupLogicalXorKHR:
2103 if (ST.canUseExtension(
2104 E: SPIRV::Extension::SPV_KHR_uniform_group_instructions)) {
2105 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_KHR_uniform_group_instructions);
2106 Reqs.addCapability(ToAdd: SPIRV::Capability::GroupUniformArithmeticKHR);
2107 }
2108 break;
2109 case SPIRV::OpReadClockKHR:
2110 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_shader_clock))
2111 report_fatal_error(reason: "OpReadClockKHR instruction requires the "
2112 "following SPIR-V extension: SPV_KHR_shader_clock",
2113 gen_crash_diag: false);
2114 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_KHR_shader_clock);
2115 Reqs.addCapability(ToAdd: SPIRV::Capability::ShaderClockKHR);
2116 break;
2117 case SPIRV::OpAbortKHR:
2118 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_abort))
2119 report_fatal_error(reason: "OpAbortKHR instruction requires the "
2120 "following SPIR-V extension: SPV_KHR_abort",
2121 gen_crash_diag: false);
2122 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_KHR_abort);
2123 Reqs.addCapability(ToAdd: SPIRV::Capability::AbortKHR);
2124 break;
2125 case SPIRV::OpPoisonKHR:
2126 case SPIRV::OpFreezeKHR:
2127 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_poison_freeze))
2128 report_fatal_error(reason: "OpPoisonKHR/OpFreezeKHR instruction requires the "
2129 "following SPIR-V extension: SPV_KHR_poison_freeze",
2130 gen_crash_diag: false);
2131 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_KHR_poison_freeze);
2132 Reqs.addCapability(ToAdd: SPIRV::Capability::PoisonFreezeKHR);
2133 break;
2134 case SPIRV::OpAtomicFAddEXT:
2135 case SPIRV::OpAtomicFMinEXT:
2136 case SPIRV::OpAtomicFMaxEXT:
2137 AddAtomicFloatRequirements(MI, Reqs, ST);
2138 break;
2139 case SPIRV::OpConvertBF16ToFINTEL:
2140 case SPIRV::OpConvertFToBF16INTEL:
2141 if (ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_bfloat16_conversion)) {
2142 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_bfloat16_conversion);
2143 Reqs.addCapability(ToAdd: SPIRV::Capability::BFloat16ConversionINTEL);
2144 }
2145 break;
2146 case SPIRV::OpRoundFToTF32INTEL:
2147 if (ST.canUseExtension(
2148 E: SPIRV::Extension::SPV_INTEL_tensor_float32_conversion)) {
2149 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_tensor_float32_conversion);
2150 Reqs.addCapability(ToAdd: SPIRV::Capability::TensorFloat32RoundingINTEL);
2151 }
2152 break;
2153 case SPIRV::OpVariableLengthArrayINTEL:
2154 case SPIRV::OpSaveMemoryINTEL:
2155 case SPIRV::OpRestoreMemoryINTEL:
2156 if (ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_variable_length_array)) {
2157 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_variable_length_array);
2158 Reqs.addCapability(ToAdd: SPIRV::Capability::VariableLengthArrayINTEL);
2159 }
2160 break;
2161 case SPIRV::OpUntypedVariableLengthArrayINTEL:
2162 if (ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_variable_length_array)) {
2163 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_variable_length_array);
2164 Reqs.addCapability(ToAdd: SPIRV::Capability::UntypedVariableLengthArrayINTEL);
2165 }
2166 break;
2167 case SPIRV::OpAsmTargetINTEL:
2168 case SPIRV::OpAsmINTEL:
2169 case SPIRV::OpAsmCallINTEL:
2170 if (ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_inline_assembly)) {
2171 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_inline_assembly);
2172 Reqs.addCapability(ToAdd: SPIRV::Capability::AsmINTEL);
2173 }
2174 break;
2175 case SPIRV::OpTypeCooperativeMatrixKHR: {
2176 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_cooperative_matrix))
2177 report_fatal_error(
2178 reason: "OpTypeCooperativeMatrixKHR type requires the "
2179 "following SPIR-V extension: SPV_KHR_cooperative_matrix",
2180 gen_crash_diag: false);
2181 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_KHR_cooperative_matrix);
2182 Reqs.addCapability(ToAdd: SPIRV::Capability::CooperativeMatrixKHR);
2183 const MachineRegisterInfo &MRI = MI.getMF()->getRegInfo();
2184 SPIRVTypeInst TypeDef = MRI.getVRegDef(Reg: MI.getOperand(i: 1).getReg());
2185 if (isBFloat16Type(TypeDef))
2186 Reqs.addCapability(ToAdd: SPIRV::Capability::BFloat16CooperativeMatrixKHR);
2187 break;
2188 }
2189 case SPIRV::OpArithmeticFenceEXT:
2190 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_EXT_arithmetic_fence))
2191 report_fatal_error(reason: "OpArithmeticFenceEXT requires the "
2192 "following SPIR-V extension: SPV_EXT_arithmetic_fence",
2193 gen_crash_diag: false);
2194 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_EXT_arithmetic_fence);
2195 Reqs.addCapability(ToAdd: SPIRV::Capability::ArithmeticFenceEXT);
2196 break;
2197 case SPIRV::OpControlBarrierArriveINTEL:
2198 case SPIRV::OpControlBarrierWaitINTEL:
2199 if (ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_split_barrier)) {
2200 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_split_barrier);
2201 Reqs.addCapability(ToAdd: SPIRV::Capability::SplitBarrierINTEL);
2202 }
2203 break;
2204 case SPIRV::OpCooperativeMatrixMulAddKHR: {
2205 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_cooperative_matrix))
2206 report_fatal_error(reason: "Cooperative matrix instructions require the "
2207 "following SPIR-V extension: "
2208 "SPV_KHR_cooperative_matrix",
2209 gen_crash_diag: false);
2210 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_KHR_cooperative_matrix);
2211 Reqs.addCapability(ToAdd: SPIRV::Capability::CooperativeMatrixKHR);
2212 constexpr unsigned MulAddMaxSize = 6;
2213 if (MI.getNumOperands() != MulAddMaxSize)
2214 break;
2215 const int64_t CoopOperands = MI.getOperand(i: MulAddMaxSize - 1).getImm();
2216 if (CoopOperands &
2217 SPIRV::CooperativeMatrixOperands::MatrixAAndBTF32ComponentsINTEL) {
2218 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_joint_matrix))
2219 report_fatal_error(reason: "MatrixAAndBTF32ComponentsINTEL type interpretation "
2220 "require the following SPIR-V extension: "
2221 "SPV_INTEL_joint_matrix",
2222 gen_crash_diag: false);
2223 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_joint_matrix);
2224 Reqs.addCapability(
2225 ToAdd: SPIRV::Capability::CooperativeMatrixTF32ComponentTypeINTEL);
2226 }
2227 if (CoopOperands & SPIRV::CooperativeMatrixOperands::
2228 MatrixAAndBBFloat16ComponentsINTEL ||
2229 CoopOperands &
2230 SPIRV::CooperativeMatrixOperands::MatrixCBFloat16ComponentsINTEL ||
2231 CoopOperands & SPIRV::CooperativeMatrixOperands::
2232 MatrixResultBFloat16ComponentsINTEL) {
2233 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_joint_matrix))
2234 report_fatal_error(reason: "***BF16ComponentsINTEL type interpretations "
2235 "require the following SPIR-V extension: "
2236 "SPV_INTEL_joint_matrix",
2237 gen_crash_diag: false);
2238 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_joint_matrix);
2239 Reqs.addCapability(
2240 ToAdd: SPIRV::Capability::CooperativeMatrixBFloat16ComponentTypeINTEL);
2241 }
2242 break;
2243 }
2244 case SPIRV::OpCooperativeMatrixLoadKHR:
2245 case SPIRV::OpCooperativeMatrixStoreKHR:
2246 case SPIRV::OpCooperativeMatrixLoadCheckedINTEL:
2247 case SPIRV::OpCooperativeMatrixStoreCheckedINTEL:
2248 case SPIRV::OpCooperativeMatrixPrefetchINTEL: {
2249 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_cooperative_matrix))
2250 report_fatal_error(reason: "Cooperative matrix instructions require the "
2251 "following SPIR-V extension: "
2252 "SPV_KHR_cooperative_matrix",
2253 gen_crash_diag: false);
2254 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_KHR_cooperative_matrix);
2255 Reqs.addCapability(ToAdd: SPIRV::Capability::CooperativeMatrixKHR);
2256
2257 // Check Layout operand in case if it's not a standard one and add the
2258 // appropriate capability.
2259 unsigned LayoutNum;
2260 switch (Op) {
2261 case SPIRV::OpCooperativeMatrixLoadKHR:
2262 LayoutNum = 3;
2263 break;
2264 case SPIRV::OpCooperativeMatrixStoreKHR:
2265 LayoutNum = 2;
2266 break;
2267 case SPIRV::OpCooperativeMatrixLoadCheckedINTEL:
2268 LayoutNum = 5;
2269 break;
2270 case SPIRV::OpCooperativeMatrixStoreCheckedINTEL:
2271 case SPIRV::OpCooperativeMatrixPrefetchINTEL:
2272 LayoutNum = 4;
2273 break;
2274 default:
2275 llvm_unreachable("unexpected cooperative matrix opcode");
2276 }
2277 Register RegLayout = MI.getOperand(i: LayoutNum).getReg();
2278 const MachineRegisterInfo &MRI = MI.getMF()->getRegInfo();
2279 MachineInstr *MILayout = MRI.getUniqueVRegDef(Reg: RegLayout);
2280 if (MILayout->getOpcode() == SPIRV::OpConstantI) {
2281 const unsigned LayoutVal = MILayout->getOperand(i: 2).getImm();
2282 if (LayoutVal ==
2283 static_cast<unsigned>(SPIRV::CooperativeMatrixLayout::PackedINTEL)) {
2284 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_joint_matrix))
2285 report_fatal_error(reason: "PackedINTEL layout require the following SPIR-V "
2286 "extension: SPV_INTEL_joint_matrix",
2287 gen_crash_diag: false);
2288 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_joint_matrix);
2289 Reqs.addCapability(ToAdd: SPIRV::Capability::PackedCooperativeMatrixINTEL);
2290 }
2291 }
2292
2293 // Nothing to do.
2294 if (Op == SPIRV::OpCooperativeMatrixLoadKHR ||
2295 Op == SPIRV::OpCooperativeMatrixStoreKHR)
2296 break;
2297
2298 std::string InstName;
2299 switch (Op) {
2300 case SPIRV::OpCooperativeMatrixPrefetchINTEL:
2301 InstName = "OpCooperativeMatrixPrefetchINTEL";
2302 break;
2303 case SPIRV::OpCooperativeMatrixLoadCheckedINTEL:
2304 InstName = "OpCooperativeMatrixLoadCheckedINTEL";
2305 break;
2306 case SPIRV::OpCooperativeMatrixStoreCheckedINTEL:
2307 InstName = "OpCooperativeMatrixStoreCheckedINTEL";
2308 break;
2309 }
2310
2311 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_joint_matrix)) {
2312 const std::string ErrorMsg =
2313 InstName + " instruction requires the "
2314 "following SPIR-V extension: SPV_INTEL_joint_matrix";
2315 report_fatal_error(reason: ErrorMsg.c_str(), gen_crash_diag: false);
2316 }
2317 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_joint_matrix);
2318 if (Op == SPIRV::OpCooperativeMatrixPrefetchINTEL) {
2319 Reqs.addCapability(ToAdd: SPIRV::Capability::CooperativeMatrixPrefetchINTEL);
2320 break;
2321 }
2322 Reqs.addCapability(
2323 ToAdd: SPIRV::Capability::CooperativeMatrixCheckedInstructionsINTEL);
2324 break;
2325 }
2326 case SPIRV::OpCooperativeMatrixConstructCheckedINTEL:
2327 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_joint_matrix))
2328 report_fatal_error(reason: "OpCooperativeMatrixConstructCheckedINTEL "
2329 "instructions require the following SPIR-V extension: "
2330 "SPV_INTEL_joint_matrix",
2331 gen_crash_diag: false);
2332 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_joint_matrix);
2333 Reqs.addCapability(
2334 ToAdd: SPIRV::Capability::CooperativeMatrixCheckedInstructionsINTEL);
2335 break;
2336 case SPIRV::OpReadPipeBlockingALTERA:
2337 case SPIRV::OpWritePipeBlockingALTERA:
2338 if (ST.canUseExtension(E: SPIRV::Extension::SPV_ALTERA_blocking_pipes)) {
2339 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_ALTERA_blocking_pipes);
2340 Reqs.addCapability(ToAdd: SPIRV::Capability::BlockingPipesALTERA);
2341 }
2342 break;
2343 case SPIRV::OpCooperativeMatrixGetElementCoordINTEL:
2344 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_joint_matrix))
2345 report_fatal_error(reason: "OpCooperativeMatrixGetElementCoordINTEL requires the "
2346 "following SPIR-V extension: SPV_INTEL_joint_matrix",
2347 gen_crash_diag: false);
2348 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_joint_matrix);
2349 Reqs.addCapability(
2350 ToAdd: SPIRV::Capability::CooperativeMatrixInvocationInstructionsINTEL);
2351 break;
2352 case SPIRV::OpConvertHandleToImageINTEL:
2353 case SPIRV::OpConvertHandleToSamplerINTEL:
2354 case SPIRV::OpConvertHandleToSampledImageINTEL: {
2355 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_bindless_images))
2356 report_fatal_error(reason: "OpConvertHandleTo[Image/Sampler/SampledImage]INTEL "
2357 "instructions require the following SPIR-V extension: "
2358 "SPV_INTEL_bindless_images",
2359 gen_crash_diag: false);
2360 SPIRVGlobalRegistry *GR = ST.getSPIRVGlobalRegistry();
2361 SPIRV::AddressingModel::AddressingModel AddrModel = MAI.Addr;
2362 SPIRVTypeInst TyDef = GR->getSPIRVTypeForVReg(VReg: MI.getOperand(i: 1).getReg());
2363 if (Op == SPIRV::OpConvertHandleToImageINTEL &&
2364 TyDef->getOpcode() != SPIRV::OpTypeImage) {
2365 report_fatal_error(reason: "Incorrect return type for the instruction "
2366 "OpConvertHandleToImageINTEL",
2367 gen_crash_diag: false);
2368 } else if (Op == SPIRV::OpConvertHandleToSamplerINTEL &&
2369 TyDef->getOpcode() != SPIRV::OpTypeSampler) {
2370 report_fatal_error(reason: "Incorrect return type for the instruction "
2371 "OpConvertHandleToSamplerINTEL",
2372 gen_crash_diag: false);
2373 } else if (Op == SPIRV::OpConvertHandleToSampledImageINTEL &&
2374 TyDef->getOpcode() != SPIRV::OpTypeSampledImage) {
2375 report_fatal_error(reason: "Incorrect return type for the instruction "
2376 "OpConvertHandleToSampledImageINTEL",
2377 gen_crash_diag: false);
2378 }
2379 SPIRVTypeInst SpvTy = GR->getSPIRVTypeForVReg(VReg: MI.getOperand(i: 2).getReg());
2380 unsigned Bitwidth = GR->getScalarOrVectorBitWidth(Type: SpvTy);
2381 if (!(Bitwidth == 32 && AddrModel == SPIRV::AddressingModel::Physical32) &&
2382 !(Bitwidth == 64 && AddrModel == SPIRV::AddressingModel::Physical64)) {
2383 report_fatal_error(
2384 reason: "Parameter value must be a 32-bit scalar in case of "
2385 "Physical32 addressing model or a 64-bit scalar in case of "
2386 "Physical64 addressing model",
2387 gen_crash_diag: false);
2388 }
2389 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_bindless_images);
2390 Reqs.addCapability(ToAdd: SPIRV::Capability::BindlessImagesINTEL);
2391 break;
2392 }
2393 case SPIRV::OpSubgroup2DBlockLoadINTEL:
2394 case SPIRV::OpSubgroup2DBlockLoadTransposeINTEL:
2395 case SPIRV::OpSubgroup2DBlockLoadTransformINTEL:
2396 case SPIRV::OpSubgroup2DBlockPrefetchINTEL:
2397 case SPIRV::OpSubgroup2DBlockStoreINTEL: {
2398 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_2d_block_io))
2399 report_fatal_error(reason: "OpSubgroup2DBlock[Load/LoadTranspose/LoadTransform/"
2400 "Prefetch/Store]INTEL instructions require the "
2401 "following SPIR-V extension: SPV_INTEL_2d_block_io",
2402 gen_crash_diag: false);
2403 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_2d_block_io);
2404 Reqs.addCapability(ToAdd: SPIRV::Capability::Subgroup2DBlockIOINTEL);
2405
2406 if (Op == SPIRV::OpSubgroup2DBlockLoadTransposeINTEL) {
2407 Reqs.addCapability(ToAdd: SPIRV::Capability::Subgroup2DBlockTransposeINTEL);
2408 break;
2409 }
2410 if (Op == SPIRV::OpSubgroup2DBlockLoadTransformINTEL) {
2411 Reqs.addCapability(ToAdd: SPIRV::Capability::Subgroup2DBlockTransformINTEL);
2412 break;
2413 }
2414 break;
2415 }
2416 case SPIRV::OpKill: {
2417 Reqs.addCapability(ToAdd: SPIRV::Capability::Shader);
2418 } break;
2419 case SPIRV::OpDemoteToHelperInvocation:
2420 Reqs.addCapability(ToAdd: SPIRV::Capability::DemoteToHelperInvocation);
2421
2422 if (ST.canUseExtension(
2423 E: SPIRV::Extension::SPV_EXT_demote_to_helper_invocation)) {
2424 if (!ST.isAtLeastSPIRVVer(VerToCompareTo: llvm::VersionTuple(1, 6)))
2425 Reqs.addExtension(
2426 ToAdd: SPIRV::Extension::SPV_EXT_demote_to_helper_invocation);
2427 }
2428 break;
2429 case SPIRV::OpSDot:
2430 case SPIRV::OpUDot:
2431 case SPIRV::OpSUDot:
2432 case SPIRV::OpSDotAccSat:
2433 case SPIRV::OpUDotAccSat:
2434 case SPIRV::OpSUDotAccSat:
2435 AddDotProductRequirements(MI, Reqs, ST);
2436 break;
2437 case SPIRV::OpImageSampleImplicitLod:
2438 case SPIRV::OpImageFetch:
2439 Reqs.addCapability(ToAdd: SPIRV::Capability::Shader);
2440 addImageOperandReqs(MI, Reqs, ST, OpIdx: 4);
2441 break;
2442 case SPIRV::OpImageSampleExplicitLod:
2443 addImageOperandReqs(MI, Reqs, ST, OpIdx: 4);
2444 break;
2445 case SPIRV::OpImageSampleDrefImplicitLod:
2446 case SPIRV::OpImageSampleDrefExplicitLod:
2447 case SPIRV::OpImageDrefGather:
2448 case SPIRV::OpImageGather:
2449 Reqs.addCapability(ToAdd: SPIRV::Capability::Shader);
2450 addImageOperandReqs(MI, Reqs, ST, OpIdx: 5);
2451 break;
2452 case SPIRV::OpImageRead: {
2453 Register ImageReg = MI.getOperand(i: 2).getReg();
2454 SPIRVTypeInst TypeDef = ST.getSPIRVGlobalRegistry()->getResultType(
2455 VReg: ImageReg, MF: const_cast<MachineFunction *>(MI.getMF()));
2456 // OpImageRead and OpImageWrite can use Unknown Image Formats
2457 // when the Kernel capability is declared. In the OpenCL environment we are
2458 // not allowed to produce
2459 // StorageImageReadWithoutFormat/StorageImageWriteWithoutFormat, see
2460 // https://github.com/KhronosGroup/SPIRV-Headers/issues/487
2461
2462 if (isImageTypeWithUnknownFormat(TypeInst: TypeDef) && ST.isShader())
2463 Reqs.addCapability(ToAdd: SPIRV::Capability::StorageImageReadWithoutFormat);
2464 break;
2465 }
2466 case SPIRV::OpImageWrite: {
2467 Register ImageReg = MI.getOperand(i: 0).getReg();
2468 SPIRVTypeInst TypeDef = ST.getSPIRVGlobalRegistry()->getResultType(
2469 VReg: ImageReg, MF: const_cast<MachineFunction *>(MI.getMF()));
2470 // OpImageRead and OpImageWrite can use Unknown Image Formats
2471 // when the Kernel capability is declared. In the OpenCL environment we are
2472 // not allowed to produce
2473 // StorageImageReadWithoutFormat/StorageImageWriteWithoutFormat, see
2474 // https://github.com/KhronosGroup/SPIRV-Headers/issues/487
2475
2476 if (isImageTypeWithUnknownFormat(TypeInst: TypeDef) && ST.isShader())
2477 Reqs.addCapability(ToAdd: SPIRV::Capability::StorageImageWriteWithoutFormat);
2478 break;
2479 }
2480 case SPIRV::OpTypeStructContinuedINTEL:
2481 case SPIRV::OpConstantCompositeContinuedINTEL:
2482 case SPIRV::OpSpecConstantCompositeContinuedINTEL:
2483 case SPIRV::OpCompositeConstructContinuedINTEL: {
2484 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_long_composites))
2485 report_fatal_error(
2486 reason: "Continued instructions require the "
2487 "following SPIR-V extension: SPV_INTEL_long_composites",
2488 gen_crash_diag: false);
2489 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_long_composites);
2490 Reqs.addCapability(ToAdd: SPIRV::Capability::LongCompositesINTEL);
2491 break;
2492 }
2493 case SPIRV::OpArbitraryFloatEQALTERA:
2494 case SPIRV::OpArbitraryFloatGEALTERA:
2495 case SPIRV::OpArbitraryFloatGTALTERA:
2496 case SPIRV::OpArbitraryFloatLEALTERA:
2497 case SPIRV::OpArbitraryFloatLTALTERA:
2498 case SPIRV::OpArbitraryFloatCbrtALTERA:
2499 case SPIRV::OpArbitraryFloatCosALTERA:
2500 case SPIRV::OpArbitraryFloatCosPiALTERA:
2501 case SPIRV::OpArbitraryFloatExp10ALTERA:
2502 case SPIRV::OpArbitraryFloatExp2ALTERA:
2503 case SPIRV::OpArbitraryFloatExpALTERA:
2504 case SPIRV::OpArbitraryFloatExpm1ALTERA:
2505 case SPIRV::OpArbitraryFloatHypotALTERA:
2506 case SPIRV::OpArbitraryFloatLog10ALTERA:
2507 case SPIRV::OpArbitraryFloatLog1pALTERA:
2508 case SPIRV::OpArbitraryFloatLog2ALTERA:
2509 case SPIRV::OpArbitraryFloatLogALTERA:
2510 case SPIRV::OpArbitraryFloatRecipALTERA:
2511 case SPIRV::OpArbitraryFloatSinCosALTERA:
2512 case SPIRV::OpArbitraryFloatSinCosPiALTERA:
2513 case SPIRV::OpArbitraryFloatSinALTERA:
2514 case SPIRV::OpArbitraryFloatSinPiALTERA:
2515 case SPIRV::OpArbitraryFloatSqrtALTERA:
2516 case SPIRV::OpArbitraryFloatACosALTERA:
2517 case SPIRV::OpArbitraryFloatACosPiALTERA:
2518 case SPIRV::OpArbitraryFloatAddALTERA:
2519 case SPIRV::OpArbitraryFloatASinALTERA:
2520 case SPIRV::OpArbitraryFloatASinPiALTERA:
2521 case SPIRV::OpArbitraryFloatATan2ALTERA:
2522 case SPIRV::OpArbitraryFloatATanALTERA:
2523 case SPIRV::OpArbitraryFloatATanPiALTERA:
2524 case SPIRV::OpArbitraryFloatCastFromIntALTERA:
2525 case SPIRV::OpArbitraryFloatCastALTERA:
2526 case SPIRV::OpArbitraryFloatCastToIntALTERA:
2527 case SPIRV::OpArbitraryFloatDivALTERA:
2528 case SPIRV::OpArbitraryFloatMulALTERA:
2529 case SPIRV::OpArbitraryFloatPowALTERA:
2530 case SPIRV::OpArbitraryFloatPowNALTERA:
2531 case SPIRV::OpArbitraryFloatPowRALTERA:
2532 case SPIRV::OpArbitraryFloatRSqrtALTERA:
2533 case SPIRV::OpArbitraryFloatSubALTERA: {
2534 if (!ST.canUseExtension(
2535 E: SPIRV::Extension::SPV_ALTERA_arbitrary_precision_floating_point))
2536 report_fatal_error(
2537 reason: "Floating point instructions can't be translated correctly without "
2538 "enabled SPV_ALTERA_arbitrary_precision_floating_point extension!",
2539 gen_crash_diag: false);
2540 Reqs.addExtension(
2541 ToAdd: SPIRV::Extension::SPV_ALTERA_arbitrary_precision_floating_point);
2542 Reqs.addCapability(
2543 ToAdd: SPIRV::Capability::ArbitraryPrecisionFloatingPointALTERA);
2544 break;
2545 }
2546 case SPIRV::OpSubgroupMatrixMultiplyAccumulateINTEL: {
2547 if (!ST.canUseExtension(
2548 E: SPIRV::Extension::SPV_INTEL_subgroup_matrix_multiply_accumulate))
2549 report_fatal_error(
2550 reason: "OpSubgroupMatrixMultiplyAccumulateINTEL instruction requires the "
2551 "following SPIR-V "
2552 "extension: SPV_INTEL_subgroup_matrix_multiply_accumulate",
2553 gen_crash_diag: false);
2554 Reqs.addExtension(
2555 ToAdd: SPIRV::Extension::SPV_INTEL_subgroup_matrix_multiply_accumulate);
2556 Reqs.addCapability(
2557 ToAdd: SPIRV::Capability::SubgroupMatrixMultiplyAccumulateINTEL);
2558 break;
2559 }
2560 case SPIRV::OpBitwiseFunctionINTEL: {
2561 if (!ST.canUseExtension(
2562 E: SPIRV::Extension::SPV_INTEL_ternary_bitwise_function))
2563 report_fatal_error(
2564 reason: "OpBitwiseFunctionINTEL instruction requires the following SPIR-V "
2565 "extension: SPV_INTEL_ternary_bitwise_function",
2566 gen_crash_diag: false);
2567 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_ternary_bitwise_function);
2568 Reqs.addCapability(ToAdd: SPIRV::Capability::TernaryBitwiseFunctionINTEL);
2569 break;
2570 }
2571 case SPIRV::OpCopyMemorySized: {
2572 Reqs.addCapability(ToAdd: SPIRV::Capability::Addresses);
2573 // TODO: Add UntypedPointersKHR when implemented.
2574 break;
2575 }
2576 case SPIRV::OpTypeUntypedPointerKHR:
2577 Reqs.getAndAddRequirements(Category: SPIRV::OperandCategory::StorageClassOperand,
2578 i: MI.getOperand(i: 1).getImm(), ST);
2579 [[fallthrough]];
2580 case SPIRV::OpUntypedVariableKHR:
2581 case SPIRV::OpUntypedAccessChainKHR:
2582 case SPIRV::OpUntypedInBoundsAccessChainKHR:
2583 case SPIRV::OpUntypedPtrAccessChainKHR:
2584 case SPIRV::OpUntypedInBoundsPtrAccessChainKHR:
2585 case SPIRV::OpUntypedPrefetchKHR:
2586 case SPIRV::OpUntypedGroupAsyncCopyKHR: {
2587 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_untyped_pointers))
2588 report_fatal_error(reason: "Untyped pointer instructions require the following "
2589 "SPIR-V extension: SPV_KHR_untyped_pointers",
2590 gen_crash_diag: false);
2591 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_KHR_untyped_pointers);
2592 Reqs.addCapability(ToAdd: SPIRV::Capability::UntypedPointersKHR);
2593 break;
2594 }
2595 case SPIRV::OpPredicatedLoadINTEL:
2596 case SPIRV::OpPredicatedStoreINTEL: {
2597 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_predicated_io))
2598 report_fatal_error(
2599 reason: "OpPredicated[Load/Store]INTEL instructions require "
2600 "the following SPIR-V extension: SPV_INTEL_predicated_io",
2601 gen_crash_diag: false);
2602 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_predicated_io);
2603 Reqs.addCapability(ToAdd: SPIRV::Capability::PredicatedIOINTEL);
2604 break;
2605 }
2606 case SPIRV::OpFAddS:
2607 case SPIRV::OpFSubS:
2608 case SPIRV::OpFMulS:
2609 case SPIRV::OpFDivS:
2610 case SPIRV::OpFRemS:
2611 case SPIRV::OpFMod:
2612 case SPIRV::OpFNegate:
2613 case SPIRV::OpFAddV:
2614 case SPIRV::OpFSubV:
2615 case SPIRV::OpFMulV:
2616 case SPIRV::OpFDivV:
2617 case SPIRV::OpFRemV:
2618 case SPIRV::OpFNegateV: {
2619 const MachineRegisterInfo &MRI = MI.getMF()->getRegInfo();
2620 SPIRVTypeInst TypeDef = MRI.getVRegDef(Reg: MI.getOperand(i: 1).getReg());
2621 if (isVectorType(SPVTy: TypeDef))
2622 TypeDef = MRI.getVRegDef(Reg: TypeDef->getOperand(i: 1).getReg());
2623 if (isBFloat16Type(TypeDef)) {
2624 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_bfloat16_arithmetic))
2625 report_fatal_error(
2626 reason: "Arithmetic instructions with bfloat16 arguments require the "
2627 "following SPIR-V extension: SPV_INTEL_bfloat16_arithmetic",
2628 gen_crash_diag: false);
2629 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_bfloat16_arithmetic);
2630 Reqs.addCapability(ToAdd: SPIRV::Capability::BFloat16ArithmeticINTEL);
2631 }
2632 break;
2633 }
2634 case SPIRV::OpOrdered:
2635 case SPIRV::OpUnordered:
2636 case SPIRV::OpFOrdEqual:
2637 case SPIRV::OpFOrdNotEqual:
2638 case SPIRV::OpFOrdLessThan:
2639 case SPIRV::OpFOrdLessThanEqual:
2640 case SPIRV::OpFOrdGreaterThan:
2641 case SPIRV::OpFOrdGreaterThanEqual:
2642 case SPIRV::OpFUnordEqual:
2643 case SPIRV::OpFUnordNotEqual:
2644 case SPIRV::OpFUnordLessThan:
2645 case SPIRV::OpFUnordLessThanEqual:
2646 case SPIRV::OpFUnordGreaterThan:
2647 case SPIRV::OpFUnordGreaterThanEqual: {
2648 const MachineRegisterInfo &MRI = MI.getMF()->getRegInfo();
2649 MachineInstr *OperandDef = MRI.getVRegDef(Reg: MI.getOperand(i: 2).getReg());
2650 SPIRVTypeInst TypeDef = MRI.getVRegDef(Reg: OperandDef->getOperand(i: 1).getReg());
2651 if (isVectorType(SPVTy: TypeDef))
2652 TypeDef = MRI.getVRegDef(Reg: TypeDef->getOperand(i: 1).getReg());
2653 if (isBFloat16Type(TypeDef)) {
2654 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_bfloat16_arithmetic))
2655 report_fatal_error(
2656 reason: "Relational instructions with bfloat16 arguments require the "
2657 "following SPIR-V extension: SPV_INTEL_bfloat16_arithmetic",
2658 gen_crash_diag: false);
2659 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_bfloat16_arithmetic);
2660 Reqs.addCapability(ToAdd: SPIRV::Capability::BFloat16ArithmeticINTEL);
2661 }
2662 break;
2663 }
2664 case SPIRV::OpDPdxCoarse:
2665 case SPIRV::OpDPdyCoarse:
2666 case SPIRV::OpDPdxFine:
2667 case SPIRV::OpDPdyFine: {
2668 Reqs.addCapability(ToAdd: SPIRV::Capability::DerivativeControl);
2669 break;
2670 }
2671 case SPIRV::OpLoopControlINTEL: {
2672 Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_unstructured_loop_controls);
2673 Reqs.addCapability(ToAdd: SPIRV::Capability::UnstructuredLoopControlsINTEL);
2674 break;
2675 }
2676
2677 default:
2678 break;
2679 }
2680
2681 // If we require capability Shader, then we can remove the requirement for
2682 // the BitInstructions capability, since Shader is a superset capability
2683 // of BitInstructions.
2684 Reqs.removeCapabilityIf(ToRemove: SPIRV::Capability::BitInstructions,
2685 IfPresent: SPIRV::Capability::Shader);
2686}
2687
2688static void collectReqs(const Module &M, SPIRV::ModuleAnalysisInfo &MAI,
2689 MachineFunctionGetter GetMF, const SPIRVSubtarget &ST) {
2690 // Collect requirements for existing instructions.
2691 for (const Function &F : M) {
2692 MachineFunction *MF = GetMF(F);
2693 if (!MF)
2694 continue;
2695 for (const MachineBasicBlock &MBB : *MF)
2696 for (const MachineInstr &MI : MBB)
2697 addInstrRequirements(MI, MAI, ST);
2698 }
2699 // Collect requirements for OpExecutionMode instructions.
2700 auto Node = M.getNamedMetadata(Name: "spirv.ExecutionMode");
2701 if (Node) {
2702 bool RequireFloatControls = false, RequireIntelFloatControls2 = false,
2703 RequireKHRFloatControls2 = false,
2704 VerLower14 = !ST.isAtLeastSPIRVVer(VerToCompareTo: VersionTuple(1, 4));
2705 bool HasIntelFloatControls2 =
2706 ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_float_controls2);
2707 bool HasKHRFloatControls2 =
2708 ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_float_controls2);
2709 for (unsigned i = 0; i < Node->getNumOperands(); i++) {
2710 MDNode *MDN = cast<MDNode>(Val: Node->getOperand(i));
2711 const MDOperand &MDOp = MDN->getOperand(I: 1);
2712 if (auto *CMeta = dyn_cast<ConstantAsMetadata>(Val: MDOp)) {
2713 Constant *C = CMeta->getValue();
2714 if (ConstantInt *Const = dyn_cast<ConstantInt>(Val: C)) {
2715 auto EM = Const->getZExtValue();
2716 // SPV_KHR_float_controls is not available until v1.4:
2717 // add SPV_KHR_float_controls if the version is too low
2718 switch (EM) {
2719 case SPIRV::ExecutionMode::DenormPreserve:
2720 case SPIRV::ExecutionMode::DenormFlushToZero:
2721 case SPIRV::ExecutionMode::RoundingModeRTE:
2722 case SPIRV::ExecutionMode::RoundingModeRTZ:
2723 RequireFloatControls = VerLower14;
2724 MAI.Reqs.getAndAddRequirements(
2725 Category: SPIRV::OperandCategory::ExecutionModeOperand, i: EM, ST);
2726 break;
2727 case SPIRV::ExecutionMode::RoundingModeRTPINTEL:
2728 case SPIRV::ExecutionMode::RoundingModeRTNINTEL:
2729 case SPIRV::ExecutionMode::FloatingPointModeALTINTEL:
2730 case SPIRV::ExecutionMode::FloatingPointModeIEEEINTEL:
2731 if (HasIntelFloatControls2) {
2732 RequireIntelFloatControls2 = true;
2733 MAI.Reqs.getAndAddRequirements(
2734 Category: SPIRV::OperandCategory::ExecutionModeOperand, i: EM, ST);
2735 }
2736 break;
2737 case SPIRV::ExecutionMode::FPFastMathDefault: {
2738 if (HasKHRFloatControls2) {
2739 RequireKHRFloatControls2 = true;
2740 MAI.Reqs.getAndAddRequirements(
2741 Category: SPIRV::OperandCategory::ExecutionModeOperand, i: EM, ST);
2742 }
2743 break;
2744 }
2745 case SPIRV::ExecutionMode::ContractionOff:
2746 case SPIRV::ExecutionMode::SignedZeroInfNanPreserve:
2747 if (HasKHRFloatControls2) {
2748 RequireKHRFloatControls2 = true;
2749 MAI.Reqs.getAndAddRequirements(
2750 Category: SPIRV::OperandCategory::ExecutionModeOperand,
2751 i: SPIRV::ExecutionMode::FPFastMathDefault, ST);
2752 } else {
2753 MAI.Reqs.getAndAddRequirements(
2754 Category: SPIRV::OperandCategory::ExecutionModeOperand, i: EM, ST);
2755 }
2756 break;
2757 default:
2758 MAI.Reqs.getAndAddRequirements(
2759 Category: SPIRV::OperandCategory::ExecutionModeOperand, i: EM, ST);
2760 }
2761 }
2762 }
2763 }
2764 if (RequireFloatControls &&
2765 ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_float_controls))
2766 MAI.Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_KHR_float_controls);
2767 if (RequireIntelFloatControls2)
2768 MAI.Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_float_controls2);
2769 if (RequireKHRFloatControls2)
2770 MAI.Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_KHR_float_controls2);
2771 }
2772 for (const Function &F : M) {
2773 if (F.isDeclaration())
2774 continue;
2775 if (F.getMetadata(Kind: "reqd_work_group_size"))
2776 MAI.Reqs.getAndAddRequirements(
2777 Category: SPIRV::OperandCategory::ExecutionModeOperand,
2778 i: SPIRV::ExecutionMode::LocalSize, ST);
2779 if (F.getFnAttribute(Kind: "hlsl.numthreads").isValid()) {
2780 MAI.Reqs.getAndAddRequirements(
2781 Category: SPIRV::OperandCategory::ExecutionModeOperand,
2782 i: SPIRV::ExecutionMode::LocalSize, ST);
2783 }
2784 if (F.getFnAttribute(Kind: "enable-maximal-reconvergence").getValueAsBool()) {
2785 MAI.Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_KHR_maximal_reconvergence);
2786 }
2787 if (F.getMetadata(Kind: "work_group_size_hint"))
2788 MAI.Reqs.getAndAddRequirements(
2789 Category: SPIRV::OperandCategory::ExecutionModeOperand,
2790 i: SPIRV::ExecutionMode::LocalSizeHint, ST);
2791 if (F.getMetadata(Kind: "intel_reqd_sub_group_size") ||
2792 F.getMetadata(Kind: "reqd_sub_group_size"))
2793 MAI.Reqs.getAndAddRequirements(
2794 Category: SPIRV::OperandCategory::ExecutionModeOperand,
2795 i: SPIRV::ExecutionMode::SubgroupSize, ST);
2796 if (F.getMetadata(Kind: "max_work_group_size"))
2797 MAI.Reqs.getAndAddRequirements(
2798 Category: SPIRV::OperandCategory::ExecutionModeOperand,
2799 i: SPIRV::ExecutionMode::MaxWorkgroupSizeINTEL, ST);
2800 if (F.getMetadata(Kind: "vec_type_hint"))
2801 MAI.Reqs.getAndAddRequirements(
2802 Category: SPIRV::OperandCategory::ExecutionModeOperand,
2803 i: SPIRV::ExecutionMode::VecTypeHint, ST);
2804
2805 if (F.hasOptNone()) {
2806 if (ST.canUseExtension(E: SPIRV::Extension::SPV_INTEL_optnone)) {
2807 MAI.Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_INTEL_optnone);
2808 MAI.Reqs.addCapability(ToAdd: SPIRV::Capability::OptNoneINTEL);
2809 } else if (ST.canUseExtension(E: SPIRV::Extension::SPV_EXT_optnone)) {
2810 MAI.Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_EXT_optnone);
2811 MAI.Reqs.addCapability(ToAdd: SPIRV::Capability::OptNoneEXT);
2812 }
2813 }
2814 }
2815}
2816
2817static unsigned getFastMathFlags(const MachineInstr &I,
2818 const SPIRVSubtarget &ST) {
2819 unsigned Flags = SPIRV::FPFastMathMode::None;
2820 bool CanUseKHRFloatControls2 =
2821 ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_float_controls2);
2822 if (I.getFlag(Flag: MachineInstr::MIFlag::FmNoNans))
2823 Flags |= SPIRV::FPFastMathMode::NotNaN;
2824 if (I.getFlag(Flag: MachineInstr::MIFlag::FmNoInfs))
2825 Flags |= SPIRV::FPFastMathMode::NotInf;
2826 if (I.getFlag(Flag: MachineInstr::MIFlag::FmNsz))
2827 Flags |= SPIRV::FPFastMathMode::NSZ;
2828 if (I.getFlag(Flag: MachineInstr::MIFlag::FmArcp))
2829 Flags |= SPIRV::FPFastMathMode::AllowRecip;
2830 if (I.getFlag(Flag: MachineInstr::MIFlag::FmContract) && CanUseKHRFloatControls2)
2831 Flags |= SPIRV::FPFastMathMode::AllowContract;
2832 if (I.getFlag(Flag: MachineInstr::MIFlag::FmReassoc)) {
2833 if (CanUseKHRFloatControls2)
2834 // LLVM reassoc maps to SPIRV transform, see
2835 // https://github.com/KhronosGroup/SPIRV-Registry/issues/326 for details.
2836 // Because we are enabling AllowTransform, we must enable AllowReassoc and
2837 // AllowContract too, as required by SPIRV spec. Also, we used to map
2838 // MIFlag::FmReassoc to FPFastMathMode::Fast, which now should instead by
2839 // replaced by turning all the other bits instead. Therefore, we're
2840 // enabling every bit here except None and Fast.
2841 Flags |= SPIRV::FPFastMathMode::NotNaN | SPIRV::FPFastMathMode::NotInf |
2842 SPIRV::FPFastMathMode::NSZ | SPIRV::FPFastMathMode::AllowRecip |
2843 SPIRV::FPFastMathMode::AllowTransform |
2844 SPIRV::FPFastMathMode::AllowReassoc |
2845 SPIRV::FPFastMathMode::AllowContract;
2846 else
2847 Flags |= SPIRV::FPFastMathMode::Fast;
2848 }
2849
2850 if (CanUseKHRFloatControls2) {
2851 // Error out if SPIRV::FPFastMathMode::Fast is enabled.
2852 assert(!(Flags & SPIRV::FPFastMathMode::Fast) &&
2853 "SPIRV::FPFastMathMode::Fast is deprecated and should not be used "
2854 "anymore.");
2855
2856 // Error out if AllowTransform is enabled without AllowReassoc and
2857 // AllowContract.
2858 assert((!(Flags & SPIRV::FPFastMathMode::AllowTransform) ||
2859 ((Flags & SPIRV::FPFastMathMode::AllowReassoc &&
2860 Flags & SPIRV::FPFastMathMode::AllowContract))) &&
2861 "SPIRV::FPFastMathMode::AllowTransform requires AllowReassoc and "
2862 "AllowContract flags to be enabled as well.");
2863 }
2864
2865 return Flags;
2866}
2867
2868static bool isFastMathModeAvailable(const SPIRVSubtarget &ST) {
2869 if (ST.isKernel())
2870 return true;
2871 if (ST.getSPIRVVersion() < VersionTuple(1, 2))
2872 return false;
2873 return ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_float_controls2);
2874}
2875
2876static void handleMIFlagDecoration(
2877 MachineInstr &I, const SPIRVSubtarget &ST, const SPIRVInstrInfo &TII,
2878 SPIRV::RequirementHandler &Reqs, const SPIRVGlobalRegistry *GR,
2879 SPIRV::FPFastMathDefaultInfoVector &FPFastMathDefaultInfoVec) {
2880 if (TII.canUseIntegerWrapDecoration(MI: I)) {
2881 if (I.getFlag(Flag: MachineInstr::MIFlag::NoSWrap) &&
2882 getSymbolicOperandRequirements(
2883 Category: SPIRV::OperandCategory::DecorationOperand,
2884 i: SPIRV::Decoration::NoSignedWrap, ST, Reqs)
2885 .IsSatisfiable)
2886 buildOpDecorate(Reg: I.getOperand(i: 0).getReg(), I, TII,
2887 Dec: SPIRV::Decoration::NoSignedWrap, DecArgs: {});
2888 if (I.getFlag(Flag: MachineInstr::MIFlag::NoUWrap) &&
2889 getSymbolicOperandRequirements(
2890 Category: SPIRV::OperandCategory::DecorationOperand,
2891 i: SPIRV::Decoration::NoUnsignedWrap, ST, Reqs)
2892 .IsSatisfiable)
2893 buildOpDecorate(Reg: I.getOperand(i: 0).getReg(), I, TII,
2894 Dec: SPIRV::Decoration::NoUnsignedWrap, DecArgs: {});
2895 }
2896 // In Kernel environments, FPFastMathMode on OpExtInst is valid per core
2897 // spec. For other instruction types, SPV_KHR_float_controls2 is required.
2898 bool CanUseFM =
2899 TII.canUseFastMathFlags(
2900 MI: I, KHRFloatControls2: ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_float_controls2)) ||
2901 (ST.isKernel() && I.getOpcode() == SPIRV::OpExtInst);
2902 if (!CanUseFM)
2903 return;
2904
2905 unsigned FMFlags = getFastMathFlags(I, ST);
2906 if (FMFlags == SPIRV::FPFastMathMode::None) {
2907 // We also need to check if any FPFastMathDefault info was set for the
2908 // types used in this instruction.
2909 if (FPFastMathDefaultInfoVec.empty())
2910 return;
2911
2912 // There are three types of instructions that can use fast math flags:
2913 // 1. Arithmetic instructions (FAdd, FMul, FSub, FDiv, FRem, etc.)
2914 // 2. Relational instructions (FCmp, FOrd, FUnord, etc.)
2915 // 3. Extended instructions (ExtInst)
2916 // For arithmetic instructions, the floating point type can be in the
2917 // result type or in the operands, but they all must be the same.
2918 // For the relational and logical instructions, the floating point type
2919 // can only be in the operands 1 and 2, not the result type. Also, the
2920 // operands must have the same type. For the extended instructions, the
2921 // floating point type can be in the result type or in the operands. It's
2922 // unclear if the operands and the result type must be the same. Let's
2923 // assume they must be. Therefore, for 1. and 2., we can check the first
2924 // operand type, and for 3. we can check the result type.
2925 assert(I.getNumOperands() >= 3 && "Expected at least 3 operands");
2926 Register ResReg = I.getOpcode() == SPIRV::OpExtInst
2927 ? I.getOperand(i: 1).getReg()
2928 : I.getOperand(i: 2).getReg();
2929 SPIRVTypeInst ResType = GR->getSPIRVTypeForVReg(VReg: ResReg, MF: I.getMF());
2930 const Type *Ty = GR->getTypeForSPIRVType(Ty: ResType);
2931 Ty = Ty->isVectorTy() ? cast<VectorType>(Val: Ty)->getElementType() : Ty;
2932
2933 // Match instruction type with the FPFastMathDefaultInfoVec.
2934 bool Emit = false;
2935 for (SPIRV::FPFastMathDefaultInfo &Elem : FPFastMathDefaultInfoVec) {
2936 if (Ty == Elem.Ty) {
2937 FMFlags = Elem.FastMathFlags;
2938 Emit = Elem.ContractionOff || Elem.SignedZeroInfNanPreserve ||
2939 Elem.FPFastMathDefault;
2940 break;
2941 }
2942 }
2943
2944 if (FMFlags == SPIRV::FPFastMathMode::None && !Emit)
2945 return;
2946 }
2947 if (isFastMathModeAvailable(ST)) {
2948 Register DstReg = I.getOperand(i: 0).getReg();
2949 buildOpDecorate(Reg: DstReg, I, TII, Dec: SPIRV::Decoration::FPFastMathMode,
2950 DecArgs: {FMFlags});
2951 }
2952}
2953
2954// Walk all functions and add decorations related to MI flags.
2955static void addDecorations(const Module &M, const SPIRVInstrInfo &TII,
2956 MachineFunctionGetter GetMF,
2957 const SPIRVSubtarget &ST,
2958 SPIRV::ModuleAnalysisInfo &MAI,
2959 const SPIRVGlobalRegistry *GR) {
2960 for (const Function &F : M) {
2961 MachineFunction *MF = GetMF(F);
2962 if (!MF)
2963 continue;
2964
2965 for (auto &MBB : *MF)
2966 for (auto &MI : MBB)
2967 handleMIFlagDecoration(I&: MI, ST, TII, Reqs&: MAI.Reqs, GR,
2968 FPFastMathDefaultInfoVec&: MAI.FPFastMathDefaultInfoMap[&F]);
2969 }
2970}
2971
2972static void addMBBNames(const Module &M, const SPIRVInstrInfo &TII,
2973 MachineFunctionGetter GetMF, const SPIRVSubtarget &ST,
2974 SPIRV::ModuleAnalysisInfo &MAI) {
2975 for (const Function &F : M) {
2976 MachineFunction *MF = GetMF(F);
2977 if (!MF)
2978 continue;
2979 if (MF->getFunction()
2980 .getFnAttribute(SPIRV_BACKEND_SERVICE_FUN_NAME)
2981 .isValid())
2982 continue;
2983 MachineRegisterInfo &MRI = MF->getRegInfo();
2984 for (auto &MBB : *MF) {
2985 if (!MBB.hasName() || MBB.empty())
2986 continue;
2987 // Emit basic block names.
2988 Register Reg = MRI.createGenericVirtualRegister(Ty: LLT::scalar(SizeInBits: 64));
2989 MRI.setRegClass(Reg, RC: &SPIRV::IDRegClass);
2990 buildOpName(Target: Reg, Name: MBB.getName(), I&: *std::prev(x: MBB.end()), TII);
2991 MCRegister GlobalReg = MAI.getOrCreateMBBRegister(MBB);
2992 MAI.setRegisterAlias(MF, Reg, AliasReg: GlobalReg);
2993 }
2994 }
2995}
2996
2997// patching Instruction::PHI to SPIRV::OpPhi
2998static void patchPhis(const Module &M, SPIRVGlobalRegistry *GR,
2999 const SPIRVInstrInfo &TII, MachineFunctionGetter GetMF) {
3000 for (const Function &F : M) {
3001 MachineFunction *MF = GetMF(F);
3002 if (!MF)
3003 continue;
3004 for (auto &MBB : *MF) {
3005 for (MachineInstr &MI : MBB.phis()) {
3006 MI.setDesc(TII.get(Opcode: SPIRV::OpPhi));
3007 Register ResTypeReg = GR->getSPIRVTypeID(
3008 SpirvType: GR->getSPIRVTypeForVReg(VReg: MI.getOperand(i: 0).getReg(), MF));
3009 MI.insert(InsertBefore: MI.operands_begin() + 1,
3010 Ops: {MachineOperand::CreateReg(Reg: ResTypeReg, isDef: false)});
3011 }
3012 }
3013
3014 MF->getProperties().setNoPHIs();
3015 }
3016}
3017
3018static SPIRV::FPFastMathDefaultInfoVector &getOrCreateFPFastMathDefaultInfoVec(
3019 const Module &M, SPIRV::ModuleAnalysisInfo &MAI, const Function *F) {
3020 auto it = MAI.FPFastMathDefaultInfoMap.find(Val: F);
3021 if (it != MAI.FPFastMathDefaultInfoMap.end())
3022 return it->second;
3023
3024 // If the map does not contain the entry, create a new one. Initialize it to
3025 // contain all 3 elements sorted by bit width of target type: {half, float,
3026 // double}.
3027 SPIRV::FPFastMathDefaultInfoVector FPFastMathDefaultInfoVec;
3028 FPFastMathDefaultInfoVec.emplace_back(Args: Type::getHalfTy(C&: M.getContext()),
3029 Args: SPIRV::FPFastMathMode::None);
3030 FPFastMathDefaultInfoVec.emplace_back(Args: Type::getFloatTy(C&: M.getContext()),
3031 Args: SPIRV::FPFastMathMode::None);
3032 FPFastMathDefaultInfoVec.emplace_back(Args: Type::getDoubleTy(C&: M.getContext()),
3033 Args: SPIRV::FPFastMathMode::None);
3034 return MAI.FPFastMathDefaultInfoMap[F] = std::move(FPFastMathDefaultInfoVec);
3035}
3036
3037static SPIRV::FPFastMathDefaultInfo &getFPFastMathDefaultInfo(
3038 SPIRV::FPFastMathDefaultInfoVector &FPFastMathDefaultInfoVec,
3039 const Type *Ty) {
3040 size_t BitWidth = Ty->getScalarSizeInBits();
3041 int Index =
3042 SPIRV::FPFastMathDefaultInfoVector::computeFPFastMathDefaultInfoVecIndex(
3043 BitWidth);
3044 assert(Index >= 0 && Index < 3 &&
3045 "Expected FPFastMathDefaultInfo for half, float, or double");
3046 assert(FPFastMathDefaultInfoVec.size() == 3 &&
3047 "Expected FPFastMathDefaultInfoVec to have exactly 3 elements");
3048 return FPFastMathDefaultInfoVec[Index];
3049}
3050
3051static void collectFPFastMathDefaults(const Module &M,
3052 SPIRV::ModuleAnalysisInfo &MAI,
3053 const SPIRVSubtarget &ST) {
3054 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_float_controls2))
3055 return;
3056
3057 // Store the FPFastMathDefaultInfo in the FPFastMathDefaultInfoMap.
3058 // We need the entry point (function) as the key, and the target
3059 // type and flags as the value.
3060 // We also need to check ContractionOff and SignedZeroInfNanPreserve
3061 // execution modes, as they are now deprecated and must be replaced
3062 // with FPFastMathDefaultInfo.
3063 auto Node = M.getNamedMetadata(Name: "spirv.ExecutionMode");
3064 if (!Node)
3065 return;
3066
3067 for (unsigned i = 0; i < Node->getNumOperands(); i++) {
3068 MDNode *MDN = cast<MDNode>(Val: Node->getOperand(i));
3069 assert(MDN->getNumOperands() >= 2 && "Expected at least 2 operands");
3070 const Function *F = cast<Function>(
3071 Val: cast<ConstantAsMetadata>(Val: MDN->getOperand(I: 0))->getValue());
3072 const auto EM =
3073 cast<ConstantInt>(
3074 Val: cast<ConstantAsMetadata>(Val: MDN->getOperand(I: 1))->getValue())
3075 ->getZExtValue();
3076 if (EM == SPIRV::ExecutionMode::FPFastMathDefault) {
3077 assert(MDN->getNumOperands() == 4 &&
3078 "Expected 4 operands for FPFastMathDefault");
3079
3080 const Type *T = cast<ValueAsMetadata>(Val: MDN->getOperand(I: 2))->getType();
3081 unsigned Flags =
3082 cast<ConstantInt>(
3083 Val: cast<ConstantAsMetadata>(Val: MDN->getOperand(I: 3))->getValue())
3084 ->getZExtValue();
3085 SPIRV::FPFastMathDefaultInfoVector &FPFastMathDefaultInfoVec =
3086 getOrCreateFPFastMathDefaultInfoVec(M, MAI, F);
3087 SPIRV::FPFastMathDefaultInfo &Info =
3088 getFPFastMathDefaultInfo(FPFastMathDefaultInfoVec, Ty: T);
3089 Info.FastMathFlags = Flags;
3090 Info.FPFastMathDefault = true;
3091 } else if (EM == SPIRV::ExecutionMode::ContractionOff) {
3092 assert(MDN->getNumOperands() == 2 &&
3093 "Expected no operands for ContractionOff");
3094
3095 // We need to save this info for every possible FP type, i.e. {half,
3096 // float, double, fp128}.
3097 SPIRV::FPFastMathDefaultInfoVector &FPFastMathDefaultInfoVec =
3098 getOrCreateFPFastMathDefaultInfoVec(M, MAI, F);
3099 for (SPIRV::FPFastMathDefaultInfo &Info : FPFastMathDefaultInfoVec) {
3100 Info.ContractionOff = true;
3101 }
3102 } else if (EM == SPIRV::ExecutionMode::SignedZeroInfNanPreserve) {
3103 assert(MDN->getNumOperands() == 3 &&
3104 "Expected 1 operand for SignedZeroInfNanPreserve");
3105 unsigned TargetWidth =
3106 cast<ConstantInt>(
3107 Val: cast<ConstantAsMetadata>(Val: MDN->getOperand(I: 2))->getValue())
3108 ->getZExtValue();
3109 // We need to save this info only for the FP type with TargetWidth.
3110 SPIRV::FPFastMathDefaultInfoVector &FPFastMathDefaultInfoVec =
3111 getOrCreateFPFastMathDefaultInfoVec(M, MAI, F);
3112 int Index = SPIRV::FPFastMathDefaultInfoVector::
3113 computeFPFastMathDefaultInfoVecIndex(BitWidth: TargetWidth);
3114 assert(Index >= 0 && Index < 3 &&
3115 "Expected FPFastMathDefaultInfo for half, float, or double");
3116 assert(FPFastMathDefaultInfoVec.size() == 3 &&
3117 "Expected FPFastMathDefaultInfoVec to have exactly 3 elements");
3118 FPFastMathDefaultInfoVec[Index].SignedZeroInfNanPreserve = true;
3119 }
3120 }
3121}
3122
3123SPIRVModuleAnalysisImpl::SPIRVModuleAnalysisImpl(const SPIRVSubtarget &ST,
3124 SPIRV::ModuleAnalysisInfo &MAI,
3125 MachineFunctionGetter GetMF)
3126 : ST(&ST), GR(ST.getSPIRVGlobalRegistry()), TII(ST.getInstrInfo()),
3127 MAI(MAI), GetMF(GetMF) {}
3128
3129void SPIRVModuleAnalysisImpl::run(const Module &M) {
3130 setBaseInfo(M);
3131
3132 patchPhis(M, GR, TII: *TII, GetMF);
3133
3134 addMBBNames(M, TII: *TII, GetMF, ST: *ST, MAI);
3135 collectFPFastMathDefaults(M, MAI, ST: *ST);
3136 addDecorations(M, TII: *TII, GetMF, ST: *ST, MAI, GR);
3137
3138 collectReqs(M, MAI, GetMF, ST: *ST);
3139
3140 // Process type/const/global var/func decl instructions, number their
3141 // destination registers from 0 to N, collect Extensions and Capabilities.
3142 collectDeclarations(M);
3143
3144 // Number rest of registers from N+1 onwards.
3145 numberRegistersGlobally(M);
3146
3147 // Collect OpName, OpEntryPoint, OpDecorate etc, process other instructions.
3148 processOtherInstrs(M);
3149
3150 // If there are no entry points, we need the Linkage capability.
3151 if (MAI.MS[SPIRV::MB_EntryPoints].empty())
3152 MAI.Reqs.addCapability(ToAdd: SPIRV::Capability::Linkage);
3153
3154 // Set maximum ID used.
3155 GR->setBound(MAI.MaxID);
3156}
3157
3158void SPIRVModuleAnalysisWrapperPass::getAnalysisUsage(AnalysisUsage &AU) const {
3159 AU.addRequired<TargetPassConfig>();
3160 AU.addRequired<MachineModuleInfoWrapperPass>();
3161}
3162
3163bool SPIRVModuleAnalysisWrapperPass::runOnModule(Module &M) {
3164 SPIRVTargetMachine &TM =
3165 getAnalysis<TargetPassConfig>().getTM<SPIRVTargetMachine>();
3166 MachineModuleInfo &MMI = getAnalysis<MachineModuleInfoWrapperPass>().getMMI();
3167 SPIRVModuleAnalysisImpl(
3168 *TM.getSubtargetImpl(), MAI,
3169 [&MMI](const Function &F) { return MMI.getMachineFunction(F); })
3170 .run(M);
3171 return false;
3172}
3173
3174AnalysisKey SPIRVModuleAnalysis::Key;
3175
3176SPIRVModuleAnalysis::Result
3177SPIRVModuleAnalysis::run(Module &M, ModuleAnalysisManager &MAM) {
3178 const auto &TM = static_cast<const SPIRVTargetMachine &>(
3179 MAM.getResult<MachineModuleAnalysis>(IR&: M).getMMI().getTarget());
3180 FunctionAnalysisManager &FAM =
3181 MAM.getResult<FunctionAnalysisManagerModuleProxy>(IR&: M).getManager();
3182 Result MAI;
3183 SPIRVModuleAnalysisImpl(*TM.getSubtargetImpl(), MAI,
3184 [&FAM](const Function &F) -> MachineFunction * {
3185 MachineFunctionAnalysis::Result *MFA =
3186 FAM.getCachedResult<MachineFunctionAnalysis>(
3187 IR&: const_cast<Function &>(F));
3188 assert((MFA || F.isDeclaration()) &&
3189 "Missing MachineFunction for definition");
3190 return MFA ? &MFA->getMF() : nullptr;
3191 })
3192 .run(M);
3193 return MAI;
3194}
3195