Proteus
Programmable JIT compilation and optimization for C/C++ using LLVM
Loading...
Searching...
No Matches
DispatcherDevice.h
Go to the documentation of this file.
1#ifndef PROTEUS_FRONTEND_DISPATCHER_DEVICE_H
2#define PROTEUS_FRONTEND_DISPATCHER_DEVICE_H
3
4#if PROTEUS_ENABLE_HIP || PROTEUS_ENABLE_CUDA
5
6#include "proteus/Error.h"
12#include "proteus/impl/Config.h"
17
18#include <llvm/Support/MemoryBuffer.h>
19
20#include <mutex>
21
22namespace proteus {
23
24template <typename JitT> class DispatcherDevice : public Dispatcher {
25public:
26 using KernelFunction_t = typename DeviceTraits<JitT>::KernelFunction_t;
27
28 std::unique_ptr<MemoryBuffer>
29 compileModule(Module &M, const CodeGenerationConfig &CGConfig) override {
30 TIMESCOPE(DispatcherDevice, compileModule);
31
32 linkDeviceLibraries(M);
33 optimizeModule(M, CGConfig);
34 return codegenModule(M, CGConfig);
35 }
36
37 void optimizeModule(Module &M,
38 const CodeGenerationConfig &CGConfig) override {
39 TIMESCOPE(DispatcherDevice, optimizeModule);
40
41 if (JitT::optimizesBeforeCodegen(CGConfig.codeGenOption())) {
42 proteus::optimizeIR(M, Jit.getDeviceArch(),
43 OptimizationPipelineConfig(CGConfig));
44 return;
45 }
46
47 if (CGConfig.codeGenOption() == CodegenOption::RTC)
48 warnOptimizationConfigIgnoredByRTC(CGConfig);
49 }
50
51 std::unique_ptr<MemoryBuffer>
52 codegenModule(Module &M, const CodeGenerationConfig &CGConfig) override {
53 TIMESCOPE(DispatcherDevice, codegenModule);
54
55 auto ObjBuf = Jit.codegenObject(M, Jit.GlobalLinkedBinaries, CGConfig);
56 if (!ObjBuf)
57 reportFatalError("Expected non-null object library");
58
59 return ObjBuf;
60 }
61
62 DispatchResult launch(void *KernelFunc, LaunchDims GridDim,
63 LaunchDims BlockDim, void *KernelArgs[],
64 uint64_t ShmemSize, void *Stream) override {
65 TIMESCOPE(DispatcherDevice, launch);
66 dim3 DevGridDim = {GridDim.X, GridDim.Y, GridDim.Z};
67 dim3 DevBlockDim = {BlockDim.X, BlockDim.Y, BlockDim.Z};
68 auto DevStream =
69 reinterpret_cast<typename DeviceTraits<JitT>::DeviceStream_t>(Stream);
70
72 reinterpret_cast<KernelFunction_t>(KernelFunc), DevGridDim, DevBlockDim,
73 KernelArgs, ShmemSize, DevStream);
74 }
75
76 StringRef getDeviceArch() const override { return Jit.getDeviceArch(); }
77
78 void *lookupFunction(StringRef BaseName, const HashT &ModuleHash) override {
79 HashT HashValue = hash(BaseName, ModuleHash);
80 return CodeCache.lookup(HashValue);
81 }
82
83 void *insertFunction(const KernelName &Name, const HashT &ModuleHash,
84 CompiledLibrary &Library) override {
85 TIMESCOPE(DispatcherDevice, insertFunction);
86 HashT HashValue = hash(StringRef(Name.base()), ModuleHash);
87
88 static const std::unordered_map<std::string, GlobalVarInfo> NoGlobals;
89 const auto &VarNameToGlobalInfo =
90 Library.VarNameToGlobalInfo ? *Library.VarNameToGlobalInfo : NoGlobals;
91
92 // Objects coming from the object cache have not been relinked against
93 // the current process' globals.
94 if (Library.VarNameToGlobalInfo && !Library.RelinkGlobalsByCopy &&
95 !Library.GlobalsRelinked) {
96 proteus::relinkGlobalsObject(Library.ObjectModule->getMemBufferRef(),
97 VarNameToGlobalInfo);
98 Library.GlobalsRelinked = true;
99 }
100
101 auto KernelFunc = proteus::getKernelFunctionFromImage(
102 Name.mangled(), Library.ObjectModule->getBufferStart(),
103 Library.RelinkGlobalsByCopy, VarNameToGlobalInfo);
104 Library.IsLoaded = true;
105
106 CodeCache.insert(HashValue, KernelFunc, Name.base());
107
108 return KernelFunc;
109 }
110
111 void registerDynamicLibrary(const HashT &, const std::string &) override {
112 reportFatalError(Label + " does not support registerDynamicLibrary");
113 }
114
115 ~DispatcherDevice() {
116 if (Config::get().traceCacheStats())
117 CodeCache.printStats();
118 CodeCache.printKernelTrace();
119 printObjectCacheStats();
120 }
121
122protected:
123 DispatcherDevice(const std::string &Label, TargetModelType TM, JitT &Jit)
124 : Dispatcher(Label, TM), Jit(Jit), CodeCache(Label) {}
125
126 virtual void linkDeviceLibraries(Module &M) = 0;
127
128 JitT &Jit;
129
130private:
131 // Skipping optimization on the RTC path leaves the runtime compiler in charge
132 // of it, which silently drops every user-configured optimization setting.
133 // Warn once per process so the setting does not appear to take effect.
134 static void
135 warnOptimizationConfigIgnoredByRTC(const CodeGenerationConfig &CGConfig) {
136 // '3' and 3 are the PROTEUS_OPT_LEVEL and PROTEUS_CODEGEN_OPT_LEVEL
137 // defaults set by CodeGenerationConfig.
138 const bool UsesDefaults =
139 !CGConfig.optPipeline() && CGConfig.optLevel() == '3' &&
140 CGConfig.codeGenOptLevel() == 3 && getJITPassPluginConfigs().empty();
141 if (UsesDefaults)
142 return;
143
144 static std::once_flag WarnOnce;
145 std::call_once(WarnOnce, [] {
146 Logger::outs("proteus")
147 << "Warning: RTC codegen optimizes internally, so Proteus ignores "
148 "PROTEUS_OPT_PIPELINE, PROTEUS_OPT_LEVEL, "
149 "PROTEUS_CODEGEN_OPT_LEVEL and JIT pass plugins, use "
150 "PROTEUS_CODEGEN=serial or PROTEUS_CODEGEN=parallel to apply "
151 "them\n";
152 });
153 }
154
155 MemoryCache<KernelFunction_t> CodeCache;
156};
157
158} // namespace proteus
159
160#endif
161
162#endif // PROTEUS_FRONTEND_DISPATCHER_DEVICE_H
void char * KernelName
Definition CompilerInterfaceDevice.cpp:62
#define TIMESCOPE(...)
Definition TimeTracing.h:66
Definition MemoryCache.h:27
void optimizeIR(Module &M, StringRef Arch, const OptimizationPipelineConfig &OptConfig)
Definition CoreLLVM.h:309
std::vector< JITPassPluginConfig > getJITPassPluginConfigs()
Definition JITPassPluginRegistry.cpp:105
HashT hash(FirstT &&First, RestTs &&...Rest)
Definition Hashing.h:268
void reportFatalError(const llvm::Twine &Reason, const char *FILE, unsigned Line)
Definition Error.cpp:14
cudaError_t launchKernelFunction(CUfunction KernelFunc, dim3 GridDim, dim3 BlockDim, void **KernelArgs, uint64_t ShmemSize, CUstream Stream)
Definition CoreDeviceCUDA.h:81
CUfunction getKernelFunctionFromImage(StringRef KernelName, const void *Image, bool RelinkGlobalsByCopy, const std::unordered_map< std::string, GlobalVarInfo > &VarNameToGlobalInfo)
Definition CoreDeviceCUDA.h:53
Definition Dispatcher.h:25
unsigned Z
Definition Dispatcher.h:26
unsigned Y
Definition Dispatcher.h:26
unsigned X
Definition Dispatcher.h:26