1//===-- SPIRVAuxDataHandler.cpp - NonSemantic.AuxData emitter -*- 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#include "SPIRVAuxDataHandler.h"
10#include "MCTargetDesc/SPIRVMCTargetDesc.h"
11#include "SPIRVSubtarget.h"
12#include "SPIRVUtils.h"
13#include "llvm/CodeGen/AsmPrinter.h"
14#include "llvm/IR/Attributes.h"
15#include "llvm/IR/Constants.h"
16#include "llvm/IR/Function.h"
17#include "llvm/IR/GlobalObject.h"
18#include "llvm/IR/GlobalVariable.h"
19#include "llvm/IR/LLVMContext.h"
20#include "llvm/IR/Metadata.h"
21#include "llvm/IR/Module.h"
22#include "llvm/MC/MCInst.h"
23#include "llvm/MC/MCStreamer.h"
24#include "llvm/Support/CommandLine.h"
25#include "llvm/Support/ErrorHandling.h"
26#include "llvm/Support/MathExtras.h"
27
28using namespace llvm;
29
30static cl::opt<bool> SPVPreserveAuxData(
31 "spirv-preserve-auxdata",
32 cl::desc("Preserve LLVM attributes and metadata as "
33 "NonSemantic.AuxData ExtInst annotations (requires "
34 "SPV_KHR_non_semantic_info)"),
35 cl::Optional, cl::Hidden, cl::init(Val: false));
36
37namespace {
38enum AuxDataLinkageType : uint32_t {
39 AvailableExternally = 0,
40};
41
42constexpr unsigned NonSemanticAuxDataSet =
43 static_cast<unsigned>(SPIRV::InstructionSet::NonSemantic_AuxData);
44
45AttributeSet getGOAttrs(const GlobalObject *GO) {
46 if (const auto *F = dyn_cast<Function>(Val: GO))
47 return F->getAttributes().getFnAttrs();
48 return cast<GlobalVariable>(Val: GO)->getAttributes();
49}
50} // namespace
51
52static bool wasAvailableExternally(const GlobalObject *GO) {
53 if (const auto *F = dyn_cast<Function>(Val: GO))
54 return F->hasFnAttribute(SPIRV_WAS_AVAILABLE_EXTERNALLY_ATTR);
55 return cast<GlobalVariable>(Val: GO)->getAttributes().hasAttribute(
56 SPIRV_WAS_AVAILABLE_EXTERNALLY_ATTR);
57}
58
59SPIRVAuxDataHandler::SPIRVAuxDataHandler(AsmPrinter &AP, const Module &M)
60 : Asm(AP), Mod(M) {
61 for (const GlobalObject &GO : M.global_objects())
62 if (wasAvailableExternally(GO: &GO))
63 LinkagePreservedGOs.push_back(Elt: &GO);
64}
65
66bool SPIRVAuxDataHandler::hasWork() const { return SPVPreserveAuxData; }
67
68void SPIRVAuxDataHandler::prepareModuleOutput(const SPIRVSubtarget &ST,
69 SPIRV::ModuleAnalysisInfo &MAI) {
70 if (!hasWork())
71 return;
72 if (!ST.canUseExtension(E: SPIRV::Extension::SPV_KHR_non_semantic_info)) {
73 if (SPVPreserveAuxData)
74 report_fatal_error(reason: "-spirv-preserve-auxdata requires the "
75 "SPV_KHR_non_semantic_info extension to be enabled.");
76 return;
77 }
78 MAI.Reqs.addExtension(ToAdd: SPIRV::Extension::SPV_KHR_non_semantic_info);
79 if (!MAI.ExtInstSetMap.count(Val: NonSemanticAuxDataSet))
80 MAI.ExtInstSetMap[NonSemanticAuxDataSet] = MAI.getNextIDRegister();
81}
82
83MCRegister
84SPIRVAuxDataHandler::getOrEmitString(StringRef S,
85 SPIRV::ModuleAnalysisInfo &MAI) {
86 auto [It, Inserted] = StringRegs.try_emplace(Key: S);
87 if (!Inserted)
88 return It->second;
89 MCRegister Reg = MAI.getNextIDRegister();
90 It->second = Reg;
91 MCInst Inst;
92 Inst.setOpcode(SPIRV::OpString);
93 Inst.addOperand(Op: MCOperand::createReg(Reg));
94 addStringImm(Str: S, Inst);
95 emitMCInst(Inst);
96 return Reg;
97}
98
99void SPIRVAuxDataHandler::collectAttributesFor(const GlobalObject *GO,
100 SPIRV::ModuleAnalysisInfo &MAI) {
101 AuxDataOpcode Opcode = isa<Function>(Val: GO) ? FunctionAttributeOpcode
102 : GlobalVariableAttributeOpcode;
103 for (const Attribute &A : getGOAttrs(GO)) {
104 if (A.isStringAttribute() &&
105 A.getKindAsString() == SPIRV_WAS_AVAILABLE_EXTERNALLY_ATTR)
106 continue;
107 ExtInstRecord Rec;
108 Rec.Opcode = Opcode;
109 Rec.Target = GO;
110 if (A.isStringAttribute()) {
111 Rec.Operands.push_back(Elt: {.Reg: getOrEmitString(S: A.getKindAsString(), MAI)});
112 StringRef Val = A.getValueAsString();
113 if (!Val.empty())
114 Rec.Operands.push_back(Elt: {.Reg: getOrEmitString(S: Val, MAI)});
115 } else {
116 Rec.Operands.push_back(
117 Elt: {.Reg: getOrEmitString(S: StringPool.save(S: A.getAsString()), MAI)});
118 }
119 PendingRecords.push_back(Elt: std::move(Rec));
120 }
121}
122
123void SPIRVAuxDataHandler::collectMetadataFor(const GlobalObject *GO,
124 ArrayRef<StringRef> MDNames,
125 SPIRV::ModuleAnalysisInfo &MAI) {
126 SmallVector<std::pair<unsigned, MDNode *>> AllMD;
127 GO->getAllMetadata(MDs&: AllMD);
128 if (AllMD.empty())
129 return;
130 AuxDataOpcode Opcode =
131 isa<Function>(Val: GO) ? FunctionMetadataOpcode : GlobalVariableMetadataOpcode;
132 // MDString operands become OpStrings; ValueAsMetadata constants (e.g.
133 // !{i32 5}) become OpConstants emitted at section 10. Any other operand
134 // kind would need full value translation, so skip the whole node.
135 auto CollectOperands =
136 [&](MDNode *MD) -> std::optional<SmallVector<Operand, 4>> {
137 SmallVector<Operand, 4> Out;
138 for (const MDOperand &MdOp : MD->operands()) {
139 Metadata *Md = MdOp.get();
140 if (auto *MDStr = dyn_cast_or_null<MDString>(Val: Md)) {
141 Out.push_back(Elt: {.Reg: getOrEmitString(S: MDStr->getString(), MAI)});
142 } else if (auto *VAM = dyn_cast_or_null<ValueAsMetadata>(Val: Md)) {
143 auto *C = dyn_cast<Constant>(Val: VAM->getValue());
144 if (!C || !(isa<ConstantInt>(Val: C) || isa<ConstantFP>(Val: C)))
145 return std::nullopt;
146 Out.push_back(Elt: {.Reg: MCRegister(), .Const: C});
147 } else {
148 return std::nullopt;
149 }
150 }
151 return Out;
152 };
153 for (const auto &MD : AllMD) {
154 if (MD.first == LLVMContext::MD_dbg)
155 continue;
156 StringRef MDName = MDNames[MD.first];
157 if (MDName == "spirv.Decorations" || MDName == "spirv.ParameterDecorations")
158 continue;
159 auto Operands = CollectOperands(MD.second);
160 if (!Operands)
161 continue;
162 ExtInstRecord Rec;
163 Rec.Opcode = Opcode;
164 Rec.Target = GO;
165 Rec.Operands.push_back(Elt: {.Reg: getOrEmitString(S: MDName, MAI)});
166 Rec.Operands.append(in_start: Operands->begin(), in_end: Operands->end());
167 PendingRecords.push_back(Elt: std::move(Rec));
168 }
169}
170
171void SPIRVAuxDataHandler::emitAuxDataStrings(SPIRV::ModuleAnalysisInfo &MAI) {
172 if (!SPVPreserveAuxData)
173 return;
174 if (!MAI.getExtInstSetReg(SetNum: NonSemanticAuxDataSet).isValid())
175 return;
176 SmallVector<StringRef> MDNames;
177 Mod.getContext().getMDKindNames(Result&: MDNames);
178 for (const GlobalObject &GO : Mod.global_objects()) {
179 if (GO.isDeclaration())
180 continue;
181 collectAttributesFor(GO: &GO, MAI);
182 collectMetadataFor(GO: &GO, MDNames, MAI);
183 }
184}
185
186void SPIRVAuxDataHandler::emitAuxData(SPIRV::ModuleAnalysisInfo &MAI) {
187 MCRegister ExtSetReg = MAI.getExtInstSetReg(SetNum: NonSemanticAuxDataSet);
188 if (!ExtSetReg.isValid())
189 return;
190
191 MCRegister VoidTypeReg = findOrEmitOpTypeVoid(MAI);
192
193 for (const ExtInstRecord &Rec : PendingRecords) {
194 MCRegister TargetReg = MAI.getGlobalObjReg(GO: Rec.Target);
195 if (!TargetReg.isValid())
196 continue;
197 SmallVector<MCRegister, 5> Operands;
198 Operands.push_back(Elt: TargetReg);
199 for (const Operand &Op : Rec.Operands)
200 Operands.push_back(Elt: Op.Const ? emitConstant(C: Op.Const, MAI) : Op.Reg);
201 emitAuxDataExtInst(Opcode: Rec.Opcode, VoidTypeReg, ExtSetReg, Operands, MAI);
202 }
203
204 if (LinkagePreservedGOs.empty())
205 return;
206
207 MCRegister UInt32TypeReg = findOrEmitOpTypeUInt32(MAI);
208 MCRegister AEConstReg;
209 for (const GlobalObject *GO : LinkagePreservedGOs) {
210 MCRegister TargetReg = MAI.getGlobalObjReg(GO);
211 if (!TargetReg.isValid())
212 continue;
213 if (!AEConstReg.isValid())
214 AEConstReg =
215 emitOpConstantUInt32(Value: AvailableExternally, UInt32TypeReg, MAI);
216 emitAuxDataExtInst(Opcode: LinkageOpcode, VoidTypeReg, ExtSetReg,
217 Operands: {TargetReg, AEConstReg}, MAI);
218 }
219}
220
221void SPIRVAuxDataHandler::emitAuxDataExtInst(AuxDataOpcode Opcode,
222 MCRegister VoidTypeReg,
223 MCRegister ExtSetReg,
224 ArrayRef<MCRegister> Operands,
225 SPIRV::ModuleAnalysisInfo &MAI) {
226 MCInst Inst;
227 Inst.setOpcode(SPIRV::OpExtInst);
228 Inst.addOperand(Op: MCOperand::createReg(Reg: MAI.getNextIDRegister()));
229 Inst.addOperand(Op: MCOperand::createReg(Reg: VoidTypeReg));
230 Inst.addOperand(Op: MCOperand::createReg(Reg: ExtSetReg));
231 Inst.addOperand(Op: MCOperand::createImm(Val: Opcode));
232 for (MCRegister R : Operands)
233 Inst.addOperand(Op: MCOperand::createReg(Reg: R));
234 emitMCInst(Inst);
235}
236
237void SPIRVAuxDataHandler::emitMCInst(MCInst &Inst) {
238 Asm.OutStreamer->emitInstruction(Inst, STI: Asm.getSubtargetInfo());
239}
240
241MCRegister
242SPIRVAuxDataHandler::findOrEmitOpTypeVoid(SPIRV::ModuleAnalysisInfo &MAI) {
243 for (const MachineInstr *MI : MAI.getMSInstrs(MSType: SPIRV::MB_TypeConstVars))
244 if (MI->getOpcode() == SPIRV::OpTypeVoid)
245 return MAI.getRegisterAlias(MF: MI->getMF(), Reg: MI->getOperand(i: 0).getReg());
246 MCRegister Reg = MAI.getNextIDRegister();
247 MCInst Inst;
248 Inst.setOpcode(SPIRV::OpTypeVoid);
249 Inst.addOperand(Op: MCOperand::createReg(Reg));
250 emitMCInst(Inst);
251 return Reg;
252}
253
254MCRegister
255SPIRVAuxDataHandler::findOrEmitOpTypeInt(unsigned BitWidth,
256 SPIRV::ModuleAnalysisInfo &MAI) {
257 // SPIR-V OpTypeInt: <width>, <signedness>. Signedness 0 = unsigned, 1 =
258 // signed; we always emit unsigned.
259 constexpr int64_t UnsignedSignedness = 0;
260 for (const MachineInstr *MI : MAI.getMSInstrs(MSType: SPIRV::MB_TypeConstVars))
261 if (MI->getOpcode() == SPIRV::OpTypeInt &&
262 MI->getOperand(i: 1).getImm() == static_cast<int64_t>(BitWidth) &&
263 MI->getOperand(i: 2).getImm() == UnsignedSignedness)
264 return MAI.getRegisterAlias(MF: MI->getMF(), Reg: MI->getOperand(i: 0).getReg());
265 MCRegister Reg = MAI.getNextIDRegister();
266 MCInst Inst;
267 Inst.setOpcode(SPIRV::OpTypeInt);
268 Inst.addOperand(Op: MCOperand::createReg(Reg));
269 Inst.addOperand(Op: MCOperand::createImm(Val: BitWidth));
270 Inst.addOperand(Op: MCOperand::createImm(Val: UnsignedSignedness));
271 emitMCInst(Inst);
272 return Reg;
273}
274
275MCRegister
276SPIRVAuxDataHandler::findOrEmitOpTypeUInt32(SPIRV::ModuleAnalysisInfo &MAI) {
277 return findOrEmitOpTypeInt(BitWidth: 32, MAI);
278}
279
280MCRegister
281SPIRVAuxDataHandler::findOrEmitOpTypeFloat(unsigned BitWidth,
282 SPIRV::ModuleAnalysisInfo &MAI) {
283 for (const MachineInstr *MI : MAI.getMSInstrs(MSType: SPIRV::MB_TypeConstVars))
284 if (MI->getOpcode() == SPIRV::OpTypeFloat &&
285 MI->getOperand(i: 1).getImm() == static_cast<int64_t>(BitWidth))
286 return MAI.getRegisterAlias(MF: MI->getMF(), Reg: MI->getOperand(i: 0).getReg());
287 MCRegister Reg = MAI.getNextIDRegister();
288 MCInst Inst;
289 Inst.setOpcode(SPIRV::OpTypeFloat);
290 Inst.addOperand(Op: MCOperand::createReg(Reg));
291 Inst.addOperand(Op: MCOperand::createImm(Val: BitWidth));
292 emitMCInst(Inst);
293 return Reg;
294}
295
296MCRegister SPIRVAuxDataHandler::emitConstant(const Constant *C,
297 SPIRV::ModuleAnalysisInfo &MAI) {
298 auto [It, Inserted] = ConstantRegs.try_emplace(Key: C);
299 if (!Inserted)
300 return It->second;
301
302 APInt Bits;
303 unsigned Opcode;
304 MCRegister TypeReg;
305 if (const auto *CI = dyn_cast<ConstantInt>(Val: C)) {
306 Bits = CI->getValue();
307 Opcode = SPIRV::OpConstantI;
308 TypeReg = findOrEmitOpTypeInt(BitWidth: Bits.getBitWidth(), MAI);
309 } else {
310 const auto *CF = cast<ConstantFP>(Val: C);
311 Bits = CF->getValueAPF().bitcastToAPInt();
312 Opcode = SPIRV::OpConstantF;
313 TypeReg = findOrEmitOpTypeFloat(BitWidth: Bits.getBitWidth(), MAI);
314 }
315
316 MCRegister Reg = MAI.getNextIDRegister();
317 It->second = Reg;
318 MCInst Inst;
319 Inst.setOpcode(Opcode);
320 Inst.addOperand(Op: MCOperand::createReg(Reg));
321 Inst.addOperand(Op: MCOperand::createReg(Reg: TypeReg));
322 // SPIR-V encodes the literal as ceil(width/32) little-endian 32-bit words.
323 unsigned NumWords = divideCeil(Numerator: Bits.getBitWidth(), Denominator: 32);
324 for (unsigned I = 0; I < NumWords; ++I)
325 Inst.addOperand(Op: MCOperand::createImm(Val: Bits.extractBitsAsZExtValue(
326 numBits: std::min(a: 32u, b: Bits.getBitWidth() - I * 32), bitPosition: I * 32)));
327 // The asm printer needs this hint to render an f16 literal correctly.
328 if (Opcode == SPIRV::OpConstantF && Bits.getBitWidth() == 16)
329 Inst.setFlags(SPIRV::INST_PRINTER_WIDTH16);
330 emitMCInst(Inst);
331 return Reg;
332}
333
334MCRegister SPIRVAuxDataHandler::emitOpConstantUInt32(
335 uint32_t Value, MCRegister UInt32TypeReg, SPIRV::ModuleAnalysisInfo &MAI) {
336 MCRegister Reg = MAI.getNextIDRegister();
337 MCInst Inst;
338 Inst.setOpcode(SPIRV::OpConstantI);
339 Inst.addOperand(Op: MCOperand::createReg(Reg));
340 Inst.addOperand(Op: MCOperand::createReg(Reg: UInt32TypeReg));
341 Inst.addOperand(Op: MCOperand::createImm(Val: static_cast<int64_t>(Value)));
342 emitMCInst(Inst);
343 return Reg;
344}
345