1#ifndef PROTEUS_FRONTEND_DISPATCHER_DEVICE_H
2#define PROTEUS_FRONTEND_DISPATCHER_DEVICE_H
4#if PROTEUS_ENABLE_HIP || PROTEUS_ENABLE_CUDA
18#include <llvm/Support/MemoryBuffer.h>
24template <
typename JitT>
class DispatcherDevice :
public Dispatcher {
26 using KernelFunction_t =
typename DeviceTraits<JitT>::KernelFunction_t;
28 std::unique_ptr<MemoryBuffer>
29 compileModule(Module &M,
const CodeGenerationConfig &CGConfig)
override {
30 TIMESCOPE(DispatcherDevice, compileModule);
32 linkDeviceLibraries(M);
33 optimizeModule(M, CGConfig);
34 return codegenModule(M, CGConfig);
37 void optimizeModule(Module &M,
38 const CodeGenerationConfig &CGConfig)
override {
39 TIMESCOPE(DispatcherDevice, optimizeModule);
41 if (JitT::optimizesBeforeCodegen(CGConfig.codeGenOption())) {
43 OptimizationPipelineConfig(CGConfig));
47 if (CGConfig.codeGenOption() == CodegenOption::RTC)
48 warnOptimizationConfigIgnoredByRTC(CGConfig);
51 std::unique_ptr<MemoryBuffer>
52 codegenModule(Module &M,
const CodeGenerationConfig &CGConfig)
override {
53 TIMESCOPE(DispatcherDevice, codegenModule);
55 auto ObjBuf = Jit.codegenObject(M, Jit.GlobalLinkedBinaries, CGConfig);
62 DispatchResult launch(
void *KernelFunc,
LaunchDims GridDim,
64 uint64_t ShmemSize,
void *Stream)
override {
66 dim3 DevGridDim = {GridDim.
X, GridDim.
Y, GridDim.
Z};
67 dim3 DevBlockDim = {BlockDim.
X, BlockDim.
Y, BlockDim.
Z};
69 reinterpret_cast<typename DeviceTraits<JitT>::DeviceStream_t
>(Stream);
72 reinterpret_cast<KernelFunction_t
>(KernelFunc), DevGridDim, DevBlockDim,
73 KernelArgs, ShmemSize, DevStream);
76 StringRef getDeviceArch()
const override {
return Jit.getDeviceArch(); }
78 void *lookupFunction(StringRef BaseName,
const HashT &ModuleHash)
override {
79 HashT HashValue =
hash(BaseName, ModuleHash);
80 return CodeCache.lookup(HashValue);
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);
88 static const std::unordered_map<std::string, GlobalVarInfo> NoGlobals;
89 const auto &VarNameToGlobalInfo =
90 Library.VarNameToGlobalInfo ? *Library.VarNameToGlobalInfo : NoGlobals;
94 if (Library.VarNameToGlobalInfo && !Library.RelinkGlobalsByCopy &&
95 !Library.GlobalsRelinked) {
96 proteus::relinkGlobalsObject(Library.ObjectModule->getMemBufferRef(),
98 Library.GlobalsRelinked =
true;
102 Name.mangled(), Library.ObjectModule->getBufferStart(),
103 Library.RelinkGlobalsByCopy, VarNameToGlobalInfo);
104 Library.IsLoaded =
true;
106 CodeCache.insert(HashValue, KernelFunc, Name.base());
111 void registerDynamicLibrary(
const HashT &,
const std::string &)
override {
115 ~DispatcherDevice() {
116 if (Config::get().traceCacheStats())
117 CodeCache.printStats();
118 CodeCache.printKernelTrace();
119 printObjectCacheStats();
123 DispatcherDevice(
const std::string &Label, TargetModelType TM, JitT &Jit)
124 : Dispatcher(Label, TM), Jit(Jit), CodeCache(Label) {}
126 virtual void linkDeviceLibraries(Module &M) = 0;
135 warnOptimizationConfigIgnoredByRTC(
const CodeGenerationConfig &CGConfig) {
138 const bool UsesDefaults =
139 !CGConfig.optPipeline() && CGConfig.optLevel() ==
'3' &&
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 "
155 MemoryCache<KernelFunction_t> CodeCache;
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