1//===-- NVPTXMarkKernelPtrsGlobal.cpp - Mark kernel pointers as global ----===//
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// For CUDA kernels, pointers loaded from byval parameters are known to be in
10// global address space. This pass inserts addrspacecast pairs to make that
11// explicit, enabling later address-space inference to propagate the global AS.
12// It also handles the pattern where a pointer is loaded as an integer and then
13// converted via inttoptr.
14//
15//===----------------------------------------------------------------------===//
16
17#include "NVPTX.h"
18#include "NVVMProperties.h"
19#include "llvm/Analysis/ValueTracking.h"
20#include "llvm/IR/InstIterator.h"
21#include "llvm/IR/Instructions.h"
22#include "llvm/Pass.h"
23#include "llvm/Support/NVPTXAddrSpace.h"
24
25using namespace llvm;
26using namespace NVPTXAS;
27
28static void markPointerAsAS(Value *Ptr, unsigned AS) {
29 if (Ptr->getType()->getPointerAddressSpace() != ADDRESS_SPACE_GENERIC)
30 return;
31
32 BasicBlock::iterator InsertPt;
33 if (auto *Arg = dyn_cast<Argument>(Val: Ptr)) {
34 InsertPt = Arg->getParent()->getEntryBlock().begin();
35 } else {
36 InsertPt = ++cast<Instruction>(Val: Ptr)->getIterator();
37 assert(InsertPt != InsertPt->getParent()->end() &&
38 "We don't call this function with Ptr being a terminator.");
39 }
40
41 Instruction *PtrInGlobal = new AddrSpaceCastInst(
42 Ptr, PointerType::get(C&: Ptr->getContext(), AddressSpace: AS), Ptr->getName(), InsertPt);
43 Value *PtrInGeneric = new AddrSpaceCastInst(PtrInGlobal, Ptr->getType(),
44 Ptr->getName(), InsertPt);
45 Ptr->replaceAllUsesWith(V: PtrInGeneric);
46 PtrInGlobal->setOperand(i: 0, Val: Ptr);
47}
48
49static void markPointerAsGlobal(Value *Ptr) {
50 markPointerAsAS(Ptr, AS: ADDRESS_SPACE_GLOBAL);
51}
52
53static void handleIntToPtr(Value &V) {
54 if (!all_of(Range: V.users(), P: [](User *U) { return isa<IntToPtrInst>(Val: U); }))
55 return;
56
57 SmallVector<User *, 16> UsersToUpdate(V.users());
58 for (User *U : UsersToUpdate)
59 markPointerAsGlobal(Ptr: U);
60}
61
62static bool markKernelPtrsGlobal(Function &F) {
63 if (!isKernelFunction(F))
64 return false;
65
66 // Copying of byval aggregates + SROA may result in pointers being loaded as
67 // integers, followed by inttoptr. We mark those as global too, but only if
68 // the loaded integer is used exclusively for conversion to a pointer.
69 for (auto &I : instructions(F)) {
70 auto *LI = dyn_cast<LoadInst>(Val: &I);
71 if (!LI)
72 continue;
73
74 if (LI->getType()->isPointerTy() || LI->getType()->isIntegerTy()) {
75 Value *UO = getUnderlyingObject(V: LI->getPointerOperand());
76 if (auto *Arg = dyn_cast<Argument>(Val: UO)) {
77 if (Arg->hasByValAttr()) {
78 if (LI->getType()->isPointerTy())
79 markPointerAsGlobal(Ptr: LI);
80 else
81 handleIntToPtr(V&: *LI);
82 }
83 }
84 }
85 }
86
87 for (Argument &Arg : F.args())
88 if (Arg.getType()->isIntegerTy())
89 handleIntToPtr(V&: Arg);
90
91 return true;
92}
93
94namespace {
95
96class NVPTXMarkKernelPtrsGlobalLegacyPass : public FunctionPass {
97public:
98 static char ID;
99 NVPTXMarkKernelPtrsGlobalLegacyPass() : FunctionPass(ID) {}
100 bool runOnFunction(Function &F) override;
101};
102
103} // namespace
104
105INITIALIZE_PASS(NVPTXMarkKernelPtrsGlobalLegacyPass,
106 "nvptx-mark-kernel-ptrs-global",
107 "NVPTX Mark Kernel Pointers Global", false, false)
108
109bool NVPTXMarkKernelPtrsGlobalLegacyPass::runOnFunction(Function &F) {
110 return markKernelPtrsGlobal(F);
111}
112
113char NVPTXMarkKernelPtrsGlobalLegacyPass::ID = 0;
114
115FunctionPass *llvm::createNVPTXMarkKernelPtrsGlobalPass() {
116 return new NVPTXMarkKernelPtrsGlobalLegacyPass();
117}
118
119PreservedAnalyses
120NVPTXMarkKernelPtrsGlobalPass::run(Function &F, FunctionAnalysisManager &) {
121 if (!markKernelPtrsGlobal(F))
122 return PreservedAnalyses::all();
123 return PreservedAnalyses::none().preserveSet<CFGAnalyses>();
124}
125