1//===-- NVPTXAsmPrinter.cpp - NVPTX LLVM assembly writer ------------------===//
2//
3// Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4// See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9// This file contains a printer that converts from our internal representation
10// of machine-dependent LLVM code to NVPTX assembly language.
11//
12//===----------------------------------------------------------------------===//
13
14#include "NVPTXAsmPrinter.h"
15#include "MCTargetDesc/NVPTXBaseInfo.h"
16#include "MCTargetDesc/NVPTXInstPrinter.h"
17#include "MCTargetDesc/NVPTXTargetStreamer.h"
18#include "NVPTX.h"
19#include "NVPTXDwarfDebug.h"
20#include "NVPTXMCExpr.h"
21#include "NVPTXMachineFunctionInfo.h"
22#include "NVPTXRegisterInfo.h"
23#include "NVPTXSubtarget.h"
24#include "NVPTXTargetMachine.h"
25#include "NVPTXUtilities.h"
26#include "NVVMProperties.h"
27#include "TargetInfo/NVPTXTargetInfo.h"
28#include "cl_common_defines.h"
29#include "llvm/ADT/APFloat.h"
30#include "llvm/ADT/APInt.h"
31#include "llvm/ADT/ArrayRef.h"
32#include "llvm/ADT/DenseMap.h"
33#include "llvm/ADT/DenseSet.h"
34#include "llvm/ADT/SCCIterator.h"
35#include "llvm/ADT/STLExtras.h"
36#include "llvm/ADT/Sequence.h"
37#include "llvm/ADT/SmallPtrSet.h"
38#include "llvm/ADT/SmallString.h"
39#include "llvm/ADT/SmallVector.h"
40#include "llvm/ADT/StringExtras.h"
41#include "llvm/ADT/StringRef.h"
42#include "llvm/ADT/Twine.h"
43#include "llvm/ADT/iterator_range.h"
44#include "llvm/Analysis/ConstantFolding.h"
45#include "llvm/CodeGen/Analysis.h"
46#include "llvm/CodeGen/AsmPrinter.h"
47#include "llvm/CodeGen/AsmPrinterAnalysis.h"
48#include "llvm/CodeGen/MachineBasicBlock.h"
49#include "llvm/CodeGen/MachineFrameInfo.h"
50#include "llvm/CodeGen/MachineFunction.h"
51#include "llvm/CodeGen/MachineInstr.h"
52#include "llvm/CodeGen/MachineJumpTableInfo.h"
53#include "llvm/CodeGen/MachineLoopInfo.h"
54#include "llvm/CodeGen/MachineModuleInfo.h"
55#include "llvm/CodeGen/MachineOperand.h"
56#include "llvm/CodeGen/MachineRegisterInfo.h"
57#include "llvm/CodeGen/TargetRegisterInfo.h"
58#include "llvm/CodeGen/ValueTypes.h"
59#include "llvm/CodeGenTypes/MachineValueType.h"
60#include "llvm/IR/Argument.h"
61#include "llvm/IR/Attributes.h"
62#include "llvm/IR/BasicBlock.h"
63#include "llvm/IR/Constant.h"
64#include "llvm/IR/Constants.h"
65#include "llvm/IR/DataLayout.h"
66#include "llvm/IR/DebugInfo.h"
67#include "llvm/IR/DebugInfoMetadata.h"
68#include "llvm/IR/DebugLoc.h"
69#include "llvm/IR/DerivedTypes.h"
70#include "llvm/IR/Function.h"
71#include "llvm/IR/GlobalAlias.h"
72#include "llvm/IR/GlobalValue.h"
73#include "llvm/IR/GlobalVariable.h"
74#include "llvm/IR/InstrTypes.h"
75#include "llvm/IR/Instruction.h"
76#include "llvm/IR/LLVMContext.h"
77#include "llvm/IR/Module.h"
78#include "llvm/IR/Operator.h"
79#include "llvm/IR/Type.h"
80#include "llvm/IR/User.h"
81#include "llvm/IR/Value.h"
82#include "llvm/MC/MCExpr.h"
83#include "llvm/MC/MCInst.h"
84#include "llvm/MC/MCInstrDesc.h"
85#include "llvm/MC/MCStreamer.h"
86#include "llvm/MC/MCSymbol.h"
87#include "llvm/MC/TargetRegistry.h"
88#include "llvm/Pass.h"
89#include "llvm/Support/Alignment.h"
90#include "llvm/Support/Casting.h"
91#include "llvm/Support/Compiler.h"
92#include "llvm/Support/Endian.h"
93#include "llvm/Support/ErrorHandling.h"
94#include "llvm/Support/NativeFormatting.h"
95#include "llvm/Support/raw_ostream.h"
96#include "llvm/Target/TargetLoweringObjectFile.h"
97#include "llvm/Target/TargetMachine.h"
98#include "llvm/Transforms/Utils/UnrollLoop.h"
99#include <algorithm>
100#include <cassert>
101#include <cstdint>
102#include <cstring>
103#include <map>
104#include <memory>
105#include <set>
106#include <string>
107#include <type_traits>
108#include <vector>
109
110using namespace llvm;
111
112#define DEPOTNAME "__local_depot"
113
114// The ptx syntax and format is very different from that usually seem in a .s
115// file,
116// therefore we are not able to use the MCAsmStreamer interface here.
117//
118// We are handcrafting the output method here.
119//
120// A better approach is to clone the MCAsmStreamer to a MCPTXAsmStreamer
121// (subclass of MCStreamer).
122
123namespace {
124
125class NVPTXAsmPrinter : public AsmPrinter {
126
127 class AggBuffer {
128 // Used to buffer the emitted string for initializing global aggregates.
129 //
130 // Normally an aggregate (array, vector, or structure) is emitted as a u8[].
131 // However, if either element/field of the aggregate is a non-NULL address,
132 // and all such addresses are properly aligned, then the aggregate is
133 // emitted as u32[] or u64[]. In the case of unaligned addresses, the
134 // aggregate is emitted as u8[], and the mask() operator is used for all
135 // pointers.
136 //
137 // We first layout the aggregate in 'buffer' in bytes, except for those
138 // symbol addresses. For the i-th symbol address in the aggregate, its
139 // corresponding 4-byte or 8-byte elements in 'buffer' are filled with 0s.
140 // symbolPosInBuffer[i-1] records its position in 'buffer', and Symbols[i-1]
141 // records the Value*.
142 //
143 // Once we have this AggBuffer setup, we can choose how to print it out.
144 public:
145 // number of symbol addresses
146 unsigned numSymbols() const { return Symbols.size(); }
147
148 bool allSymbolsAligned(unsigned ptrSize) const {
149 return llvm::all_of(Range: symbolPosInBuffer,
150 P: [=](unsigned pos) { return pos % ptrSize == 0; });
151 }
152
153 private:
154 const unsigned Size; // size of the buffer in bytes
155 std::vector<unsigned char> buffer; // the buffer
156 SmallVector<unsigned, 4> symbolPosInBuffer;
157 SmallVector<const Value *, 4> Symbols;
158 // SymbolsBeforeStripping[i] is the original form of Symbols[i] before
159 // stripping pointer casts, i.e.,
160 // Symbols[i] == SymbolsBeforeStripping[i]->stripPointerCasts().
161 //
162 // We need to keep these values because AggBuffer::print decides whether to
163 // emit a "generic()" cast for Symbols[i] depending on the address space of
164 // SymbolsBeforeStripping[i].
165 SmallVector<const Value *, 4> SymbolsBeforeStripping;
166 unsigned curpos;
167 const NVPTXAsmPrinter &AP;
168 const bool EmitGeneric;
169
170 public:
171 AggBuffer(unsigned Size, const NVPTXAsmPrinter &AP)
172 : Size(Size), buffer(Size), curpos(0), AP(AP),
173 EmitGeneric(AP.EmitGeneric) {}
174
175 unsigned getBufferSize() const { return Size; }
176
177 // Number of bytes written so far.
178 unsigned getCurpos() const { return curpos; }
179
180 // Copy Num bytes from Ptr.
181 // if Bytes > Num, zero fill up to Bytes.
182 void addBytes(const unsigned char *Ptr, unsigned Num, unsigned Bytes) {
183 for (unsigned I : llvm::seq(Size: Num))
184 addByte(Byte: Ptr[I]);
185 if (Bytes > Num)
186 addZeros(Num: Bytes - Num);
187 }
188
189 void addByte(uint8_t Byte) {
190 assert(curpos < Size);
191 buffer[curpos] = Byte;
192 curpos++;
193 }
194
195 void addZeros(unsigned Num) {
196 for ([[maybe_unused]] unsigned _ : llvm::seq(Size: Num)) {
197 addByte(Byte: 0);
198 }
199 }
200
201 void addSymbol(const Value *GVar, const Value *GVarBeforeStripping) {
202 symbolPosInBuffer.push_back(Elt: curpos);
203 Symbols.push_back(Elt: GVar);
204 SymbolsBeforeStripping.push_back(Elt: GVarBeforeStripping);
205 }
206
207 void printBytes(raw_ostream &os);
208 void printWords(raw_ostream &os);
209
210 private:
211 void printSymbol(unsigned nSym, raw_ostream &os);
212 };
213
214 friend class AggBuffer;
215
216public:
217 static char ID;
218
219 StringRef getPassName() const override { return "NVPTX Assembly Printer"; }
220
221private:
222 const Function *F;
223
224 NVPTXTargetStreamer *getTargetStreamer() const;
225
226 void emitStartOfAsmFile(Module &M) override;
227 void emitBasicBlockStart(const MachineBasicBlock &MBB) override;
228 void emitFunctionEntryLabel() override;
229 void emitFunctionBodyStart() override;
230 void emitFunctionBodyEnd() override;
231 void emitImplicitDef(const MachineInstr *MI) const override;
232
233 void emitInstruction(const MachineInstr *) override;
234 void lowerToMCInst(const MachineInstr *MI, MCInst &OutMI);
235 MCOperand lowerOperand(const MachineOperand &MO);
236 MCOperand GetSymbolRef(const MCSymbol *Symbol);
237 MCRegister encodeVirtualRegister(Register Reg);
238
239 /// The number \p Reg was assigned within its register class, as declared by
240 /// this function's .reg directives.
241 unsigned getVirtualRegisterNumber(Register Reg) const;
242
243 void printMemOperand(const MachineInstr *MI, unsigned OpNum, raw_ostream &O,
244 const char *Modifier = nullptr);
245 void printModuleLevelGV(const GlobalVariable *GVar, raw_ostream &O,
246 bool processDemoted, const NVPTXSubtarget &STI);
247 void emitGlobals(const Module &M);
248 void emitGlobalAlias(const Module &M, const GlobalAlias &GA) override;
249 void emitHeader(Module &M, const NVPTXSubtarget &STI);
250 void emitKernelFunctionDirectives(const Function &F, raw_ostream &O) const;
251 void emitFunctionParamList(const Function *, raw_ostream &O);
252 void setAndEmitFunctionVirtualRegisters(const MachineFunction &MF);
253 void encodeDebugInfoRegisterNumbers(const MachineFunction &MF);
254 void emitCallPrototype(const CallBase &CB, MCSymbol *PrototypeSymbol) const;
255 void emitJumpTable(const MachineJumpTableEntry &MJT, unsigned MJTI) const;
256
257 /// Should a .noreturn directive be emitted for \p V, which is either a
258 /// function or a call site?
259 template <typename T> bool shouldEmitPTXNoReturn(const T &V) const {
260 static_assert(std::is_same_v<Function, T> || std::is_base_of_v<CallBase, T>,
261 "expected a function or a call site");
262
263 const auto &NTM = static_cast<const NVPTXTargetMachine &>(TM);
264 if (!NTM.getSubtargetImpl()->hasNoReturn())
265 return false;
266
267 if (!V.doesNotReturn() || !V.getFunctionType()->getReturnType()->isVoidTy())
268 return false;
269
270 if constexpr (std::is_same_v<Function, T>)
271 return !isKernelFunction(V);
272 else
273 return true;
274 }
275
276 bool PrintAsmOperand(const MachineInstr *MI, unsigned OpNo,
277 const char *ExtraCode, raw_ostream &) override;
278 void printOperand(const MachineInstr *MI, unsigned OpNum, raw_ostream &O);
279 bool PrintAsmMemoryOperand(const MachineInstr *MI, unsigned OpNo,
280 const char *ExtraCode, raw_ostream &) override;
281
282 const MCExpr *lowerConstantForGV(const Constant *CV,
283 bool ProcessingGeneric) const;
284 void printMCExpr(const MCExpr &Expr, raw_ostream &OS) const;
285 /// Emit a blob of inline asm to the output streamer.
286 void emitInlineAsm(StringRef Str, const MCSubtargetInfo &STI,
287 const MCTargetOptions &MCOptions, const MDNode *LocMDNode,
288 InlineAsm::AsmDialect Dialect,
289 const MachineInstr *MI) override;
290
291protected:
292 bool doInitialization(Module &M) override;
293 bool doFinalization(Module &M) override;
294
295 /// Create NVPTX-specific DwarfDebug handler.
296 DwarfDebug *createDwarfDebug() override;
297
298private:
299 bool GlobalsEmitted;
300
301 // This is specific per MachineFunction.
302 const MachineRegisterInfo *MRI;
303
304 // The number assigned to each virtual register within its class, populated
305 // by setAndEmitFunctionVirtualRegisters and cleared between functions.
306 using VRegMap = DenseMap<Register, unsigned>;
307 using VRegRCMap = DenseMap<const TargetRegisterClass *, VRegMap>;
308 VRegRCMap VRegMapping;
309
310 // List of variables demoted to a function scope.
311 std::map<const Function *, std::vector<const GlobalVariable *>> localDecls;
312
313 /// Print the state space, alignment, type, name, and — when
314 /// \p EmitInitializer is set — the initializer of \p GVar. Passing false
315 /// prints a declaration whose type still matches the definition, as an
316 /// `.extern` forward declaration requires.
317 void emitPTXGlobalVariableDefinition(const GlobalVariable *GVar,
318 raw_ostream &O,
319 const NVPTXSubtarget &STI,
320 bool EmitInitializer);
321 void emitPTXAddressSpace(unsigned int AddressSpace, raw_ostream &O) const;
322 std::string getPTXFundamentalTypeStr(Type *Ty) const;
323 void printScalarConstant(const Constant *CPV, raw_ostream &O);
324 void printFPConstant(const ConstantFP *Fp, raw_ostream &O) const;
325 void bufferLEByte(const Constant *CPV, int Bytes, AggBuffer *aggBuffer);
326 void bufferAggregateConstant(const Constant *CV, AggBuffer *aggBuffer);
327 void bufferAggregateConstVec(const ConstantVector *CV, AggBuffer *aggBuffer);
328
329 void emitLinkageDirective(const GlobalValue *V, raw_ostream &O);
330 void emitDeclarations(const Module &, raw_ostream &O);
331 void emitDeclaration(const Function *, raw_ostream &O);
332 void emitAliasDeclaration(const GlobalAlias *, raw_ostream &O);
333 void emitDeclarationWithName(const Function *, MCSymbol *, raw_ostream &O);
334 void emitDemotedVars(const Function *, raw_ostream &);
335
336 bool isLoopHeaderOfNoUnroll(const MachineBasicBlock &MBB) const;
337
338 // Used to control the need to emit .generic() in the initializer of
339 // module scope variables.
340 // Although ptx supports the hybrid mode like the following,
341 // .global .u32 a;
342 // .global .u32 b;
343 // .global .u32 addr[] = {a, generic(b)}
344 // we have difficulty representing the difference in the NVVM IR.
345 //
346 // Since the address value should always be generic in CUDA C and always
347 // be specific in OpenCL, we use this simple control here.
348 //
349 const bool EmitGeneric;
350
351public:
352 NVPTXAsmPrinter(TargetMachine &TM, std::unique_ptr<MCStreamer> Streamer)
353 : AsmPrinter(TM, std::move(Streamer), ID),
354 EmitGeneric(static_cast<NVPTXTargetMachine &>(TM).getDrvInterface() ==
355 NVPTX::CUDA) {}
356
357 bool runOnMachineFunction(MachineFunction &F) override;
358
359 void getAnalysisUsage(AnalysisUsage &AU) const override {
360 AU.addRequired<MachineLoopInfoWrapperPass>();
361 AsmPrinter::getAnalysisUsage(AU);
362 }
363
364 std::string getVirtualRegisterName(Register Reg) const;
365
366 const MCSymbol *getFunctionFrameSymbol() const override;
367
368 // Make emitGlobalVariable() no-op for NVPTX.
369 // Global variables have been already emitted by the time the base AsmPrinter
370 // attempts to do so in doFinalization() (see NVPTXAsmPrinter::emitGlobals()).
371 void emitGlobalVariable(const GlobalVariable *GV) override {}
372};
373
374} // end anonymous namespace
375
376/// Emits initial debug location directive.
377static void emitInitialRawDwarfLocDirective(const MachineFunction &MF,
378 DwarfDebug *DD,
379 MCStreamer &OutStreamer) {
380 if (!DD)
381 return;
382
383 assert(OutStreamer.hasRawTextSupport() && "Expected assembly output mode.");
384 // This is NVPTX specific and it's unclear why.
385 // PR51079: If we have code without debug information we need to give up.
386 const DISubprogram *SP = MF.getFunction().getSubprogram();
387 if (!SP)
388 return;
389 assert(SP->getUnit());
390 // NoDebug and DebugDirectivesOnly do not require emitting the initial loc
391 // directive. NoDebug does not require any debug directives and the initial
392 // loc directive is not needed for DebugDirectivesOnly as it is redundant
393 // assuming this is a non-empty function.
394 if (SP->getUnit()->isDebugDirectivesOnly() || SP->getUnit()->isNoDebug())
395 return;
396
397 (void)DD->emitInitialLocDirective(MF, /*CUID=*/0);
398}
399
400namespace {
401
402/// Return a list of GlobalVariables on which \p V depends.
403static void
404discoverDependentGlobals(const Value *V,
405 SmallVectorImpl<const GlobalVariable *> &Globals,
406 SmallPtrSetImpl<const GlobalVariable *> &Seen) {
407 if (const GlobalVariable *GV = dyn_cast<GlobalVariable>(Val: V)) {
408 if (Seen.insert(Ptr: GV).second)
409 Globals.push_back(Elt: GV);
410 return;
411 }
412
413 // Global values are emitted as symbols. Their operands do not contribute to
414 // the initializer expression that refers to that symbol.
415 if (isa<GlobalValue>(Val: V))
416 return;
417
418 // lowerConstantForGV emits a GEP as its base symbol plus a constant byte
419 // offset. Symbols used to compute an index are not part of that expression.
420 if (const GEPOperator *GEP = dyn_cast<GEPOperator>(Val: V)) {
421 discoverDependentGlobals(V: GEP->getPointerOperand(), Globals, Seen);
422 return;
423 }
424
425 if (const User *U = dyn_cast<User>(Val: V))
426 for (const auto &O : U->operands())
427 discoverDependentGlobals(V: O, Globals, Seen);
428}
429
430struct GlobalVariableDependencyNode {
431 const GlobalVariable *GV = nullptr;
432 unsigned ModuleOrder = 0;
433 SmallVector<const GlobalVariableDependencyNode *, 4> Dependencies;
434};
435
436class GlobalVariableDependencyGraph {
437 // scc_iterator needs a single entry node. Global initializer dependencies
438 // may be disconnected, so use a synthetic root with an edge to every global.
439 GlobalVariableDependencyNode SyntheticRoot;
440 // Edges store pointers into Nodes, so node addresses must remain stable while
441 // the graph is constructed.
442 std::map<const GlobalVariable *, GlobalVariableDependencyNode> Nodes;
443
444public:
445 explicit GlobalVariableDependencyGraph(const Module &M) {
446 unsigned ModuleOrder = 0;
447 for (const GlobalVariable &GV : M.globals()) {
448 GlobalVariableDependencyNode &Node = Nodes.try_emplace(k: &GV).first->second;
449 Node.GV = &GV;
450 Node.ModuleOrder = ModuleOrder++;
451 SyntheticRoot.Dependencies.push_back(Elt: &Node);
452 }
453
454 for (auto &[GV, Node] : Nodes) {
455 SmallVector<const GlobalVariable *, 4> Dependencies;
456 SmallPtrSet<const GlobalVariable *, 4> Seen;
457 for (const Use &Operand : GV->operands())
458 discoverDependentGlobals(V: Operand, Globals&: Dependencies, Seen);
459
460 for (const GlobalVariable *Dependency : Dependencies) {
461 auto It = Nodes.find(x: Dependency);
462 if (It != Nodes.end())
463 Node.Dependencies.push_back(Elt: &It->second);
464 }
465 }
466 }
467
468 const GlobalVariableDependencyNode *getEntryNode() const {
469 return &SyntheticRoot;
470 }
471};
472
473struct GlobalVariableDependencyGraphTraits {
474 using NodeRef = const GlobalVariableDependencyNode *;
475 using ChildIteratorType =
476 SmallVectorImpl<const GlobalVariableDependencyNode *>::const_iterator;
477
478 static NodeRef getEntryNode(NodeRef Node) { return Node; }
479 static ChildIteratorType child_begin(NodeRef Node) {
480 return Node->Dependencies.begin();
481 }
482 static ChildIteratorType child_end(NodeRef Node) {
483 return Node->Dependencies.end();
484 }
485};
486
487using GlobalVariableSCCIterator =
488 scc_iterator<const GlobalVariableDependencyNode *,
489 GlobalVariableDependencyGraphTraits>;
490
491static bool shouldSkipModuleLevelGlobal(const GlobalVariable &GV) {
492 if (GV.hasSection() && GV.getSection() == "llvm.metadata")
493 return true;
494 return GV.getName().starts_with(Prefix: "llvm.") || GV.getName().starts_with(Prefix: "nvvm.");
495}
496
497static bool isForwardDeclarableGlobal(const GlobalVariable *GVar) {
498 if (shouldSkipModuleLevelGlobal(GV: *GVar) || GVar->isDeclaration() ||
499 getPTXOpaqueType(*GVar) != PTXOpaqueType::None)
500 return false;
501
502 // A PTX .extern declaration can be resolved by a later .visible, .weak, or
503 // .common definition, but not by a static definition.
504 if (GVar->hasExternalLinkage())
505 return GVar->hasInitializer();
506
507 if (GVar->hasLinkOnceLinkage() || GVar->hasWeakLinkage() ||
508 GVar->hasAvailableExternallyLinkage() || GVar->hasCommonLinkage())
509 return true;
510
511 return false;
512}
513
514/// Order definitions after treating references to forward-declared globals as
515/// already satisfied. A remaining cycle cannot be emitted portably because it
516/// requires an undeclared forward reference.
517static SmallVector<const GlobalVariable *, 4> orderDefinitionsInSCC(
518 ArrayRef<const GlobalVariableDependencyNode *> SCC,
519 const DenseSet<const GlobalVariableDependencyNode *> &ForwardDeclared) {
520 using Node = GlobalVariableDependencyNode;
521
522 DenseSet<const Node *> SCCSet;
523 SCCSet.insert_range(R&: SCC);
524
525 DenseMap<const Node *, unsigned> DependencyCount;
526 DenseMap<const Node *, SmallVector<const Node *, 4>> Dependents;
527 std::set<std::pair<unsigned, const Node *>> Ready;
528
529 // Dependencies outside this SCC have already been emitted. Forward-declared
530 // dependencies are also satisfied, so only count the remaining SCC edges.
531 for (const Node *N : SCC) {
532 unsigned &Count = DependencyCount[N];
533 for (const Node *Dependency : N->Dependencies) {
534 if (!SCCSet.count(V: Dependency) || ForwardDeclared.count(V: Dependency))
535 continue;
536 ++Count;
537 Dependents[Dependency].push_back(Elt: N);
538 }
539 if (Count == 0)
540 Ready.emplace(args: N->ModuleOrder, args&: N);
541 }
542
543 SmallVector<const GlobalVariable *, 4> Order;
544 while (!Ready.empty()) {
545 const Node *N = Ready.begin()->second;
546 Ready.erase(position: Ready.begin());
547 Order.push_back(Elt: N->GV);
548
549 auto It = Dependents.find(Val: N);
550 if (It == Dependents.end())
551 continue;
552 for (const Node *Dependent : It->second) {
553 assert(DependencyCount[Dependent] && "Dependency already satisfied");
554 if (--DependencyCount[Dependent] == 0)
555 Ready.emplace(args: Dependent->ModuleOrder, args&: Dependent);
556 }
557 }
558
559 if (Order.size() != SCC.size())
560 report_fatal_error(reason: "Circular dependency found in global variable set");
561 return Order;
562}
563
564} // namespace
565
566void NVPTXAsmPrinter::emitInstruction(const MachineInstr *MI) {
567 NVPTX_MC::verifyInstructionPredicates(Opcode: MI->getOpcode(),
568 Features: getSubtargetInfo().getFeatureBits());
569
570 MCInst Inst;
571 lowerToMCInst(MI, OutMI&: Inst);
572 EmitToStreamer(S&: *OutStreamer, Inst);
573}
574
575void NVPTXAsmPrinter::lowerToMCInst(const MachineInstr *MI, MCInst &OutMI) {
576 OutMI.setOpcode(MI->getOpcode());
577 for (const auto MO : MI->operands())
578 OutMI.addOperand(Op: lowerOperand(MO));
579}
580
581MCOperand NVPTXAsmPrinter::lowerOperand(const MachineOperand &MO) {
582 switch (MO.getType()) {
583 default:
584 llvm_unreachable("unknown operand type");
585 case MachineOperand::MO_Register:
586 return MCOperand::createReg(Reg: encodeVirtualRegister(Reg: MO.getReg()));
587 case MachineOperand::MO_Immediate:
588 return MCOperand::createImm(Val: MO.getImm());
589 case MachineOperand::MO_MachineBasicBlock:
590 return MCOperand::createExpr(
591 Val: MCSymbolRefExpr::create(Symbol: MO.getMBB()->getSymbol(), Ctx&: OutContext));
592 case MachineOperand::MO_ExternalSymbol:
593 return GetSymbolRef(Symbol: GetExternalSymbolSymbol(Sym: MO.getSymbolName()));
594 case MachineOperand::MO_MCSymbol:
595 return GetSymbolRef(Symbol: MO.getMCSymbol());
596 case MachineOperand::MO_JumpTableIndex:
597 // The jump table index names the .branchtargets list emitted for a brx.idx
598 // (see emitJumpTable); reference it by that label.
599 return GetSymbolRef(Symbol: GetJTISymbol(JTID: MO.getIndex()));
600 case MachineOperand::MO_GlobalAddress:
601 return GetSymbolRef(Symbol: getSymbol(GV: MO.getGlobal()));
602 case MachineOperand::MO_FPImmediate: {
603 const ConstantFP *Cnt = MO.getFPImm();
604 const APFloat &Val = Cnt->getValueAPF();
605
606 switch (Cnt->getType()->getTypeID()) {
607 default:
608 report_fatal_error(reason: "Unsupported FP type");
609 break;
610 case Type::HalfTyID:
611 return MCOperand::createExpr(
612 Val: NVPTXFloatMCExpr::createConstantFPHalf(Flt: Val, Ctx&: OutContext));
613 case Type::BFloatTyID:
614 return MCOperand::createExpr(
615 Val: NVPTXFloatMCExpr::createConstantBFPHalf(Flt: Val, Ctx&: OutContext));
616 case Type::FloatTyID:
617 return MCOperand::createExpr(
618 Val: NVPTXFloatMCExpr::createConstantFPSingle(Flt: Val, Ctx&: OutContext));
619 case Type::DoubleTyID:
620 return MCOperand::createExpr(
621 Val: NVPTXFloatMCExpr::createConstantFPDouble(Flt: Val, Ctx&: OutContext));
622 }
623 break;
624 }
625 }
626}
627
628static NVPTX::VirtualRegisterKind
629getVirtualRegisterKind(const TargetRegisterClass *RC) {
630 if (RC == &NVPTX::B1RegClass)
631 return NVPTX::VirtualRegisterKind::B1;
632 if (RC == &NVPTX::B16RegClass)
633 return NVPTX::VirtualRegisterKind::B16;
634 if (RC == &NVPTX::B32RegClass)
635 return NVPTX::VirtualRegisterKind::B32;
636 if (RC == &NVPTX::B64RegClass)
637 return NVPTX::VirtualRegisterKind::B64;
638 if (RC == &NVPTX::B128RegClass)
639 return NVPTX::VirtualRegisterKind::B128;
640 llvm_unreachable("Bad register class");
641}
642
643unsigned NVPTXAsmPrinter::getVirtualRegisterNumber(Register Reg) const {
644 const auto It = VRegMapping.find(Val: MRI->getRegClass(Reg));
645 assert(It != VRegMapping.end() && "Bad register class");
646
647 const unsigned Num = It->second.lookup(Val: Reg);
648 assert(Num && "Bad virtual register");
649 return Num;
650}
651
652MCRegister NVPTXAsmPrinter::encodeVirtualRegister(Register Reg) {
653 if (Reg.isVirtual()) {
654 // Pack the register class into the upper bits so that
655 // NVPTXInstPrinter::printRegName can recover the declared name.
656 const auto Kind = getVirtualRegisterKind(RC: MRI->getRegClass(Reg));
657 const unsigned Num = getVirtualRegisterNumber(Reg);
658 assert(Num <= NVPTX::VirtualRegisterNumMask &&
659 "Too many virtual registers");
660 return (static_cast<unsigned>(Kind) << NVPTX::VirtualRegisterKindShift) |
661 Num;
662 }
663
664 // Some special-use registers are actually physical registers.
665 // Encode this as the register class ID of 0 and the real register ID.
666 assert(Reg.id() <= NVPTX::VirtualRegisterNumMask &&
667 "Physical register would decode as a virtual register");
668 return Reg.asMCReg();
669}
670
671MCOperand NVPTXAsmPrinter::GetSymbolRef(const MCSymbol *Symbol) {
672 const MCExpr *Expr;
673 Expr = MCSymbolRefExpr::create(Symbol, Ctx&: OutContext);
674 return MCOperand::createExpr(Val: Expr);
675}
676
677template <typename OwnerT>
678static void printParam(const OwnerT *Owner, Type *Ty, unsigned AttrIdx,
679 bool IsByVal, bool IsKernel, StringRef Name,
680 const DataLayout &DL, raw_ostream &O) {
681 O << ".param ";
682
683 if (IsByVal || shouldPassAsArray(Ty)) {
684 const Align ParamAlign =
685 IsByVal && !IsKernel ? getDeviceByValParamAlign(Owner, Ty, AttrIdx, DL)
686 : getPTXParamAlign(Owner, Ty, AttrIdx, DL);
687 O << ".align " << ParamAlign.value() << " .b8 " << Name << "["
688 << DL.getTypeAllocSize(Ty) << "]";
689 return;
690 }
691
692 assert((Ty->isFloatingPointTy() || Ty->isIntOrPtrTy()) &&
693 "Unknown parameter type");
694 const unsigned Size = DL.getTypeSizeInBits(Ty).getFixedValue();
695 O << ".b"
696 << (IsKernel ? promoteScalarKernelArgumentSize(Size)
697 : promoteScalarArgumentSize(Size))
698 << " " << Name;
699}
700
701template <typename OwnerT>
702static void printReturnValClause(const OwnerT *Owner, StringRef Name,
703 const DataLayout &DL, raw_ostream &O) {
704 Type *RetTy = Owner->getFunctionType()->getReturnType();
705
706 // A void or zero-sized return type (e.g. an empty struct) produces no return
707 // parameter.
708 if (RetTy->isVoidTy() || RetTy->isEmptyTy())
709 return;
710
711 // Only device functions return a value, so no kernel promotion applies.
712 O << "(";
713 printParam(Owner, RetTy, AttributeList::ReturnIndex, /*IsByVal=*/false,
714 /*IsKernel=*/false, Name, DL, O);
715 O << ") ";
716}
717
718void NVPTXAsmPrinter::emitCallPrototype(const CallBase &CB,
719 MCSymbol *PrototypeSymbol) const {
720 const DataLayout &DL = getDataLayout();
721 const NVPTXSubtarget &STI = MF->getSubtarget<NVPTXSubtarget>();
722
723 OutStreamer->emitLabel(Symbol: PrototypeSymbol);
724
725 SmallString<128> Str;
726 raw_svector_ostream O(Str);
727
728 O << ".callprototype ";
729 printReturnValClause(Owner: &CB, Name: "_", DL, O);
730 O << "_ (";
731
732 auto MakeArg = [&](const unsigned I) {
733 const bool IsByVal = CB.isByValArgument(ArgNo: I);
734 Type *Ty =
735 IsByVal ? CB.getParamByValType(ArgNo: I) : CB.getArgOperand(i: I)->getType();
736
737 printParam(Owner: &CB, Ty, AttrIdx: I + AttributeList::FirstArgIndex, IsByVal,
738 /*IsKernel=*/false, Name: "_", DL, O);
739 };
740
741 const FunctionType *FTy = CB.getFunctionType();
742 const unsigned NumArgs = FTy->getNumParams();
743
744 // Zero-sized arguments (e.g. empty structs) are not passed and so do not
745 // appear in the prototype.
746 const auto NonEmptyArgs = make_filter_range(Range: seq(Size: NumArgs), Pred: [&](unsigned I) {
747 return !CB.getArgOperand(i: I)->getType()->isEmptyTy();
748 });
749
750 interleave(c: NonEmptyArgs, os&: O, each_fn: MakeArg, separator: ", ");
751
752 if (FTy->isVarArg() && CB.arg_size() > NumArgs)
753 O << (NonEmptyArgs.empty() ? "" : ",") << " .param .align "
754 << STI.getMaxRequiredAlignment() << " .b8 _[]";
755
756 O << ")";
757 if (shouldEmitPTXNoReturn(V: CB))
758 O << " .noreturn";
759 O << ";\n";
760
761 OutStreamer->emitRawText(String: O.str());
762}
763
764void NVPTXAsmPrinter::emitJumpTable(const MachineJumpTableEntry &MJT,
765 unsigned MJTI) const {
766 OutStreamer->emitLabel(Symbol: GetJTISymbol(JTID: MJTI));
767
768 if (MJT.MBBs.empty())
769 return;
770
771 const auto Targets = to_vector(
772 Range: map_range(C: MJT.MBBs, F: [](const MachineBasicBlock *MBB) -> const MCSymbol * {
773 return MBB->getSymbol();
774 }));
775 getTargetStreamer()->emitBranchTargetsDirective(Targets);
776}
777
778// Return true if MBB is the header of a loop marked with
779// llvm.loop.unroll.disable or llvm.loop.unroll.count=1.
780bool NVPTXAsmPrinter::isLoopHeaderOfNoUnroll(
781 const MachineBasicBlock &MBB) const {
782 const MachineLoopInfo *LI = GetMLI(*MF);
783 assert(LI && "NVPTXAsmPrinter requires MachineLoopInfo");
784 // We insert .pragma "nounroll" only to the loop header.
785 if (!LI->isLoopHeader(BB: &MBB))
786 return false;
787
788 // llvm.loop.unroll.disable is marked on the back edges of a loop. Therefore,
789 // we iterate through each back edge of the loop with header MBB, and check
790 // whether its metadata contains llvm.loop.unroll.disable.
791 for (const MachineBasicBlock *PMBB : MBB.predecessors()) {
792 if (LI->getLoopFor(BB: PMBB) != LI->getLoopFor(BB: &MBB)) {
793 // Edges from other loops to MBB are not back edges.
794 continue;
795 }
796 if (const BasicBlock *PBB = PMBB->getBasicBlock()) {
797 if (MDNode *LoopID =
798 PBB->getTerminator()->getMetadata(KindID: LLVMContext::MD_loop)) {
799 if (GetUnrollMetadata(LoopID, Name: "llvm.loop.unroll.disable"))
800 return true;
801 if (MDNode *UnrollCountMD =
802 GetUnrollMetadata(LoopID, Name: "llvm.loop.unroll.count")) {
803 if (mdconst::extract<ConstantInt>(MD: UnrollCountMD->getOperand(I: 1))
804 ->isOne())
805 return true;
806 }
807 }
808 }
809 }
810 return false;
811}
812
813void NVPTXAsmPrinter::emitBasicBlockStart(const MachineBasicBlock &MBB) {
814 AsmPrinter::emitBasicBlockStart(MBB);
815 if (isLoopHeaderOfNoUnroll(MBB))
816 getTargetStreamer()->emitPragmaDirective(Pragma: "nounroll");
817}
818
819void NVPTXAsmPrinter::emitFunctionEntryLabel() {
820 SmallString<128> Str;
821 raw_svector_ostream O(Str);
822
823 if (!GlobalsEmitted) {
824 emitGlobals(M: *MF->getFunction().getParent());
825 GlobalsEmitted = true;
826 }
827
828 // Set up
829 MRI = &MF->getRegInfo();
830 F = &MF->getFunction();
831 emitLinkageDirective(V: F, O);
832 if (isKernelFunction(F: *F))
833 O << ".entry ";
834 else {
835 O << ".func ";
836 printReturnValClause(Owner: F, Name: "func_retval0", DL: getDataLayout(), O);
837 }
838
839 CurrentFnSym->print(OS&: O, MAI);
840
841 emitFunctionParamList(F, O);
842 O << "\n";
843
844 if (isKernelFunction(F: *F))
845 emitKernelFunctionDirectives(F: *F, O);
846
847 if (shouldEmitPTXNoReturn(V: *F))
848 O << ".noreturn";
849
850 OutStreamer->emitRawText(String: O.str());
851
852 VRegMapping.clear();
853 // Emit open brace for function body.
854 OutStreamer->emitRawText(String: StringRef("{\n"));
855 setAndEmitFunctionVirtualRegisters(*MF);
856 encodeDebugInfoRegisterNumbers(MF: *MF);
857 // Emit initial .loc debug directive for correct relocation symbol data.
858 emitInitialRawDwarfLocDirective(MF: *MF, DD: getDwarfDebug(), OutStreamer&: *OutStreamer);
859}
860
861bool NVPTXAsmPrinter::runOnMachineFunction(MachineFunction &F) {
862 bool Result = AsmPrinter::runOnMachineFunction(MF&: F);
863 // Emit closing brace for the body of function F.
864 // The closing brace must be emitted here because we need to emit additional
865 // debug labels/data after the last basic block.
866 // We need to emit the closing brace here because we don't have function that
867 // finished emission of the function body.
868 OutStreamer->emitRawText(String: StringRef("}\n"));
869 return Result;
870}
871
872void NVPTXAsmPrinter::emitFunctionBodyStart() {
873 SmallString<128> Str;
874 raw_svector_ostream O(Str);
875 emitDemotedVars(&MF->getFunction(), O);
876 OutStreamer->emitRawText(String: O.str());
877
878 const auto *MFI = MF->getInfo<NVPTXMachineFunctionInfo>();
879 for (const auto &[CB, Symbol] : MFI->getCallPrototypes())
880 emitCallPrototype(CB: *CB, PrototypeSymbol: Symbol);
881
882 if (const MachineJumpTableInfo *MJTI = MF->getJumpTableInfo())
883 for (const auto &[Idx, JT] : enumerate(First: MJTI->getJumpTables()))
884 emitJumpTable(MJT: JT, MJTI: Idx);
885}
886
887void NVPTXAsmPrinter::emitFunctionBodyEnd() {
888 VRegMapping.clear();
889}
890
891const MCSymbol *NVPTXAsmPrinter::getFunctionFrameSymbol() const {
892 return OutContext.getOrCreateSymbol(DEPOTNAME + Twine(getFunctionNumber()));
893}
894
895void NVPTXAsmPrinter::emitImplicitDef(const MachineInstr *MI) const {
896 Register RegNo = MI->getOperand(i: 0).getReg();
897 if (RegNo.isVirtual())
898 OutStreamer->AddComment(T: Twine("implicit-def: ") +
899 getVirtualRegisterName(Reg: RegNo));
900 else
901 OutStreamer->AddComment(T: Twine("implicit-def: ") +
902 NVPTXInstPrinter::getRegisterName(Reg: RegNo));
903 OutStreamer->addBlankLine();
904}
905
906void NVPTXAsmPrinter::emitKernelFunctionDirectives(const Function &F,
907 raw_ostream &O) const {
908 // If the NVVM IR has some of reqntid* specified, then output
909 // the reqntid directive, and set the unspecified ones to 1.
910 // If none of Reqntid* is specified, don't output reqntid directive.
911 const auto ReqNTID = getReqNTID(F);
912 if (!ReqNTID.empty())
913 O << formatv(Fmt: ".reqntid {0:$[, ]}\n",
914 Vals: make_range(x: ReqNTID.begin(), y: ReqNTID.end()));
915
916 const auto MaxNTID = getMaxNTID(F);
917 if (!MaxNTID.empty())
918 O << formatv(Fmt: ".maxntid {0:$[, ]}\n",
919 Vals: make_range(x: MaxNTID.begin(), y: MaxNTID.end()));
920
921 if (const auto Mincta = getMinCTASm(F))
922 O << ".minnctapersm " << *Mincta << "\n";
923
924 if (const auto Maxnreg = getMaxNReg(F))
925 O << ".maxnreg " << *Maxnreg << "\n";
926
927 // .maxclusterrank directive requires SM_90 or higher, make sure that we
928 // filter it out for lower SM versions, as it causes a hard ptxas crash.
929 const NVPTXTargetMachine &NTM = static_cast<const NVPTXTargetMachine &>(TM);
930 const NVPTXSubtarget *STI = &NTM.getSubtarget<NVPTXSubtarget>(F);
931
932 if (STI->hasFeature(Feature: NVPTX::SM90)) {
933 const auto ClusterDim = getClusterDim(F);
934 const bool BlocksAreClusters = hasBlocksAreClusters(F);
935
936 if (!ClusterDim.empty()) {
937
938 if (!BlocksAreClusters)
939 O << ".explicitcluster\n";
940
941 if (ClusterDim[0] != 0) {
942 assert(llvm::all_of(ClusterDim, not_equal_to(0)) &&
943 "cluster_dim_x != 0 implies cluster_dim_y and cluster_dim_z "
944 "should be non-zero as well");
945
946 O << formatv(Fmt: ".reqnctapercluster {0:$[, ]}\n",
947 Vals: make_range(x: ClusterDim.begin(), y: ClusterDim.end()));
948 } else {
949 assert(llvm::all_of(ClusterDim, equal_to(0)) &&
950 "cluster_dim_x == 0 implies cluster_dim_y and cluster_dim_z "
951 "should be 0 as well");
952 }
953 }
954
955 if (BlocksAreClusters) {
956 LLVMContext &Ctx = F.getContext();
957 if (ReqNTID.empty() || ClusterDim.empty())
958 Ctx.diagnose(DI: DiagnosticInfoUnsupported(
959 F, "blocksareclusters requires reqntid and cluster_dim attributes",
960 F.getSubprogram()));
961 else if (!STI->hasFeature(Feature: NVPTX::PTX90))
962 Ctx.diagnose(DI: DiagnosticInfoUnsupported(
963 F, "blocksareclusters requires PTX version >= 9.0",
964 F.getSubprogram()));
965 else
966 O << ".blocksareclusters\n";
967 }
968
969 if (const auto Maxclusterrank = getMaxClusterRank(F))
970 O << ".maxclusterrank " << *Maxclusterrank << "\n";
971 }
972}
973
974std::string NVPTXAsmPrinter::getVirtualRegisterName(Register Reg) const {
975 const auto Kind = getVirtualRegisterKind(RC: MRI->getRegClass(Reg));
976
977 std::string Name;
978 raw_string_ostream(Name) << NVPTX::getVirtualRegisterPrefix(Kind)
979 << getVirtualRegisterNumber(Reg);
980 return Name;
981}
982
983void NVPTXAsmPrinter::emitAliasDeclaration(const GlobalAlias *GA,
984 raw_ostream &O) {
985 const Function *F = dyn_cast_or_null<Function>(Val: GA->getAliaseeObject());
986 if (!F || isKernelFunction(F: *F) || F->isDeclaration())
987 report_fatal_error(
988 reason: "NVPTX aliasee must be a non-kernel function definition");
989
990 if (GA->hasLinkOnceLinkage() || GA->hasWeakLinkage() ||
991 GA->hasAvailableExternallyLinkage() || GA->hasCommonLinkage())
992 report_fatal_error(reason: "NVPTX aliasee must not be '.weak'");
993
994 emitDeclarationWithName(F, getSymbol(GV: GA), O);
995}
996
997void NVPTXAsmPrinter::emitDeclaration(const Function *F, raw_ostream &O) {
998 emitDeclarationWithName(F, getSymbol(GV: F), O);
999}
1000
1001void NVPTXAsmPrinter::emitDeclarationWithName(const Function *F, MCSymbol *S,
1002 raw_ostream &O) {
1003 emitLinkageDirective(V: F, O);
1004 if (isKernelFunction(F: *F)) {
1005 O << ".entry ";
1006 } else {
1007 O << ".func ";
1008 printReturnValClause(Owner: F, Name: "func_retval0", DL: getDataLayout(), O);
1009 }
1010 S->print(OS&: O, MAI);
1011 O << "\n";
1012 emitFunctionParamList(F, O);
1013 O << "\n";
1014 if (shouldEmitPTXNoReturn(V: *F))
1015 O << ".noreturn";
1016 O << ";\n";
1017}
1018
1019static bool usedInGlobalVarDef(const Constant *C) {
1020 if (!C)
1021 return false;
1022
1023 if (const GlobalVariable *GV = dyn_cast<GlobalVariable>(Val: C))
1024 return GV->getName() != "llvm.used";
1025
1026 for (const User *U : C->users())
1027 if (const Constant *C = dyn_cast<Constant>(Val: U))
1028 if (usedInGlobalVarDef(C))
1029 return true;
1030
1031 return false;
1032}
1033
1034static bool usedInOneFunc(const User *U, Function const *&OneFunc) {
1035 if (const GlobalVariable *OtherGV = dyn_cast<GlobalVariable>(Val: U))
1036 if (OtherGV->getName() == "llvm.used")
1037 return true;
1038
1039 if (const Instruction *I = dyn_cast<Instruction>(Val: U)) {
1040 if (const Function *CurFunc = I->getFunction()) {
1041 if (OneFunc && (CurFunc != OneFunc))
1042 return false;
1043 OneFunc = CurFunc;
1044 return true;
1045 }
1046 return false;
1047 }
1048
1049 for (const User *UU : U->users())
1050 if (!usedInOneFunc(U: UU, OneFunc))
1051 return false;
1052
1053 return true;
1054}
1055
1056/* Find out if a global variable can be demoted to local scope.
1057 * Currently, this is valid for CUDA shared variables, which have local
1058 * scope and global lifetime. So the conditions to check are :
1059 * 1. Is the global variable in shared address space?
1060 * 2. Does it have local linkage?
1061 * 3. Is the global variable referenced only in one function?
1062 */
1063static bool canDemoteGlobalVar(const GlobalVariable *GV, Function const *&f) {
1064 if (!GV->hasLocalLinkage())
1065 return false;
1066 if (GV->getAddressSpace() != ADDRESS_SPACE_SHARED)
1067 return false;
1068
1069 const Function *oneFunc = nullptr;
1070
1071 bool flag = usedInOneFunc(U: GV, OneFunc&: oneFunc);
1072 if (!flag)
1073 return false;
1074 if (!oneFunc)
1075 return false;
1076 f = oneFunc;
1077 return true;
1078}
1079
1080static bool useFuncSeen(const Constant *C,
1081 const SmallPtrSetImpl<const Function *> &SeenSet) {
1082 for (const User *U : C->users()) {
1083 if (const Constant *cu = dyn_cast<Constant>(Val: U)) {
1084 if (useFuncSeen(C: cu, SeenSet))
1085 return true;
1086 } else if (const Instruction *I = dyn_cast<Instruction>(Val: U)) {
1087 if (const Function *Caller = I->getFunction())
1088 if (SeenSet.contains(Ptr: Caller))
1089 return true;
1090 }
1091 }
1092 return false;
1093}
1094
1095void NVPTXAsmPrinter::emitDeclarations(const Module &M, raw_ostream &O) {
1096 SmallPtrSet<const Function *, 32> SeenSet;
1097 for (const Function &F : M) {
1098 if (F.getAttributes().hasFnAttr(Kind: "nvptx-libcall-callee")) {
1099 emitDeclaration(F: &F, O);
1100 continue;
1101 }
1102
1103 if (F.isDeclaration()) {
1104 if (F.use_empty())
1105 continue;
1106 if (F.getIntrinsicID())
1107 continue;
1108 // An unrecognized intrinsic would produce an invalid PTX declaration. Let
1109 // the user know that, and skip it.
1110 if (F.isIntrinsic()) {
1111 LLVMContext &Ctx = F.getContext();
1112 Ctx.diagnose(DI: DiagnosticInfoUnsupported(
1113 F, "unknown intrinsic '" + F.getName() +
1114 "' cannot be lowered by the NVPTX backend"));
1115 continue;
1116 }
1117 emitDeclaration(F: &F, O);
1118 continue;
1119 }
1120 for (const User *U : F.users()) {
1121 if (const Constant *C = dyn_cast<Constant>(Val: U)) {
1122 if (usedInGlobalVarDef(C)) {
1123 // The use is in the initialization of a global variable
1124 // that is a function pointer, so print a declaration
1125 // for the original function
1126 emitDeclaration(F: &F, O);
1127 break;
1128 }
1129 // Emit a declaration of this function if the function that
1130 // uses this constant expr has already been seen.
1131 if (useFuncSeen(C, SeenSet)) {
1132 emitDeclaration(F: &F, O);
1133 break;
1134 }
1135 }
1136
1137 if (!isa<Instruction>(Val: U))
1138 continue;
1139 const Function *Caller = cast<Instruction>(Val: U)->getFunction();
1140 if (!Caller)
1141 continue;
1142
1143 // If a caller has already been seen, then the caller is
1144 // appearing in the module before the callee. so print out
1145 // a declaration for the callee.
1146 if (SeenSet.contains(Ptr: Caller)) {
1147 emitDeclaration(F: &F, O);
1148 break;
1149 }
1150 }
1151 SeenSet.insert(Ptr: &F);
1152 }
1153 for (const GlobalAlias &GA : M.aliases())
1154 emitAliasDeclaration(GA: &GA, O);
1155}
1156
1157void NVPTXAsmPrinter::emitStartOfAsmFile(Module &M) {
1158 // Construct a default subtarget off of the TargetMachine defaults. The
1159 // rest of NVPTX isn't friendly to change subtargets per function and
1160 // so the default TargetMachine will have all of the options.
1161 const NVPTXTargetMachine &NTM = static_cast<const NVPTXTargetMachine &>(TM);
1162 const NVPTXSubtarget *STI = NTM.getSubtargetImpl();
1163
1164 // Emit header before any dwarf directives are emitted below.
1165 emitHeader(M, STI: *STI);
1166}
1167
1168/// Create NVPTX-specific DwarfDebug handler.
1169DwarfDebug *NVPTXAsmPrinter::createDwarfDebug() {
1170 return new NVPTXDwarfDebug(this);
1171}
1172
1173bool NVPTXAsmPrinter::doInitialization(Module &M) {
1174 const NVPTXTargetMachine &NTM = static_cast<const NVPTXTargetMachine &>(TM);
1175 const NVPTXSubtarget &STI = *NTM.getSubtargetImpl();
1176 if (M.alias_size() &&
1177 (!STI.hasFeature(Feature: NVPTX::PTX63) || !STI.hasFeature(Feature: NVPTX::SM30)))
1178 report_fatal_error(reason: ".alias requires PTX version >= 6.3 and sm_30");
1179
1180 // We need to call the parent's one explicitly.
1181 bool Result = AsmPrinter::doInitialization(M);
1182
1183 GlobalsEmitted = false;
1184
1185 // Ensure globals are in the symbol table before ISel so any temp symbols are
1186 // guaranteed not to collide with user symbols
1187 for (const GlobalValue &GV : M.global_values())
1188 getSymbol(GV: &GV);
1189
1190 return Result;
1191}
1192
1193void NVPTXAsmPrinter::emitGlobals(const Module &M) {
1194 SmallString<128> Str2;
1195 raw_svector_ostream OS2(Str2);
1196
1197 emitDeclarations(M, O&: OS2);
1198
1199 const NVPTXTargetMachine &NTM = static_cast<const NVPTXTargetMachine &>(TM);
1200 const NVPTXSubtarget &STI = *NTM.getSubtargetImpl();
1201
1202 // ptxas requires global symbols referenced by initializers to be known
1203 // before use. Acyclic dependencies can be handled by dependency-first
1204 // emission. Cyclic SCCs need compatible .extern declarations first.
1205 // Edges point from each global to the globals used by its initializer.
1206 // Reverse-topological SCC iteration therefore emits dependencies first.
1207 GlobalVariableDependencyGraph DependencyGraph(M);
1208 for (GlobalVariableSCCIterator I =
1209 GlobalVariableSCCIterator::begin(G: DependencyGraph.getEntryNode());
1210 !I.isAtEnd(); ++I) {
1211 SmallVector<const GlobalVariableDependencyNode *, 4> SCC(I->begin(),
1212 I->end());
1213
1214 // Nothing points to the synthetic root, so it is always in its own SCC.
1215 if (!SCC.front()->GV) {
1216 assert(SCC.size() == 1 && "Synthetic root must be in its own SCC");
1217 continue;
1218 }
1219
1220 llvm::sort(C&: SCC, Comp: [](const auto *LHS, const auto *RHS) {
1221 return LHS->ModuleOrder < RHS->ModuleOrder;
1222 });
1223
1224 const bool IsCyclic = I.hasCycle();
1225 DenseSet<const GlobalVariableDependencyNode *> ForwardDeclared;
1226 if (IsCyclic)
1227 for (const auto *Node : SCC)
1228 if (isForwardDeclarableGlobal(GVar: Node->GV))
1229 ForwardDeclared.insert(V: Node);
1230
1231 // Check that declarations break every cycle before writing any output.
1232 SmallVector<const GlobalVariable *, 4> OrderedGlobals =
1233 IsCyclic ? orderDefinitionsInSCC(SCC, ForwardDeclared)
1234 : SmallVector<const GlobalVariable *, 4>{SCC.front()->GV};
1235
1236 for (const auto *Node : SCC) {
1237 if (!ForwardDeclared.count(V: Node))
1238 continue;
1239 OS2 << ".extern ";
1240 emitPTXGlobalVariableDefinition(GVar: Node->GV, O&: OS2, STI,
1241 /*EmitInitializer=*/false);
1242 OS2 << ";\n";
1243 }
1244
1245 for (const GlobalVariable *GV : OrderedGlobals)
1246 printModuleLevelGV(GVar: GV, O&: OS2, /*ProcessDemoted=*/processDemoted: false, STI);
1247 }
1248
1249 OS2 << '\n';
1250
1251 OutStreamer->emitRawText(String: OS2.str());
1252}
1253
1254void NVPTXAsmPrinter::emitGlobalAlias(const Module &M, const GlobalAlias &GA) {
1255 getTargetStreamer()->emitAliasDirective(Name: getSymbol(GV: &GA),
1256 Aliasee: getSymbol(GV: GA.getAliaseeObject()));
1257}
1258
1259NVPTXTargetStreamer *NVPTXAsmPrinter::getTargetStreamer() const {
1260 return static_cast<NVPTXTargetStreamer *>(OutStreamer->getTargetStreamer());
1261}
1262
1263static bool hasFullDebugInfo(Module &M) {
1264 for (DICompileUnit *CU : M.debug_compile_units()) {
1265 switch(CU->getEmissionKind()) {
1266 case DICompileUnit::NoDebug:
1267 case DICompileUnit::DebugDirectivesOnly:
1268 break;
1269 case DICompileUnit::LineTablesOnly:
1270 case DICompileUnit::FullDebug:
1271 return true;
1272 }
1273 }
1274
1275 return false;
1276}
1277
1278void NVPTXAsmPrinter::emitHeader(Module &M, const NVPTXSubtarget &STI) {
1279 auto *TS = getTargetStreamer();
1280
1281 TS->emitBanner();
1282
1283 const unsigned PTXVersion = STI.getPTXVersion();
1284 TS->emitVersionDirective(PTXVersion);
1285
1286 const NVPTXTargetMachine &NTM = static_cast<const NVPTXTargetMachine &>(TM);
1287 bool TexModeIndependent = NTM.getDrvInterface() == NVPTX::NVCL;
1288
1289 TS->emitTargetDirective(Target: STI.getTargetName(), TexModeIndependent,
1290 HasDebug: hasFullDebugInfo(M));
1291 TS->emitAddressSizeDirective(AddrSize: M.getDataLayout().getPointerSizeInBits());
1292}
1293
1294bool NVPTXAsmPrinter::doFinalization(Module &M) {
1295 // If we did not emit any functions, then the global declarations have not
1296 // yet been emitted.
1297 if (!GlobalsEmitted) {
1298 emitGlobals(M);
1299 GlobalsEmitted = true;
1300 }
1301
1302 // call doFinalization
1303 bool ret = AsmPrinter::doFinalization(M);
1304
1305 clearAnnotationCache(&M);
1306
1307 auto *TS =
1308 static_cast<NVPTXTargetStreamer *>(OutStreamer->getTargetStreamer());
1309 // Close the last emitted section
1310 if (hasDebugInfo()) {
1311 TS->closeLastSection();
1312 // Emit empty .debug_macinfo section for better support of the empty files.
1313 TS->emitEmptySectionDirective(Name: ".debug_macinfo");
1314 }
1315
1316 // Output last DWARF .file directives, if any.
1317 TS->outputDwarfFileDirectives();
1318
1319 return ret;
1320}
1321
1322// This function emits appropriate linkage directives for
1323// functions and global variables.
1324//
1325// extern function declaration -> .extern
1326// extern function definition -> .visible
1327// external global variable with init -> .visible
1328// external without init -> .extern
1329// appending -> not allowed, assert.
1330// for any linkage other than
1331// internal, private, linker_private,
1332// linker_private_weak, linker_private_weak_def_auto,
1333// we emit -> .weak.
1334
1335void NVPTXAsmPrinter::emitLinkageDirective(const GlobalValue *V,
1336 raw_ostream &O) {
1337 if (static_cast<NVPTXTargetMachine &>(TM).getDrvInterface() == NVPTX::CUDA) {
1338 if (V->hasExternalLinkage()) {
1339 if (const auto *GVar = dyn_cast<GlobalVariable>(Val: V))
1340 O << (GVar->hasInitializer() ? ".visible " : ".extern ");
1341 else if (V->isDeclaration())
1342 O << ".extern ";
1343 else
1344 O << ".visible ";
1345 } else if (V->hasAppendingLinkage()) {
1346 report_fatal_error(reason: "Symbol '" + (V->hasName() ? V->getName() : "") +
1347 "' has unsupported appending linkage type");
1348 } else if (!V->hasInternalLinkage() && !V->hasPrivateLinkage()) {
1349 O << ".weak ";
1350 }
1351 }
1352}
1353
1354void NVPTXAsmPrinter::printModuleLevelGV(const GlobalVariable *GVar,
1355 raw_ostream &O, bool ProcessDemoted,
1356 const NVPTXSubtarget &STI) {
1357 // Skip metadata and LLVM intrinsic global variables.
1358 if (shouldSkipModuleLevelGlobal(GV: *GVar))
1359 return;
1360
1361 if (GVar->hasExternalLinkage()) {
1362 if (GVar->hasInitializer())
1363 O << ".visible ";
1364 else
1365 O << ".extern ";
1366 } else if (STI.hasFeature(Feature: NVPTX::PTX50) && GVar->hasCommonLinkage() &&
1367 GVar->getAddressSpace() == ADDRESS_SPACE_GLOBAL) {
1368 O << ".common ";
1369 } else if (GVar->hasLinkOnceLinkage() || GVar->hasWeakLinkage() ||
1370 GVar->hasAvailableExternallyLinkage() ||
1371 GVar->hasCommonLinkage()) {
1372 O << ".weak ";
1373 }
1374
1375 const PTXOpaqueType OpaqueType = getPTXOpaqueType(*GVar);
1376
1377 if (OpaqueType == PTXOpaqueType::Texture) {
1378 O << ".global .texref ";
1379 getSymbol(GV: GVar)->print(OS&: O, MAI);
1380 O << ";\n";
1381 return;
1382 }
1383
1384 if (OpaqueType == PTXOpaqueType::Surface) {
1385 O << ".global .surfref ";
1386 getSymbol(GV: GVar)->print(OS&: O, MAI);
1387 O << ";\n";
1388 return;
1389 }
1390
1391 if (GVar->isDeclaration()) {
1392 // (extern) declarations, no definition or initializer
1393 // Currently the only known declaration is for an automatic __local
1394 // (.shared) promoted to global.
1395 emitPTXGlobalVariableDefinition(GVar, O, STI, /*EmitInitializer=*/false);
1396 O << ";\n";
1397 return;
1398 }
1399
1400 if (OpaqueType == PTXOpaqueType::Sampler) {
1401 O << ".global .samplerref ";
1402 getSymbol(GV: GVar)->print(OS&: O, MAI);
1403
1404 const Constant *Initializer = nullptr;
1405 if (GVar->hasInitializer())
1406 Initializer = GVar->getInitializer();
1407 const ConstantInt *CI = nullptr;
1408 if (Initializer)
1409 CI = dyn_cast<ConstantInt>(Val: Initializer);
1410 if (CI) {
1411 unsigned sample = CI->getZExtValue();
1412
1413 O << " = { ";
1414
1415 for (int i = 0,
1416 addr = ((sample & __CLK_ADDRESS_MASK) >> __CLK_ADDRESS_BASE);
1417 i < 3; i++) {
1418 O << "addr_mode_" << i << " = ";
1419 switch (addr) {
1420 case 0:
1421 O << "wrap";
1422 break;
1423 case 1:
1424 O << "clamp_to_border";
1425 break;
1426 case 2:
1427 O << "clamp_to_edge";
1428 break;
1429 case 3:
1430 O << "wrap";
1431 break;
1432 case 4:
1433 O << "mirror";
1434 break;
1435 }
1436 O << ", ";
1437 }
1438 O << "filter_mode = ";
1439 switch ((sample & __CLK_FILTER_MASK) >> __CLK_FILTER_BASE) {
1440 case 0:
1441 O << "nearest";
1442 break;
1443 case 1:
1444 O << "linear";
1445 break;
1446 case 2:
1447 llvm_unreachable("Anisotropic filtering is not supported");
1448 default:
1449 O << "nearest";
1450 break;
1451 }
1452 if (!((sample & __CLK_NORMALIZED_MASK) >> __CLK_NORMALIZED_BASE)) {
1453 O << ", force_unnormalized_coords = 1";
1454 }
1455 O << " }";
1456 }
1457
1458 O << ";\n";
1459 return;
1460 }
1461
1462 if (GVar->hasPrivateLinkage()) {
1463 if (GVar->getName().starts_with(Prefix: "unrollpragma"))
1464 return;
1465
1466 // FIXME - need better way (e.g. Metadata) to avoid generating this global
1467 if (GVar->getName().starts_with(Prefix: "filename"))
1468 return;
1469 if (GVar->use_empty())
1470 return;
1471 }
1472
1473 const Function *DemotedFunc = nullptr;
1474 if (!ProcessDemoted && canDemoteGlobalVar(GV: GVar, f&: DemotedFunc)) {
1475 O << "// " << GVar->getName() << " has been demoted\n";
1476 localDecls[DemotedFunc].push_back(x: GVar);
1477 return;
1478 }
1479
1480 emitPTXGlobalVariableDefinition(GVar, O, STI, /*EmitInitializer=*/true);
1481 O << ";\n";
1482}
1483
1484void NVPTXAsmPrinter::emitPTXGlobalVariableDefinition(
1485 const GlobalVariable *GVar, raw_ostream &O, const NVPTXSubtarget &STI,
1486 bool EmitInitializer) {
1487 const DataLayout &DL = getDataLayout();
1488
1489 Type *ETy = GVar->getValueType();
1490
1491 emitPTXAddressSpace(AddressSpace: GVar->getAddressSpace(), O);
1492
1493 if (isManaged(*GVar)) {
1494 if (!STI.hasFeature(Feature: NVPTX::PTX40) || !STI.hasFeature(Feature: NVPTX::SM30))
1495 report_fatal_error(
1496 reason: ".attribute(.managed) requires PTX version >= 4.0 and sm_30");
1497 O << " .attribute(.managed)";
1498 }
1499
1500 O << " .align "
1501 << GVar->getAlign().value_or(u: DL.getPrefTypeAlign(Ty: ETy)).value();
1502
1503 const Constant *Initializer = nullptr;
1504 if (GVar->hasInitializer()) {
1505 const Constant *Init = GVar->getInitializer();
1506 if (!Init->isNullValue() && !isa<UndefValue>(Val: Init)) {
1507 if (GVar->getAddressSpace() != ADDRESS_SPACE_GLOBAL &&
1508 GVar->getAddressSpace() != ADDRESS_SPACE_CONST)
1509 report_fatal_error(reason: "initial value of '" + GVar->getName() +
1510 "' is not allowed in addrspace(" +
1511 Twine(GVar->getAddressSpace()) + ")");
1512 Initializer = Init;
1513 }
1514 }
1515
1516 if (ETy->isPointerTy() || ((ETy->isIntegerTy() || ETy->isFloatingPointTy()) &&
1517 ETy->getScalarSizeInBits() <= 64)) {
1518 O << " ." << getPTXFundamentalTypeStr(Ty: ETy) << " ";
1519 getSymbol(GV: GVar)->print(OS&: O, MAI);
1520
1521 if (EmitInitializer && Initializer) {
1522 O << " = ";
1523 printScalarConstant(CPV: Initializer, O);
1524 }
1525 return;
1526 }
1527
1528 // Although PTX has direct support for struct type and array type and LLVM IR
1529 // is very similar to PTX, the LLVM CodeGen does not support for targets that
1530 // support these high level field accesses. Structs, arrays and vectors are
1531 // lowered into arrays of bytes.
1532 assert((ETy->isIntegerTy() || ETy->isFP128Ty() || ETy->isAggregateType() ||
1533 isa<FixedVectorType>(ETy)) &&
1534 "type not supported yet");
1535
1536 const uint64_t ElementSize = DL.getTypeStoreSize(Ty: ETy);
1537
1538 if (!Initializer) {
1539 O << " .b8 ";
1540 getSymbol(GV: GVar)->print(OS&: O, MAI);
1541 if (ElementSize)
1542 O << "[" << ElementSize << "]";
1543 else if (!EmitInitializer)
1544 O << "[]";
1545 return;
1546 }
1547
1548 AggBuffer aggBuffer(ElementSize, *this);
1549 bufferAggregateConstant(CV: Initializer, aggBuffer: &aggBuffer);
1550 if (aggBuffer.numSymbols()) {
1551 const unsigned int ptrSize = MAI.getCodePointerSize();
1552 if (ElementSize % ptrSize || !aggBuffer.allSymbolsAligned(ptrSize)) {
1553 // Print in bytes and use the mask() operator for pointers.
1554 if (!STI.hasMaskOperator())
1555 report_fatal_error(reason: "initialized packed aggregate with pointers '" +
1556 GVar->getName() +
1557 "' requires at least PTX ISA version 7.1");
1558 O << " .u8 ";
1559 getSymbol(GV: GVar)->print(OS&: O, MAI);
1560 O << "[" << ElementSize << "]";
1561 if (EmitInitializer) {
1562 O << " = {";
1563 aggBuffer.printBytes(os&: O);
1564 O << "}";
1565 }
1566 } else {
1567 O << " .u" << ptrSize * 8 << " ";
1568 getSymbol(GV: GVar)->print(OS&: O, MAI);
1569 O << "[" << ElementSize / ptrSize << "]";
1570 if (EmitInitializer) {
1571 O << " = {";
1572 aggBuffer.printWords(os&: O);
1573 O << "}";
1574 }
1575 }
1576 } else {
1577 O << " .b8 ";
1578 getSymbol(GV: GVar)->print(OS&: O, MAI);
1579 O << "[" << ElementSize << "]";
1580 if (EmitInitializer) {
1581 O << " = {";
1582 aggBuffer.printBytes(os&: O);
1583 O << "}";
1584 }
1585 }
1586}
1587
1588void NVPTXAsmPrinter::AggBuffer::printSymbol(unsigned nSym, raw_ostream &os) {
1589 const Value *v = Symbols[nSym];
1590 const Value *v0 = SymbolsBeforeStripping[nSym];
1591 if (const GlobalValue *GVar = dyn_cast<GlobalValue>(Val: v)) {
1592 MCSymbol *Name = AP.getSymbol(GV: GVar);
1593 PointerType *PTy = dyn_cast<PointerType>(Val: v0->getType());
1594 // Is v0 a generic pointer?
1595 bool isGenericPointer = PTy && PTy->getAddressSpace() == 0;
1596 if (EmitGeneric && isGenericPointer && !isa<Function>(Val: v)) {
1597 os << "generic(";
1598 Name->print(OS&: os, MAI: AP.MAI);
1599 os << ")";
1600 } else {
1601 Name->print(OS&: os, MAI: AP.MAI);
1602 }
1603 } else if (const ConstantExpr *CExpr = dyn_cast<ConstantExpr>(Val: v0)) {
1604 const MCExpr *Expr = AP.lowerConstantForGV(CV: CExpr, ProcessingGeneric: false);
1605 AP.printMCExpr(Expr: *Expr, OS&: os);
1606 } else
1607 llvm_unreachable("symbol type unknown");
1608}
1609
1610void NVPTXAsmPrinter::AggBuffer::printBytes(raw_ostream &os) {
1611 unsigned int ptrSize = AP.MAI.getCodePointerSize();
1612 // Do not emit trailing zero initializers. They will be zero-initialized by
1613 // ptxas. This saves on both space requirements for the generated PTX and on
1614 // memory use by ptxas. (See:
1615 // https://docs.nvidia.com/cuda/parallel-thread-execution/index.html#global-state-space)
1616 unsigned int InitializerCount = Size;
1617 // TODO: symbols make this harder, but it would still be good to trim trailing
1618 // 0s for aggs with symbols as well.
1619 if (numSymbols() == 0)
1620 while (InitializerCount >= 1 && !buffer[InitializerCount - 1])
1621 InitializerCount--;
1622
1623 symbolPosInBuffer.push_back(Elt: InitializerCount);
1624 unsigned int nSym = 0;
1625 unsigned int nextSymbolPos = symbolPosInBuffer[nSym];
1626 for (unsigned int pos = 0; pos < InitializerCount;) {
1627 if (pos)
1628 os << ", ";
1629 if (pos != nextSymbolPos) {
1630 os << (unsigned int)buffer[pos];
1631 ++pos;
1632 continue;
1633 }
1634 // Generate a per-byte mask() operator for the symbol, which looks like:
1635 // .global .u8 addr[] = {0xFF(foo), 0xFF00(foo), 0xFF0000(foo), ...};
1636 // See https://docs.nvidia.com/cuda/parallel-thread-execution/index.html#initializers
1637 std::string symText;
1638 llvm::raw_string_ostream oss(symText);
1639 printSymbol(nSym, os&: oss);
1640 for (unsigned i = 0; i < ptrSize; ++i) {
1641 if (i)
1642 os << ", ";
1643 llvm::write_hex(S&: os, N: 0xFFULL << i * 8, Style: HexPrintStyle::PrefixUpper);
1644 os << "(" << symText << ")";
1645 }
1646 pos += ptrSize;
1647 nextSymbolPos = symbolPosInBuffer[++nSym];
1648 assert(nextSymbolPos >= pos);
1649 }
1650}
1651
1652void NVPTXAsmPrinter::AggBuffer::printWords(raw_ostream &os) {
1653 unsigned int ptrSize = AP.MAI.getCodePointerSize();
1654 symbolPosInBuffer.push_back(Elt: Size);
1655 unsigned int nSym = 0;
1656 unsigned int nextSymbolPos = symbolPosInBuffer[nSym];
1657 assert(nextSymbolPos % ptrSize == 0);
1658 for (unsigned int pos = 0; pos < Size; pos += ptrSize) {
1659 if (pos)
1660 os << ", ";
1661 if (pos == nextSymbolPos) {
1662 printSymbol(nSym, os);
1663 nextSymbolPos = symbolPosInBuffer[++nSym];
1664 assert(nextSymbolPos % ptrSize == 0);
1665 assert(nextSymbolPos >= pos + ptrSize);
1666 } else if (ptrSize == 4)
1667 os << support::endian::read32le(P: &buffer[pos]);
1668 else
1669 os << support::endian::read64le(P: &buffer[pos]);
1670 }
1671}
1672
1673void NVPTXAsmPrinter::emitDemotedVars(const Function *F, raw_ostream &O) {
1674 auto It = localDecls.find(x: F);
1675 if (It == localDecls.end())
1676 return;
1677
1678 ArrayRef<const GlobalVariable *> GVars = It->second;
1679
1680 const NVPTXTargetMachine &NTM = static_cast<const NVPTXTargetMachine &>(TM);
1681 const NVPTXSubtarget &STI = *NTM.getSubtargetImpl();
1682
1683 for (const GlobalVariable *GV : GVars) {
1684 O << "\t// demoted variable\n\t";
1685 printModuleLevelGV(GVar: GV, O, /*processDemoted=*/ProcessDemoted: true, STI);
1686 }
1687}
1688
1689/// The PTX state space directive for \p AddressSpace, or an empty string if it
1690/// does not name one, as is the case for the generic address space.
1691static StringRef getPTXAddressSpaceName(unsigned AddressSpace) {
1692 switch (AddressSpace) {
1693 case ADDRESS_SPACE_LOCAL:
1694 return ".local";
1695 case ADDRESS_SPACE_GLOBAL:
1696 return ".global";
1697 case ADDRESS_SPACE_CONST:
1698 return ".const";
1699 case ADDRESS_SPACE_SHARED:
1700 return ".shared";
1701 default:
1702 return {};
1703 }
1704}
1705
1706/// The PTX opaque type directive for an image or sampler handle, or an empty
1707/// string for PTXOpaqueType::None.
1708static StringRef getPTXOpaqueTypeName(PTXOpaqueType OpaqueType) {
1709 switch (OpaqueType) {
1710 case PTXOpaqueType::Sampler:
1711 return ".samplerref";
1712 case PTXOpaqueType::Texture:
1713 return ".texref";
1714 case PTXOpaqueType::Surface:
1715 return ".surfref";
1716 case PTXOpaqueType::None:
1717 return {};
1718 }
1719 llvm_unreachable("unexpected PTXOpaqueType");
1720}
1721
1722void NVPTXAsmPrinter::emitPTXAddressSpace(unsigned int AddressSpace,
1723 raw_ostream &O) const {
1724 const StringRef Name = getPTXAddressSpaceName(AddressSpace);
1725 if (Name.empty())
1726 report_fatal_error(reason: "Bad address space found while emitting PTX: " +
1727 llvm::Twine(AddressSpace));
1728 O << Name;
1729}
1730
1731std::string NVPTXAsmPrinter::getPTXFundamentalTypeStr(Type *Ty) const {
1732 switch (Ty->getTypeID()) {
1733 case Type::IntegerTyID:
1734 case Type::PointerTyID: {
1735 const uint64_t NumBits = getDataLayout().getTypeStoreSizeInBits(Ty);
1736 assert(NumBits <= 64 && "type too large");
1737 return "u" + utostr(X: promoteScalarKernelArgumentSize(Size: NumBits));
1738 }
1739 case Type::BFloatTyID:
1740 case Type::HalfTyID:
1741 case Type::FloatTyID:
1742 case Type::DoubleTyID:
1743 return "b" + utostr(X: Ty->getScalarSizeInBits());
1744 default:
1745 break;
1746 }
1747 llvm_unreachable("unexpected type");
1748}
1749
1750void NVPTXAsmPrinter::emitFunctionParamList(const Function *F, raw_ostream &O) {
1751 const DataLayout &DL = getDataLayout();
1752 const NVPTXSubtarget &STI = TM.getSubtarget<NVPTXSubtarget>(F: *F);
1753 const auto *TLI = cast<NVPTXTargetLowering>(Val: STI.getTargetLowering());
1754 const NVPTXMachineFunctionInfo *MFI =
1755 MF ? MF->getInfo<NVPTXMachineFunctionInfo>() : nullptr;
1756
1757 const bool IsKernelFunc = isKernelFunction(F: *F);
1758
1759 // Zero-sized arguments (e.g. empty structs) do not produce a parameter.
1760 // Number the emitted parameters contiguously, skipping the zero-sized ones,
1761 // so that the names match those used in LowerFormalArguments and the
1762 // contiguous numbering used by callers (see LowerCall).
1763 const auto NonEmptyArgs =
1764 make_filter_range(Range: F->args(), Pred: [](const Argument &Arg) {
1765 return !Arg.getType()->isEmptyTy();
1766 });
1767
1768 if (NonEmptyArgs.empty() && !F->isVarArg()) {
1769 O << "()";
1770 return;
1771 }
1772
1773 O << "(\n";
1774
1775 auto MakeParam = [&](const auto &IndexedArg) {
1776 const auto &[ParamIndex, Arg] = IndexedArg;
1777 Type *Ty = Arg.getType();
1778 MCSymbol *const ParamSym = TLI->getParamSymbol(Ctx&: OutContext, F, Idx: ParamIndex);
1779
1780 O << "\t";
1781
1782 // A byval param is passed as a copy of the pointee and an aggregate is
1783 // passed as a blob of bytes; both are declared as a byte array.
1784 const bool IsByVal = Arg.hasByValAttr();
1785 const bool AsArray = IsByVal || shouldPassAsArray(Ty);
1786
1787 // Kernels declare image/sampler handles and the address space of a
1788 // pointee. Both of those are scalar handles, so a byte-array param is
1789 // neither.
1790 if (IsKernelFunc && !AsArray) {
1791 const StringRef OpaqueType = getPTXOpaqueTypeName(getPTXOpaqueType(Arg));
1792 if (!OpaqueType.empty()) {
1793 O << ".param ";
1794 if (!MFI || !MFI->checkImageHandleSymbol(Symbol: ParamSym))
1795 O << ".u64 .ptr ";
1796
1797 O << OpaqueType << " " << *ParamSym;
1798 return;
1799 }
1800
1801 if (auto *PTy = dyn_cast<PointerType>(Val: Ty)) {
1802 const unsigned AS = PTy->getAddressSpace();
1803 O << ".param .u" << DL.getPointerSizeInBits(AS) << " .ptr";
1804
1805 const StringRef Space = getPTXAddressSpaceName(AddressSpace: AS);
1806 if (!Space.empty())
1807 O << " " << Space;
1808
1809 O << " .align " << Arg.getParamAlign().valueOrOne().value() << " "
1810 << *ParamSym;
1811 return;
1812 }
1813 }
1814
1815 printParam(F, IsByVal ? Arg.getParamByValType() : Ty,
1816 Arg.getArgNo() + AttributeList::FirstArgIndex, IsByVal,
1817 IsKernelFunc, ParamSym->getName(), DL, O);
1818 };
1819
1820 interleave(c: enumerate(First: NonEmptyArgs), os&: O, each_fn: MakeParam, separator: ",\n");
1821
1822 if (F->isVarArg())
1823 O << (NonEmptyArgs.empty() ? "" : ",\n") << "\t.param .align "
1824 << STI.getMaxRequiredAlignment() << " .b8 "
1825 << *TLI->getParamSymbol(Ctx&: OutContext, F, /* vararg */ Idx: -1) << "[]";
1826
1827 O << "\n)";
1828}
1829
1830void NVPTXAsmPrinter::setAndEmitFunctionVirtualRegisters(
1831 const MachineFunction &MF) {
1832 auto *TS = getTargetStreamer();
1833
1834 // Emit the Fake Stack Object
1835 const MachineFrameInfo &MFI = MF.getFrameInfo();
1836 if (const int64_t NumBytes = MFI.getStackSize()) {
1837 TS->emitLocalDirective(Alignment: MFI.getMaxAlign(), Name: getFunctionFrameSymbol(),
1838 Size: NumBytes);
1839
1840 // Declare the frame pointers that NVPTXFrameLowering's prologue defines.
1841 const NVPTXRegisterInfo *NRI =
1842 MF.getSubtarget<NVPTXSubtarget>().getRegisterInfo();
1843 for (const Register FrameReg :
1844 {NRI->getFrameRegister(MF), NRI->getFrameLocalRegister(MF)})
1845 TS->emitRegDirective(
1846 SizeInBits: NRI->getRegSizeInBits(Reg: FrameReg, MRI: *MRI).getFixedValue(),
1847 Name: NVPTXInstPrinter::getRegisterName(Reg: FrameReg));
1848 }
1849
1850 // Go through all virtual registers to establish the mapping between the
1851 // global virtual
1852 // register number and the per class virtual register number.
1853 // We use the per class virtual register number in the ptx output.
1854 for (unsigned I : llvm::seq(Size: MRI->getNumVirtRegs())) {
1855 Register VR = Register::index2VirtReg(Index: I);
1856 if (MRI->use_empty(RegNo: VR) && MRI->def_empty(RegNo: VR))
1857 continue;
1858 auto &RCRegMap = VRegMapping[MRI->getRegClass(Reg: VR)];
1859 RCRegMap[VR] = RCRegMap.size() + 1;
1860 }
1861
1862 // Emit declaration of the virtual registers or 'physical' registers for
1863 // each register class
1864 const TargetRegisterInfo *TRI = MF.getSubtarget().getRegisterInfo();
1865 for (const TargetRegisterClass &RC : TRI->regclasses()) {
1866 // Only declare those registers that may be used.
1867 const auto It = VRegMapping.find(Val: &RC);
1868 if (It == VRegMapping.end() || It->second.empty())
1869 continue;
1870
1871 TS->emitRegDirective(
1872 SizeInBits: TRI->getRegSizeInBits(RC).getFixedValue(),
1873 Name: NVPTX::getVirtualRegisterPrefix(Kind: getVirtualRegisterKind(RC: &RC)),
1874 Count: It->second.size() + 1);
1875 }
1876}
1877
1878/// Translate virtual register numbers in DebugInfo locations to their printed
1879/// encodings, as used by CUDA-GDB.
1880void NVPTXAsmPrinter::encodeDebugInfoRegisterNumbers(
1881 const MachineFunction &MF) {
1882 const NVPTXSubtarget &STI = MF.getSubtarget<NVPTXSubtarget>();
1883 const NVPTXRegisterInfo *NRI = STI.getRegisterInfo();
1884
1885 // Clear the old mapping, and add the new one. This mapping is used after the
1886 // printing of the current function is complete, but before the next function
1887 // is printed.
1888 NRI->clearDebugRegisterMap();
1889
1890 for (const VRegMap &RegMap : make_second_range(c&: VRegMapping))
1891 for (const Register Reg : make_first_range(c: RegMap))
1892 NRI->addToDebugRegisterMap(VirtReg: Reg, RegisterName: getVirtualRegisterName(Reg));
1893}
1894
1895void NVPTXAsmPrinter::printFPConstant(const ConstantFP *Fp,
1896 raw_ostream &O) const {
1897 if (Fp->getType()->isFloatTy())
1898 O << "0f";
1899 else if (Fp->getType()->isDoubleTy())
1900 O << "0d";
1901 else
1902 llvm_unreachable("unsupported fp type");
1903
1904 const APInt API = Fp->getValueAPF().bitcastToAPInt();
1905 O << format_hex_no_prefix(N: API.getZExtValue(), Width: API.getBitWidth() / 4,
1906 /*Upper=*/true);
1907}
1908
1909void NVPTXAsmPrinter::printScalarConstant(const Constant *CPV, raw_ostream &O) {
1910 if (const ConstantInt *CI = dyn_cast<ConstantInt>(Val: CPV)) {
1911 O << CI->getValue();
1912 return;
1913 }
1914 if (const ConstantFP *CFP = dyn_cast<ConstantFP>(Val: CPV)) {
1915 const APInt API = CFP->getValueAPF().bitcastToAPInt();
1916 O << "0x"
1917 << format_hex_no_prefix(N: API.getZExtValue(), Width: API.getBitWidth() / 4,
1918 /*Upper=*/true);
1919 return;
1920 }
1921 if (isa<ConstantPointerNull>(Val: CPV)) {
1922 O << "0";
1923 return;
1924 }
1925 if (const GlobalValue *GVar = dyn_cast<GlobalValue>(Val: CPV)) {
1926 const bool IsNonGenericPointer = GVar->getAddressSpace() != 0;
1927 if (EmitGeneric && !isa<Function>(Val: CPV) && !IsNonGenericPointer) {
1928 O << "generic(";
1929 getSymbol(GV: GVar)->print(OS&: O, MAI);
1930 O << ")";
1931 } else {
1932 getSymbol(GV: GVar)->print(OS&: O, MAI);
1933 }
1934 return;
1935 }
1936 if (const ConstantExpr *Cexpr = dyn_cast<ConstantExpr>(Val: CPV)) {
1937 const MCExpr *E = lowerConstantForGV(CV: cast<Constant>(Val: Cexpr), ProcessingGeneric: false);
1938 printMCExpr(Expr: *E, OS&: O);
1939 return;
1940 }
1941 llvm_unreachable("Not scalar type found in printScalarConstant()");
1942}
1943
1944void NVPTXAsmPrinter::bufferLEByte(const Constant *CPV, int Bytes,
1945 AggBuffer *AggBuffer) {
1946 const DataLayout &DL = getDataLayout();
1947 int AllocSize = DL.getTypeAllocSize(Ty: CPV->getType());
1948 if (isa<UndefValue>(Val: CPV) || CPV->isNullValue()) {
1949 // Non-zero Bytes indicates that we need to zero-fill everything. Otherwise,
1950 // only the space allocated by CPV.
1951 AggBuffer->addZeros(Num: Bytes ? Bytes : AllocSize);
1952 return;
1953 }
1954
1955 // Helper for filling AggBuffer with APInts.
1956 auto AddIntToBuffer = [AggBuffer, Bytes](const APInt &Val) {
1957 size_t NumBytes = (Val.getBitWidth() + 7) / 8;
1958 SmallVector<unsigned char, 16> Buf(NumBytes);
1959 // `extractBitsAsZExtValue` does not allow the extraction of bits beyond the
1960 // input's bit width, and i1 arrays may not have a length that is a multuple
1961 // of 8. We handle the last byte separately, so we never request out of
1962 // bounds bits.
1963 for (unsigned I = 0; I < NumBytes - 1; ++I) {
1964 Buf[I] = Val.extractBitsAsZExtValue(numBits: 8, bitPosition: I * 8);
1965 }
1966 size_t LastBytePosition = (NumBytes - 1) * 8;
1967 size_t LastByteBits = Val.getBitWidth() - LastBytePosition;
1968 Buf[NumBytes - 1] =
1969 Val.extractBitsAsZExtValue(numBits: LastByteBits, bitPosition: LastBytePosition);
1970 AggBuffer->addBytes(Ptr: Buf.data(), Num: NumBytes, Bytes);
1971 };
1972
1973 switch (CPV->getType()->getTypeID()) {
1974 case Type::IntegerTyID:
1975 if (const auto *CI = dyn_cast<ConstantInt>(Val: CPV)) {
1976 AddIntToBuffer(CI->getValue());
1977 break;
1978 }
1979 if (const auto *Cexpr = dyn_cast<ConstantExpr>(Val: CPV)) {
1980 if (const auto *CI =
1981 dyn_cast<ConstantInt>(Val: ConstantFoldConstant(C: Cexpr, DL))) {
1982 AddIntToBuffer(CI->getValue());
1983 break;
1984 }
1985 if (Cexpr->getOpcode() == Instruction::PtrToInt) {
1986 Value *V = Cexpr->getOperand(i_nocapture: 0)->stripPointerCasts();
1987 AggBuffer->addSymbol(GVar: V, GVarBeforeStripping: Cexpr->getOperand(i_nocapture: 0));
1988 AggBuffer->addZeros(Num: AllocSize);
1989 break;
1990 }
1991 // A symbol-relative integer whose offset is applied outside the
1992 // ptrtoint, e.g. add(ptrtoint(@g), C). It can't fold to a ConstantInt
1993 // because it references a symbol; emit it through lowerConstantForGV, the
1994 // same path scalar symbol-relative integer globals use.
1995 AggBuffer->addSymbol(GVar: Cexpr, GVarBeforeStripping: Cexpr);
1996 AggBuffer->addZeros(Num: AllocSize);
1997 break;
1998 }
1999 llvm_unreachable("unsupported integer const type");
2000 break;
2001
2002 case Type::HalfTyID:
2003 case Type::BFloatTyID:
2004 case Type::FloatTyID:
2005 case Type::DoubleTyID:
2006 case Type::FP128TyID:
2007 AddIntToBuffer(cast<ConstantFP>(Val: CPV)->getValueAPF().bitcastToAPInt());
2008 break;
2009
2010 case Type::PointerTyID: {
2011 if (const GlobalValue *GVar = dyn_cast<GlobalValue>(Val: CPV)) {
2012 AggBuffer->addSymbol(GVar, GVarBeforeStripping: GVar);
2013 } else if (const ConstantExpr *Cexpr = dyn_cast<ConstantExpr>(Val: CPV)) {
2014 const Value *v = Cexpr->stripPointerCasts();
2015 AggBuffer->addSymbol(GVar: v, GVarBeforeStripping: Cexpr);
2016 }
2017 AggBuffer->addZeros(Num: AllocSize);
2018 break;
2019 }
2020
2021 case Type::ArrayTyID:
2022 case Type::FixedVectorTyID:
2023 case Type::StructTyID: {
2024 if (isa<ConstantAggregate>(Val: CPV) || isa<ConstantDataSequential>(Val: CPV)) {
2025 // bufferAggregateConstant doesn't emit tail-padding, i.e. it writes
2026 // `store_size` bytes, not `alloc_size` bytes. Do it ourselves here.
2027 unsigned StartPos = AggBuffer->getCurpos();
2028 bufferAggregateConstant(CV: CPV, aggBuffer: AggBuffer);
2029 unsigned Written = AggBuffer->getCurpos() - StartPos;
2030 unsigned SlotSize = std::max<int>(a: Bytes, b: AllocSize);
2031 if (SlotSize > Written)
2032 AggBuffer->addZeros(Num: SlotSize - Written);
2033 } else if (isa<ConstantAggregateZero>(Val: CPV))
2034 AggBuffer->addZeros(Num: Bytes);
2035 else
2036 llvm_unreachable("Unexpected Constant type");
2037 break;
2038 }
2039
2040 default:
2041 llvm_unreachable("unsupported type");
2042 }
2043}
2044
2045void NVPTXAsmPrinter::bufferAggregateConstant(const Constant *CPV,
2046 AggBuffer *aggBuffer) {
2047 const DataLayout &DL = getDataLayout();
2048
2049 auto ExtendBuffer = [](APInt Val, AggBuffer *Buffer) {
2050 unsigned NumBytes = divideCeil(Numerator: Val.getBitWidth(), Denominator: 8);
2051 for (unsigned I : llvm::seq(Size: NumBytes)) {
2052 unsigned NumBits = std::min(a: 8u, b: Val.getBitWidth() - I * 8);
2053 Buffer->addByte(Byte: Val.extractBitsAsZExtValue(numBits: NumBits, bitPosition: I * 8));
2054 }
2055 };
2056
2057 // Integer or floating point vector splats.
2058 if (isa<ConstantInt, ConstantFP>(Val: CPV)) {
2059 if (auto *VTy = dyn_cast<FixedVectorType>(Val: CPV->getType())) {
2060 for (unsigned I : llvm::seq(Size: VTy->getNumElements()))
2061 bufferLEByte(CPV: CPV->getAggregateElement(Elt: I), Bytes: 0, AggBuffer: aggBuffer);
2062 return;
2063 }
2064 }
2065
2066 // Integers of arbitrary width
2067 if (const ConstantInt *CI = dyn_cast<ConstantInt>(Val: CPV)) {
2068 assert(CI->getType()->isIntegerTy() && "Expected integer constant!");
2069 ExtendBuffer(CI->getValue(), aggBuffer);
2070 return;
2071 }
2072
2073 // f128
2074 if (const ConstantFP *CFP = dyn_cast<ConstantFP>(Val: CPV)) {
2075 assert(CFP->getType()->isFloatingPointTy() && "Expected fp constant!");
2076 if (CFP->getType()->isFP128Ty()) {
2077 ExtendBuffer(CFP->getValueAPF().bitcastToAPInt(), aggBuffer);
2078 return;
2079 }
2080 }
2081
2082 // Buffer arrays one element at a time.
2083 if (isa<ConstantArray>(Val: CPV)) {
2084 for (const auto &Op : CPV->operands())
2085 bufferLEByte(CPV: cast<Constant>(Val: Op), Bytes: 0, AggBuffer: aggBuffer);
2086 return;
2087 }
2088
2089 // Constant vectors
2090 if (const auto *CVec = dyn_cast<ConstantVector>(Val: CPV)) {
2091 bufferAggregateConstVec(CV: CVec, aggBuffer);
2092 return;
2093 }
2094
2095 if (const auto *CDS = dyn_cast<ConstantDataSequential>(Val: CPV)) {
2096 for (unsigned I : llvm::seq(Size: CDS->getNumElements()))
2097 bufferLEByte(CPV: cast<Constant>(Val: CDS->getElementAsConstant(i: I)), Bytes: 0, AggBuffer: aggBuffer);
2098 return;
2099 }
2100
2101 if (isa<ConstantStruct>(Val: CPV)) {
2102 if (CPV->getNumOperands()) {
2103 StructType *ST = cast<StructType>(Val: CPV->getType());
2104 for (unsigned I : llvm::seq(Size: CPV->getNumOperands())) {
2105 int EndOffset = (I + 1 == CPV->getNumOperands())
2106 ? DL.getStructLayout(Ty: ST)->getElementOffset(Idx: 0) +
2107 DL.getTypeAllocSize(Ty: ST)
2108 : DL.getStructLayout(Ty: ST)->getElementOffset(Idx: I + 1);
2109 int Bytes = EndOffset - DL.getStructLayout(Ty: ST)->getElementOffset(Idx: I);
2110 bufferLEByte(CPV: cast<Constant>(Val: CPV->getOperand(i: I)), Bytes, AggBuffer: aggBuffer);
2111 }
2112 }
2113 return;
2114 }
2115 llvm_unreachable("unsupported constant type in printAggregateConstant()");
2116}
2117
2118void NVPTXAsmPrinter::bufferAggregateConstVec(const ConstantVector *CV,
2119 AggBuffer *aggBuffer) {
2120 unsigned NumElems = CV->getType()->getNumElements();
2121 const unsigned BuffSize = aggBuffer->getBufferSize();
2122
2123 // Buffer one element at a time if we have allocated enough buffer space.
2124 if (BuffSize >= NumElems) {
2125 for (const auto &Op : CV->operands())
2126 bufferLEByte(CPV: cast<Constant>(Val: Op), Bytes: 0, AggBuffer: aggBuffer);
2127 return;
2128 }
2129
2130 // Sub-byte datatypes will have more elements than bytes allocated for the
2131 // buffer. Merge consecutive elements to form a full byte. We expect that 8 %
2132 // sub-byte-elem-size should be 0 and current expected usage is for i4 (for
2133 // e2m1-fp4 types).
2134 Type *ElemTy = CV->getType()->getElementType();
2135 assert(ElemTy->isIntegerTy() && "Expected integer data type.");
2136 unsigned ElemTySize = ElemTy->getPrimitiveSizeInBits();
2137 assert(ElemTySize < 8 && "Expected sub-byte data type.");
2138 assert(8 % ElemTySize == 0 && "Element type size must evenly divide a byte.");
2139 // Number of elements to merge to form a full byte.
2140 unsigned NumElemsPerByte = 8 / ElemTySize;
2141 unsigned NumCompleteBytes = NumElems / NumElemsPerByte;
2142 unsigned NumTailElems = NumElems % NumElemsPerByte;
2143
2144 // Helper lambda to constant-fold sub-vector of sub-byte type elements into
2145 // i8. Start and end indices of the sub-vector is provided, along with number
2146 // of padding zeros if required.
2147 auto ConvertSubCVtoInt8 = [this, &ElemTy](const ConstantVector *CV,
2148 unsigned Start, unsigned End,
2149 unsigned NumPaddingZeros = 0) {
2150 // Collect elements to create sub-vector.
2151 SmallVector<Constant *, 8> SubCVElems;
2152 for (unsigned I : llvm::seq(Begin: Start, End))
2153 SubCVElems.push_back(Elt: CV->getAggregateElement(Elt: I));
2154
2155 // Optionally pad with zeros.
2156 if (NumPaddingZeros)
2157 SubCVElems.append(NumInputs: NumPaddingZeros, Elt: ConstantInt::getNullValue(Ty: ElemTy));
2158
2159 auto SubCV = ConstantVector::get(V: SubCVElems);
2160 Type *Int8Ty = IntegerType::get(C&: SubCV->getContext(), NumBits: 8);
2161
2162 // Merge elements of the sub-vector using ConstantFolding.
2163 ConstantInt *MergedElem =
2164 dyn_cast_or_null<ConstantInt>(Val: ConstantFoldConstant(
2165 C: ConstantExpr::getBitCast(C: const_cast<Constant *>(SubCV), Ty: Int8Ty),
2166 DL: getDataLayout()));
2167
2168 if (!MergedElem)
2169 report_fatal_error(
2170 reason: "Cannot lower vector global with unusual element type");
2171
2172 return MergedElem;
2173 };
2174
2175 // Iterate through elements of vector one chunk at a time and buffer that
2176 // chunk.
2177 for (unsigned ByteIdx : llvm::seq(Size: NumCompleteBytes))
2178 bufferLEByte(CPV: ConvertSubCVtoInt8(CV, ByteIdx * NumElemsPerByte,
2179 (ByteIdx + 1) * NumElemsPerByte),
2180 Bytes: 0, AggBuffer: aggBuffer);
2181
2182 // For unevenly sized vectors add tail padding zeros.
2183 if (NumTailElems > 0)
2184 bufferLEByte(CPV: ConvertSubCVtoInt8(CV, NumElems - NumTailElems, NumElems,
2185 NumElemsPerByte - NumTailElems),
2186 Bytes: 0, AggBuffer: aggBuffer);
2187}
2188
2189/// lowerConstantForGV - Return an MCExpr for the given Constant. This is mostly
2190/// a copy from AsmPrinter::lowerConstant, except customized to only handle
2191/// expressions that are representable in PTX and create
2192/// NVPTXGenericMCSymbolRefExpr nodes for addrspacecast instructions.
2193const MCExpr *
2194NVPTXAsmPrinter::lowerConstantForGV(const Constant *CV,
2195 bool ProcessingGeneric) const {
2196 MCContext &Ctx = OutContext;
2197
2198 if (CV->isNullValue() || isa<UndefValue>(Val: CV))
2199 return MCConstantExpr::create(Value: 0, Ctx);
2200
2201 if (const ConstantInt *CI = dyn_cast<ConstantInt>(Val: CV))
2202 return MCConstantExpr::create(Value: CI->getZExtValue(), Ctx);
2203
2204 if (const GlobalValue *GV = dyn_cast<GlobalValue>(Val: CV)) {
2205 const MCSymbolRefExpr *Expr = MCSymbolRefExpr::create(Symbol: getSymbol(GV), Ctx);
2206 if (ProcessingGeneric)
2207 return NVPTXGenericMCSymbolRefExpr::create(SymExpr: Expr, Ctx);
2208 return Expr;
2209 }
2210
2211 const ConstantExpr *CE = dyn_cast<ConstantExpr>(Val: CV);
2212 if (!CE) {
2213 llvm_unreachable("Unknown constant value to lower!");
2214 }
2215
2216 switch (CE->getOpcode()) {
2217 default:
2218 break; // Error
2219
2220 case Instruction::AddrSpaceCast: {
2221 // Strip the addrspacecast and pass along the operand
2222 PointerType *DstTy = cast<PointerType>(Val: CE->getType());
2223 if (DstTy->getAddressSpace() == 0)
2224 return lowerConstantForGV(CV: cast<const Constant>(Val: CE->getOperand(i_nocapture: 0)), ProcessingGeneric: true);
2225
2226 break; // Error
2227 }
2228
2229 case Instruction::GetElementPtr: {
2230 const DataLayout &DL = getDataLayout();
2231
2232 // Generate a symbolic expression for the byte address
2233 APInt OffsetAI(DL.getPointerTypeSizeInBits(CE->getType()), 0);
2234 cast<GEPOperator>(Val: CE)->accumulateConstantOffset(DL, Offset&: OffsetAI);
2235
2236 const MCExpr *Base = lowerConstantForGV(CV: CE->getOperand(i_nocapture: 0),
2237 ProcessingGeneric);
2238 if (!OffsetAI)
2239 return Base;
2240
2241 int64_t Offset = OffsetAI.getSExtValue();
2242 return MCBinaryExpr::createAdd(LHS: Base, RHS: MCConstantExpr::create(Value: Offset, Ctx),
2243 Ctx);
2244 }
2245
2246 case Instruction::Trunc:
2247 // We emit the value and depend on the assembler to truncate the generated
2248 // expression properly. This is important for differences between
2249 // blockaddress labels. Since the two labels are in the same function, it
2250 // is reasonable to treat their delta as a 32-bit value.
2251 [[fallthrough]];
2252 case Instruction::BitCast:
2253 return lowerConstantForGV(CV: CE->getOperand(i_nocapture: 0), ProcessingGeneric);
2254
2255 case Instruction::IntToPtr: {
2256 const DataLayout &DL = getDataLayout();
2257
2258 // Handle casts to pointers by changing them into casts to the appropriate
2259 // integer type. This promotes constant folding and simplifies this code.
2260 Constant *Op = CE->getOperand(i_nocapture: 0);
2261 Op = ConstantFoldIntegerCast(C: Op, DestTy: DL.getIntPtrType(CV->getType()),
2262 /*IsSigned*/ false, DL);
2263 if (Op)
2264 return lowerConstantForGV(CV: Op, ProcessingGeneric);
2265
2266 break; // Error
2267 }
2268
2269 case Instruction::PtrToInt: {
2270 const DataLayout &DL = getDataLayout();
2271
2272 // Support only foldable casts to/from pointers that can be eliminated by
2273 // changing the pointer to the appropriately sized integer type.
2274 Constant *Op = CE->getOperand(i_nocapture: 0);
2275 Type *Ty = CE->getType();
2276
2277 const MCExpr *OpExpr = lowerConstantForGV(CV: Op, ProcessingGeneric);
2278
2279 // We can emit the pointer value into this slot if the slot is an
2280 // integer slot equal to the size of the pointer.
2281 if (DL.getTypeAllocSize(Ty) == DL.getTypeAllocSize(Ty: Op->getType()))
2282 return OpExpr;
2283
2284 // Otherwise the pointer is smaller than the resultant integer, mask off
2285 // the high bits so we are sure to get a proper truncation if the input is
2286 // a constant expr.
2287 unsigned InBits = DL.getTypeAllocSizeInBits(Ty: Op->getType());
2288 const MCExpr *MaskExpr = MCConstantExpr::create(Value: ~0ULL >> (64-InBits), Ctx);
2289 return MCBinaryExpr::createAnd(LHS: OpExpr, RHS: MaskExpr, Ctx);
2290 }
2291
2292 // The MC library also has a right-shift operator, but it isn't consistently
2293 // signed or unsigned between different targets.
2294 case Instruction::Add: {
2295 const MCExpr *LHS = lowerConstantForGV(CV: CE->getOperand(i_nocapture: 0), ProcessingGeneric);
2296 const MCExpr *RHS = lowerConstantForGV(CV: CE->getOperand(i_nocapture: 1), ProcessingGeneric);
2297 switch (CE->getOpcode()) {
2298 default: llvm_unreachable("Unknown binary operator constant cast expr");
2299 case Instruction::Add: return MCBinaryExpr::createAdd(LHS, RHS, Ctx);
2300 }
2301 }
2302 }
2303
2304 // If the code isn't optimized, there may be outstanding folding
2305 // opportunities. Attempt to fold the expression using DataLayout as a
2306 // last resort before giving up.
2307 Constant *C = ConstantFoldConstant(C: CE, DL: getDataLayout());
2308 if (C != CE)
2309 return lowerConstantForGV(CV: C, ProcessingGeneric);
2310
2311 // Otherwise report the problem to the user.
2312 std::string S;
2313 raw_string_ostream OS(S);
2314 OS << "Unsupported expression in static initializer: ";
2315 CE->printAsOperand(O&: OS, /*PrintType=*/false,
2316 M: !MF ? nullptr : MF->getFunction().getParent());
2317 report_fatal_error(reason: Twine(OS.str()));
2318}
2319
2320void NVPTXAsmPrinter::printMCExpr(const MCExpr &Expr, raw_ostream &OS) const {
2321 OutContext.getAsmInfo().printExpr(OS, Expr);
2322}
2323
2324/// PrintAsmOperand - Print out an operand for an inline asm expression.
2325///
2326bool NVPTXAsmPrinter::PrintAsmOperand(const MachineInstr *MI, unsigned OpNo,
2327 const char *ExtraCode, raw_ostream &O) {
2328 if (ExtraCode && ExtraCode[0]) {
2329 if (ExtraCode[1] != 0)
2330 return true; // Unknown modifier.
2331
2332 switch (ExtraCode[0]) {
2333 default:
2334 // See if this is a generic print operand
2335 return AsmPrinter::PrintAsmOperand(MI, OpNo, ExtraCode, OS&: O);
2336 case 'r':
2337 break;
2338 }
2339 }
2340
2341 printOperand(MI, OpNum: OpNo, O);
2342
2343 return false;
2344}
2345
2346bool NVPTXAsmPrinter::PrintAsmMemoryOperand(const MachineInstr *MI,
2347 unsigned OpNo,
2348 const char *ExtraCode,
2349 raw_ostream &O) {
2350 if (ExtraCode && ExtraCode[0])
2351 return true; // Unknown modifier
2352
2353 O << '[';
2354 printMemOperand(MI, OpNum: OpNo, O);
2355 O << ']';
2356
2357 return false;
2358}
2359
2360void NVPTXAsmPrinter::printOperand(const MachineInstr *MI, unsigned OpNum,
2361 raw_ostream &O) {
2362 const MachineOperand &MO = MI->getOperand(i: OpNum);
2363 switch (MO.getType()) {
2364 case MachineOperand::MO_Register:
2365 if (MO.getReg().isPhysical()) {
2366 if (MO.getReg() == NVPTX::VRDepot)
2367 getFunctionFrameSymbol()->print(OS&: O, MAI);
2368 else
2369 O << NVPTXInstPrinter::getRegisterName(Reg: MO.getReg());
2370 } else {
2371 O << getVirtualRegisterName(Reg: MO.getReg());
2372 }
2373 break;
2374
2375 case MachineOperand::MO_Immediate:
2376 O << MO.getImm();
2377 break;
2378
2379 case MachineOperand::MO_FPImmediate:
2380 printFPConstant(Fp: MO.getFPImm(), O);
2381 break;
2382
2383 case MachineOperand::MO_GlobalAddress:
2384 PrintSymbolOperand(MO, OS&: O);
2385 break;
2386
2387 case MachineOperand::MO_MCSymbol:
2388 MO.getMCSymbol()->print(OS&: O, MAI);
2389 break;
2390
2391 case MachineOperand::MO_MachineBasicBlock:
2392 MO.getMBB()->getSymbol()->print(OS&: O, MAI);
2393 break;
2394
2395 default:
2396 llvm_unreachable("Operand type not supported.");
2397 }
2398}
2399
2400void NVPTXAsmPrinter::printMemOperand(const MachineInstr *MI, unsigned OpNum,
2401 raw_ostream &O, const char *Modifier) {
2402 printOperand(MI, OpNum, O);
2403
2404 if (Modifier && strcmp(s1: Modifier, s2: "add") == 0) {
2405 O << ", ";
2406 printOperand(MI, OpNum: OpNum + 1, O);
2407 } else {
2408 if (MI->getOperand(i: OpNum + 1).isImm() &&
2409 MI->getOperand(i: OpNum + 1).getImm() == 0)
2410 return; // don't print ',0' or '+0'
2411 O << "+";
2412 printOperand(MI, OpNum: OpNum + 1, O);
2413 }
2414}
2415
2416/// Returns true if \p Line begins with an alphabetic character or underscore,
2417/// indicating it is a PTX instruction that should receive a .loc directive.
2418static bool isPTXInstruction(StringRef Line) {
2419 StringRef Trimmed = Line.ltrim();
2420 return !Trimmed.empty() &&
2421 (std::isalpha(static_cast<unsigned char>(Trimmed[0])) ||
2422 Trimmed[0] == '_');
2423}
2424
2425/// Returns the DILocation for an inline asm MachineInstr if debug line info
2426/// should be emitted, or nullptr otherwise.
2427static const DILocation *getInlineAsmDebugLoc(const MachineInstr *MI) {
2428 if (!MI || !MI->getDebugLoc())
2429 return nullptr;
2430 const DISubprogram *SP = MI->getMF()->getFunction().getSubprogram();
2431 if (!SP || SP->getUnit()->getEmissionKind() == DICompileUnit::NoDebug)
2432 return nullptr;
2433 const DILocation *DL = MI->getDebugLoc();
2434 if (!DL->getFile() || !DL->getLine() || DL->isImplicitCode())
2435 return nullptr;
2436 return DL;
2437}
2438
2439namespace {
2440struct InlineAsmInliningContext {
2441 MCSymbol *FuncNameSym = nullptr;
2442 unsigned FileIA = 0;
2443 unsigned LineIA = 0;
2444 unsigned ColIA = 0;
2445
2446 bool hasInlinedAt() const { return FuncNameSym != nullptr; }
2447};
2448} // namespace
2449
2450/// Resolves the enhanced-lineinfo inlining context for an inline asm debug
2451/// location. Returns a default (empty) context if inlining info is unavailable.
2452static InlineAsmInliningContext
2453getInlineAsmInliningContext(const DILocation *DL, const MachineFunction &MF,
2454 NVPTXDwarfDebug *NVDD, MCStreamer &Streamer,
2455 unsigned CUID) {
2456 InlineAsmInliningContext Ctx;
2457 const DILocation *InlinedAt = DL->getInlinedAt();
2458 if (!InlinedAt || !InlinedAt->getFile() || !NVDD ||
2459 !NVDD->isEnhancedLineinfo(MF))
2460 return Ctx;
2461 const auto *SubProg = getDISubprogram(Scope: DL->getScope());
2462 if (!SubProg)
2463 return Ctx;
2464 Ctx.FuncNameSym = NVDD->getOrCreateFuncNameSymbol(LinkageName: SubProg->getLinkageName());
2465 Ctx.FileIA = Streamer.emitDwarfFileDirective(
2466 FileNo: 0, Directory: InlinedAt->getFile()->getDirectory(),
2467 Filename: InlinedAt->getFile()->getFilename(), Checksum: std::nullopt, Source: std::nullopt, CUID);
2468 Ctx.LineIA = InlinedAt->getLine();
2469 Ctx.ColIA = InlinedAt->getColumn();
2470 return Ctx;
2471}
2472
2473void NVPTXAsmPrinter::emitInlineAsm(StringRef Str, const MCSubtargetInfo &STI,
2474 const MCTargetOptions &MCOptions,
2475 const MDNode *LocMDNode,
2476 InlineAsm::AsmDialect Dialect,
2477 const MachineInstr *MI) {
2478 assert(!Str.empty() && "Can't emit empty inline asm block");
2479 if (Str.back() == 0)
2480 Str = Str.substr(Start: 0, N: Str.size() - 1);
2481
2482 auto emitAsmStr = [&](StringRef AsmStr) {
2483 emitInlineAsmStart();
2484 OutStreamer->emitRawText(String: AsmStr);
2485 emitInlineAsmEnd(StartInfo: STI, EndInfo: nullptr, MI);
2486 };
2487
2488 const DILocation *DL = getInlineAsmDebugLoc(MI);
2489 if (!DL) {
2490 emitAsmStr(Str);
2491 return;
2492 }
2493
2494 const DIFile *File = DL->getFile();
2495 unsigned Line = DL->getLine();
2496 const unsigned Column = DL->getColumn();
2497 const unsigned CUID = OutStreamer->getContext().getDwarfCompileUnitID();
2498 const unsigned FileNumber = OutStreamer->emitDwarfFileDirective(
2499 FileNo: 0, Directory: File->getDirectory(), Filename: File->getFilename(), Checksum: std::nullopt, Source: std::nullopt,
2500 CUID);
2501
2502 auto *NVDD = static_cast<NVPTXDwarfDebug *>(getDwarfDebug());
2503 InlineAsmInliningContext InlineCtx =
2504 getInlineAsmInliningContext(DL, MF: *MI->getMF(), NVDD, Streamer&: *OutStreamer, CUID);
2505
2506 SmallVector<StringRef, 16> Lines;
2507 Str.split(A&: Lines, Separator: '\n');
2508 emitInlineAsmStart();
2509 for (const StringRef &L : Lines) {
2510 StringRef RTrimmed = L.rtrim(Char: '\r');
2511 if (isPTXInstruction(Line: L)) {
2512 if (InlineCtx.hasInlinedAt()) {
2513 OutStreamer->emitDwarfLocDirectiveWithInlinedAt(
2514 FileNo: FileNumber, Line, Column, FileIA: InlineCtx.FileIA, LineIA: InlineCtx.LineIA,
2515 ColumnIA: InlineCtx.ColIA, Sym: InlineCtx.FuncNameSym, DWARF2_FLAG_IS_STMT, Isa: 0, Discriminator: 0,
2516 FileName: File->getFilename());
2517 } else {
2518 OutStreamer->emitDwarfLocDirective(FileNo: FileNumber, Line, Column,
2519 DWARF2_FLAG_IS_STMT, Isa: 0, Discriminator: 0,
2520 FileName: File->getFilename());
2521 }
2522 }
2523 OutStreamer->emitRawText(String: RTrimmed);
2524 ++Line;
2525 }
2526 emitInlineAsmEnd(StartInfo: STI, EndInfo: nullptr, MI);
2527}
2528
2529char NVPTXAsmPrinter::ID = 0;
2530
2531INITIALIZE_PASS(NVPTXAsmPrinter, "nvptx-asm-printer", "NVPTX Assembly Printer",
2532 false, false)
2533
2534// Force static initialization.
2535extern "C" LLVM_ABI LLVM_EXTERNAL_VISIBILITY void
2536LLVMInitializeNVPTXAsmPrinter() {
2537 RegisterAsmPrinter<NVPTXAsmPrinter> X(getTheNVPTXTarget32());
2538 RegisterAsmPrinter<NVPTXAsmPrinter> Y(getTheNVPTXTarget64());
2539}
2540
2541PreservedAnalyses NVPTXAsmPrinterBeginPass::run(Module &M,
2542 ModuleAnalysisManager &MAM) {
2543 AsmPrinter &Printer = MAM.getResult<AsmPrinterAnalysis>(IR&: M).getPrinter();
2544 setupModuleAsmPrinter(M, MAM, AsmPrinter&: Printer);
2545 Printer.doInitialization(M);
2546 return PreservedAnalyses::all();
2547}
2548
2549PreservedAnalyses
2550NVPTXAsmPrinterPass::run(MachineFunction &MF,
2551 MachineFunctionAnalysisManager &MFAM) {
2552 AsmPrinter &Printer =
2553 MFAM.getResult<ModuleAnalysisManagerMachineFunctionProxy>(IR&: MF)
2554 .getCachedResult<AsmPrinterAnalysis>(IR&: *MF.getFunction().getParent())
2555 ->getPrinter();
2556 setupMachineFunctionAsmPrinter(MFAM, MF, AsmPrinter&: Printer);
2557 Printer.runOnMachineFunction(MF);
2558 return PreservedAnalyses::all();
2559}
2560
2561PreservedAnalyses NVPTXAsmPrinterEndPass::run(Module &M,
2562 ModuleAnalysisManager &MAM) {
2563 AsmPrinter &Printer = MAM.getResult<AsmPrinterAnalysis>(IR&: M).getPrinter();
2564 setupModuleAsmPrinter(M, MAM, AsmPrinter&: Printer);
2565 Printer.doFinalization(M);
2566 return PreservedAnalyses::all();
2567}
2568