Proteus
Programmable JIT compilation and optimization for C/C++ using LLVM
Loading...
Searching...
No Matches
LambdaRegistry.h
Go to the documentation of this file.
1#ifndef PROTEUS_LAMBDA_INTERFACE_H
2#define PROTEUS_LAMBDA_INTERFACE_H
3
5#include "proteus/Error.h"
10
11#include <llvm/ADT/DenseMap.h>
12#include <llvm/ADT/STLExtras.h>
13#include <llvm/ADT/SmallVector.h>
14#include <llvm/ADT/StringRef.h>
15#include <llvm/Demangle/Demangle.h>
16#include <mutex>
17#include <optional>
18
19namespace proteus {
20
21using namespace llvm;
22
23// The LambdaRegistry stores the unique lambda type symbol and the values of Jit
24// member variables in a map for retrieval by the Jit engines.
26public:
28 static LambdaRegistry Singleton;
29 return Singleton;
30 }
31
32 LambdaRegistry(const LambdaRegistry &) = delete;
36
37 using JitVariantMap = DenseMap<int32_t, RuntimeConstant>;
38 using JitVariantVec = SmallVector<JitVariantMap, 4>;
43
44 std::optional<JitVariantMap> getHostJitVariant(uint64_t FunctorID) const {
45 const auto &CurrentVariants = currentHostJitVariables();
46 if (CurrentVariants.empty())
47 return std::nullopt;
48
49 auto It = CurrentVariants.find(FunctorID);
50 if (It == CurrentVariants.end() || It->second.empty())
51 return std::nullopt;
52
53 return It->second;
54 }
55
56 void appendHostJitVariable(uint64_t FunctorID, const RuntimeConstant &RC) {
57 pendingHostJitVariables()[FunctorID][RC.Pos] = RC;
58 }
59
60 void commitHostJitVariables(uint64_t FunctorID) {
61 auto &PendingVariants = pendingHostJitVariables();
62 auto PendingIt = PendingVariants.find(FunctorID);
63 if (PendingIt == PendingVariants.end())
64 return;
65 if (PendingIt->second.empty())
66 return;
67
68 currentHostJitVariables()[FunctorID] = std::move(PendingIt->second);
69 PendingVariants.erase(PendingIt);
70 }
71
72 void eraseHostJitVariables(uint64_t FunctorID) {
73 currentHostJitVariables().erase(FunctorID);
74 }
75
76 bool emptyHost() const { return currentHostJitVariables().empty(); }
77
79 auto &PendingLaunchInfo = pendingDeviceLaunchInfo();
80 auto &CurrentLaunchInfo = currentDeviceLaunchInfo();
81 PendingLaunchInfo.LambdaCalleeInfo.clear();
82 PendingLaunchInfo.CallsiteRuntimeConstants.clear();
83 CurrentLaunchInfo.LambdaCalleeInfo.clear();
84 CurrentLaunchInfo.CallsiteRuntimeConstants.clear();
85 }
86
88 uint32_t CallsiteIndex,
89 const RuntimeConstant &RC) {
90 auto &PendingLaunchInfo = pendingDeviceLaunchInfo();
91 auto &RCVec = PendingLaunchInfo.CallsiteRuntimeConstants[CallsiteIndex];
92 RCVec.push_back(RC);
93 if (llvm::find(PendingLaunchInfo.LambdaCalleeInfo, LambdaID) ==
94 PendingLaunchInfo.LambdaCalleeInfo.end()) {
95 PendingLaunchInfo.LambdaCalleeInfo.push_back(LambdaID);
96 }
97 }
98
100 auto &PendingLaunchInfo = pendingDeviceLaunchInfo();
101 auto &CurrentLaunchInfo = currentDeviceLaunchInfo();
102 for (auto &KV : PendingLaunchInfo.CallsiteRuntimeConstants) {
103 llvm::sort(KV.second,
104 [](const RuntimeConstant &L, const RuntimeConstant &R) {
105 return L.Pos < R.Pos;
106 });
107 }
108 CurrentLaunchInfo = std::move(PendingLaunchInfo);
109 PendingLaunchInfo.LambdaCalleeInfo.clear();
110 PendingLaunchInfo.CallsiteRuntimeConstants.clear();
111 }
112
114 auto &CurrentLaunchInfo = currentDeviceLaunchInfo();
115 DeviceLaunchInfo LaunchInfo = std::move(CurrentLaunchInfo);
116 CurrentLaunchInfo.LambdaCalleeInfo.clear();
117 CurrentLaunchInfo.CallsiteRuntimeConstants.clear();
118 return LaunchInfo;
119 }
120
122 void *RegistrationFunc) {
123 std::lock_guard<std::mutex> Lock(DeviceRegistrationMutex);
124 KernelToLambdaRegistration[Kernel] = RegistrationFunc;
125 }
126
128 void *RegistrationFunc = nullptr;
129 {
130 std::lock_guard<std::mutex> Lock(DeviceRegistrationMutex);
131 auto It = KernelToLambdaRegistration.find(Kernel);
132 if (It != KernelToLambdaRegistration.end())
133 RegistrationFunc = It->second;
134 }
135 if (!RegistrationFunc) {
137 return;
138 }
139 auto RegisterFunc = reinterpret_cast<void (*)(void **)>(RegistrationFunc);
140 RegisterFunc(Args);
141 }
142
143private:
144 explicit LambdaRegistry() = default;
145
146 static DenseMap<uint64_t, JitVariantMap> &pendingHostJitVariables() {
147 static thread_local DenseMap<uint64_t, JitVariantMap> PendingHostJitVars;
148 return PendingHostJitVars;
149 }
150
151 static DenseMap<uint64_t, JitVariantMap> &currentHostJitVariables() {
152 static thread_local DenseMap<uint64_t, JitVariantMap> CurrentHostJitVars;
153 return CurrentHostJitVars;
154 }
155
156 static DeviceLaunchInfo &pendingDeviceLaunchInfo() {
157 static thread_local DeviceLaunchInfo PendingLaunchInfo;
158 return PendingLaunchInfo;
159 }
160
161 static DeviceLaunchInfo &currentDeviceLaunchInfo() {
162 static thread_local DeviceLaunchInfo CurrentLaunchInfo;
163 return CurrentLaunchInfo;
164 }
165
166 mutable std::mutex DeviceRegistrationMutex;
167 DenseMap<void *, void *> KernelToLambdaRegistration;
168};
169
170} // namespace proteus
171
172#endif
void * RegistrationFunc
Definition CompilerInterfaceDevice.cpp:90
uint64_t uint32_t CallsiteIndex
Definition CompilerInterfaceDevice.cpp:73
uint64_t LambdaID
Definition CompilerInterfaceDevice.cpp:72
void * Kernel
Definition CompilerInterfaceDevice.cpp:62
char int void ** Args
Definition CompilerInterfaceHost.cpp:23
Definition LambdaRegistry.h:25
std::optional< JitVariantMap > getHostJitVariant(uint64_t FunctorID) const
Definition LambdaRegistry.h:44
void beginDeviceLaunch()
Definition LambdaRegistry.h:78
void populateLambdaRegistrationCodeCache(void *Kernel, void *RegistrationFunc)
Definition LambdaRegistry.h:121
void appendDeviceCallsiteRuntimeConstant(uint64_t LambdaID, uint32_t CallsiteIndex, const RuntimeConstant &RC)
Definition LambdaRegistry.h:87
LambdaRegistry & operator=(LambdaRegistry &&)=delete
SmallVector< JitVariantMap, 4 > JitVariantVec
Definition LambdaRegistry.h:38
LambdaRegistry(const LambdaRegistry &)=delete
void appendHostJitVariable(uint64_t FunctorID, const RuntimeConstant &RC)
Definition LambdaRegistry.h:56
void commitHostJitVariables(uint64_t FunctorID)
Definition LambdaRegistry.h:60
DenseMap< int32_t, RuntimeConstant > JitVariantMap
Definition LambdaRegistry.h:37
LambdaRegistry & operator=(const LambdaRegistry &)=delete
void eraseHostJitVariables(uint64_t FunctorID)
Definition LambdaRegistry.h:72
void invokeRegisterLambdaConstants(void *Kernel, void **Args)
Definition LambdaRegistry.h:127
static LambdaRegistry & instance()
Definition LambdaRegistry.h:27
LambdaRegistry(LambdaRegistry &&)=delete
DeviceLaunchInfo takeDeviceLaunchInfo()
Definition LambdaRegistry.h:113
void finalizeDeviceLaunch()
Definition LambdaRegistry.h:99
bool emptyHost() const
Definition LambdaRegistry.h:76
Definition CompiledLibrary.h:8
Definition MemoryCache.h:27
DenseMap< uint64_t, LambdaCallsiteRuntimeConstants > LambdaCallsiteRuntimeConstantsMap
Definition LambdaCallsite.h:32
Definition LambdaRegistry.h:39
LambdaCallsiteRuntimeConstantsMap CallsiteRuntimeConstants
Definition LambdaRegistry.h:41
SmallVector< uint64_t, 4 > LambdaCalleeInfo
Definition LambdaRegistry.h:40
Definition CompilerInterfaceTypes.h:72
int32_t Pos
Definition CompilerInterfaceTypes.h:75