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