Proteus
Programmable JIT compilation and optimization for C/C++ using LLVM
Loading...
Searching...
No Matches
JitEngineDevice.h
Go to the documentation of this file.
1//===-- JitEngineDevice.cpp -- Base JIT Engine Device header impl. --===//
2//
3// Part of the Proteus Project, under the Apache License v2.0 with LLVM
4// Exceptions. See https://llvm.org/LICENSE.txt for license information.
5// SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6//
7//===----------------------------------------------------------------------===//
8//
9//===----------------------------------------------------------------------===//
10
11#ifndef PROTEUS_JITENGINEDEVICE_H
12#define PROTEUS_JITENGINEDEVICE_H
13
16#include "proteus/Init.h"
17#include "proteus/TimeTracing.h"
24#include "proteus/impl/Debug.h"
31#include "proteus/impl/Utils.h"
32
33#include <llvm/ADT/SmallPtrSet.h>
34#include <llvm/ADT/SmallVector.h>
35#include <llvm/ADT/StringRef.h>
36#include <llvm/Analysis/CallGraph.h>
37#include <llvm/Analysis/TargetTransformInfo.h>
38#include <llvm/Bitcode/BitcodeWriter.h>
39#include <llvm/CodeGen/CommandFlags.h>
40#include <llvm/CodeGen/MachineModuleInfo.h>
41#include <llvm/Config/llvm-config.h>
42#include <llvm/Demangle/Demangle.h>
43#include <llvm/ExecutionEngine/Orc/ThreadSafeModule.h>
44#include <llvm/IR/Constants.h>
45#include <llvm/IR/GlobalVariable.h>
46#include <llvm/IR/Instruction.h>
47#include <llvm/IR/Instructions.h>
48#include <llvm/IR/LLVMContext.h>
49#include <llvm/IR/LegacyPassManager.h>
50#include <llvm/IR/Module.h>
51#include <llvm/IR/ReplaceConstant.h>
52#include <llvm/IR/Type.h>
53#include <llvm/IR/Verifier.h>
54#include <llvm/IRReader/IRReader.h>
55#include <llvm/Linker/Linker.h>
56#include <llvm/MC/TargetRegistry.h>
57#include <llvm/Object/ELFObjectFile.h>
58#include <llvm/Passes/PassBuilder.h>
59#include <llvm/Support/Error.h>
60#include <llvm/Support/MemoryBuffer.h>
61#include <llvm/Support/MemoryBufferRef.h>
62#include <llvm/Target/TargetMachine.h>
63#include <llvm/Transforms/IPO/Internalize.h>
64#include <llvm/Transforms/Utils/Cloning.h>
65#include <llvm/Transforms/Utils/ModuleUtils.h>
66
67#include <cstdint>
68#include <functional>
69#include <memory>
70#include <optional>
71#include <string>
72
73namespace proteus {
74
75using namespace llvm;
76
78 int32_t Magic;
79 int32_t Version;
80 const char *Binary;
82};
83
85private:
86 FatbinWrapperT *FatbinWrapper;
87 std::unique_ptr<LLVMContext> Ctx;
88 SmallVector<std::string> LinkedModuleIds;
89 Module *LinkedModule;
90 std::optional<SmallVector<std::unique_ptr<Module>>> ExtractedModules;
91 std::optional<HashT> ExtractedModuleHash;
92 std::optional<CallGraph> ModuleCallGraph;
93 std::unique_ptr<MemoryBuffer> DeviceBinary;
94 std::unordered_map<std::string, GlobalVarInfo> VarNameToGlobalInfo;
95 std::once_flag Flag;
96
97public:
98 BinaryInfo() = default;
99 BinaryInfo(FatbinWrapperT *FatbinWrapper,
100 SmallVector<std::string> &&LinkedModuleIds)
101 : FatbinWrapper(FatbinWrapper), Ctx(std::make_unique<LLVMContext>()),
102 LinkedModuleIds(LinkedModuleIds), LinkedModule(nullptr),
103 ExtractedModules(std::nullopt), ModuleCallGraph(std::nullopt),
104 DeviceBinary(nullptr) {}
105
107
108 std::unique_ptr<LLVMContext> &getLLVMContext() { return Ctx; }
109
110 bool hasLinkedModule() const { return (LinkedModule != nullptr); }
111 Module &getLinkedModule() {
113 if (!LinkedModule) {
114 if (!hasExtractedModules())
115 reportFatalError("Expected extracted modules");
116
117 Timer T(Config::get().ProteusEnableTimers);
118 // Avoid linking when there's a single module by moving it instead and
119 // making sure it's materialized for call graph analysis.
120 if (ExtractedModules->size() == 1) {
121 LinkedModule = ExtractedModules->front().get();
122 if (auto E = LinkedModule->materializeAll())
123 reportFatalError("Error materializing " + toString(std::move(E)));
124 } else {
125 // By the LLVM API, linkModules takes ownership of module pointers in
126 // ExtractedModules and returns a new unique ptr to the linked module.
127 // We update ExtractedModules to contain and own only the generated
128 // LinkedModule.
129 auto GeneratedLinkedModule =
130 proteus::linkModules(*Ctx, std::move(ExtractedModules.value()));
131 SmallVector<std::unique_ptr<Module>> NewExtractedModules;
132 NewExtractedModules.emplace_back(std::move(GeneratedLinkedModule));
133 setExtractedModules(NewExtractedModules);
134
135 LinkedModule = ExtractedModules->front().get();
136 }
137
139 << "getLinkedModule " << T.elapsed() << " ms\n");
140 }
141
142 return *LinkedModule;
143 }
144
145 bool hasExtractedModules() const { return ExtractedModules.has_value(); }
146 const SmallVector<std::reference_wrapper<Module>>
148 // This should be called only once when cloning the kernel module to
149 // cache.
150 SmallVector<std::reference_wrapper<Module>> ModulesRef;
151 for (auto &M : ExtractedModules.value())
152 ModulesRef.emplace_back(*M);
153
154 return ModulesRef;
155 }
156 void setExtractedModules(SmallVector<std::unique_ptr<Module>> &Modules) {
157 ExtractedModules = std::move(Modules);
158 }
159
160 bool hasModuleHash() const { return ExtractedModuleHash.has_value(); }
162 if (!hasModuleHash())
163 reportFatalError("Expected module hash to be set");
164
165 return ExtractedModuleHash.value();
166 }
167 void setModuleHash(HashT HashValue) { ExtractedModuleHash = HashValue; }
168 void updateModuleHash(HashT HashValue) {
169 if (ExtractedModuleHash)
170 ExtractedModuleHash = hashCombine(ExtractedModuleHash.value(), HashValue);
171 else
172 ExtractedModuleHash = HashValue;
173 }
174
175 CallGraph &getCallGraph() {
176 if (!ModuleCallGraph.has_value()) {
177 if (!LinkedModule)
178 reportFatalError("Expected non-null linked module");
179 ModuleCallGraph.emplace(CallGraph(*LinkedModule));
180 }
181 return ModuleCallGraph.value();
182 }
183
184 bool hasDeviceBinary() { return (DeviceBinary != nullptr); }
185 MemoryBufferRef getDeviceBinary() {
186 if (!hasDeviceBinary())
187 reportFatalError("Expected non-null device binary");
188 return DeviceBinary->getMemBufferRef();
189 }
190 void setDeviceBinary(std::unique_ptr<MemoryBuffer> DeviceBinaryBuffer) {
191 DeviceBinary = std::move(DeviceBinaryBuffer);
192 }
193
194 void addModuleId(const char *ModuleId) {
195 LinkedModuleIds.push_back(ModuleId);
196 }
197
198 void insertGlobalVar(const char *VarName, const void *HostAddr,
199 const void *DeviceAddr, uint64_t VarSize) {
200 auto KV = VarNameToGlobalInfo.emplace(
201 VarName, GlobalVarInfo(HostAddr, DeviceAddr, VarSize));
202
203 auto TraceOut = [&KV]() {
204 auto GlobalName = KV.first->first;
205 auto &GVI = KV.first->second;
206
207 SmallString<128> S;
208 raw_svector_ostream OS(S);
209 OS << "[GVarInfo]: " << GlobalName << " HAddr:" << GVI.HostAddr
210 << " DevAddr:" << GVI.DevAddr << " VarSize:" << GVI.VarSize << "\n";
211
212 return S;
213 };
214
215 if (Config::get().traceSpecializations())
216 Logger::trace(TraceOut());
217 }
218
219 std::unordered_map<std::string, GlobalVarInfo> &getVarNameToGlobalInfo() {
220 return VarNameToGlobalInfo;
221 }
222
223 auto &getModuleIds() { return LinkedModuleIds; }
224};
225
227 std::optional<void *> Kernel;
228 std::unique_ptr<LLVMContext> Ctx;
229 std::string Name;
230 ArrayRef<RuntimeConstantInfo *> RCInfoArray;
231 std::optional<std::unique_ptr<Module>> ExtractedModule;
232 std::optional<std::unique_ptr<MemoryBuffer>> Bitcode;
233 std::optional<std::reference_wrapper<BinaryInfo>> BinInfo;
234 std::optional<HashT> StaticHash;
235 // LambdaCalleeInfo optionally contains a vector with all annotated
236 // LambdaFunctorWrapper annotators and thei unique IDs
237 std::optional<SmallVector<uint64_t>> LambdaCalleeInfo;
238 std::optional<LambdaCallsiteLocationMap> LambdaCallsiteLocationInfo;
239
240public:
241 JITKernelInfo(void *Kernel, BinaryInfo &BinInfo, char const *Name,
242 ArrayRef<RuntimeConstantInfo *> RCInfoArray)
243 : Kernel(Kernel), Ctx(std::make_unique<LLVMContext>()), Name(Name),
244 RCInfoArray(RCInfoArray), ExtractedModule(std::nullopt),
245 Bitcode{std::nullopt}, BinInfo(BinInfo), LambdaCalleeInfo(std::nullopt),
246 LambdaCallsiteLocationInfo(std::nullopt) {}
247
248 JITKernelInfo() = default;
249 void *getKernel() const {
250 assert(Kernel.has_value() && "Expected Kernel is inited");
251 return Kernel.value();
252 }
253 std::unique_ptr<LLVMContext> &getLLVMContext() { return Ctx; }
254 const std::string &getName() const { return Name; }
255 ArrayRef<RuntimeConstantInfo *> getRCInfoArray() const { return RCInfoArray; }
256 bool hasModule() const { return ExtractedModule.has_value(); }
257 Module &getModule() const { return *ExtractedModule->get(); }
258 BinaryInfo &getBinaryInfo() const { return BinInfo.value(); }
259 void setModule(std::unique_ptr<llvm::Module> Mod) {
260 ExtractedModule = std::move(Mod);
261 }
262
263 bool hasBitcode() { return Bitcode.has_value(); }
264 void setBitcode(std::unique_ptr<MemoryBuffer> ExtractedBitcode) {
265 Bitcode = std::move(ExtractedBitcode);
266 }
267 MemoryBufferRef getBitcode() { return Bitcode.value()->getMemBufferRef(); }
268
269 bool hasStaticHash() const { return StaticHash.has_value(); }
270 const HashT getStaticHash() const { return StaticHash.value(); }
271 void createStaticHash(HashT ModuleHash) {
272 StaticHash = hash(Name);
273 StaticHash = hashCombine(StaticHash.value(), ModuleHash);
274 }
275
276 bool hasLambdaCalleeInfo() { return LambdaCalleeInfo.has_value(); }
277 const auto &getLambdaCalleeInfo() { return LambdaCalleeInfo.value(); }
278 void setLambdaCalleeInfo(SmallVector<uint64_t> &&LambdaInfo) {
279 LambdaCalleeInfo = std::move(LambdaInfo);
280 }
281
283 return LambdaCallsiteLocationInfo.has_value();
284 }
285 const auto &getLambdaCallsiteLocationInfo() const {
286 return LambdaCallsiteLocationInfo.value();
287 }
289 LambdaCallsiteLocationInfo = std::move(Info);
290 }
292 LambdaKernelArgLocation Location) {
293 if (!LambdaCallsiteLocationInfo)
294 LambdaCallsiteLocationInfo.emplace();
295
296 auto &Callsites = (*LambdaCallsiteLocationInfo)[LambdaID];
297 auto It = Callsites.find(CallsiteIndex);
298 if (It != Callsites.end()) {
299 if (It->second.KernelArgIndex != Location.KernelArgIndex ||
300 It->second.Offset != Location.Offset ||
301 It->second.StorageType != Location.StorageType)
302 reportFatalError("Conflicting lambda callsite location for lambda " +
303 std::to_string(LambdaID));
304 return;
305 }
306 Callsites[CallsiteIndex] = Location;
307 }
308};
309
310template <typename ImplT> struct DeviceTraits;
311
312template <typename ImplT> class JitEngineDevice : public JitEngine {
313public:
317
319 compileAndRun(JITKernelInfo &KernelInfo, dim3 GridDim, dim3 BlockDim,
320 void **KernelArgs, uint64_t ShmemSize,
322
323 std::pair<std::unique_ptr<Module>, std::unique_ptr<MemoryBuffer>>
325 LLVMContext &Ctx) {
327 std::unique_ptr<Module> KernelModule =
328 static_cast<ImplT &>(*this).tryExtractKernelModule(BinInfo, KernelName,
329 Ctx);
330 std::unique_ptr<MemoryBuffer> Bitcode = nullptr;
331
332 // If there is no ready-made kernel module from AOT, extract per-TU or the
333 // single linked module and clone the kernel module.
334 if (!KernelModule) {
335 Timer T(Config::get().ProteusEnableTimers);
336 if (!BinInfo.hasExtractedModules())
337 static_cast<ImplT &>(*this).extractModules(BinInfo);
338
339 std::unique_ptr<Module> KernelModuleTmp = nullptr;
340 switch (Config::get().ProteusKernelClone) {
342 auto &LinkedModule = BinInfo.getLinkedModule();
343 KernelModule = llvm::CloneModule(LinkedModule);
344 break;
345 }
347 auto &LinkedModule = BinInfo.getLinkedModule();
348 KernelModule =
350 break;
351 }
353 KernelModule = proteus::cloneKernelFromModules(
355 break;
356 }
357 default:
358 reportFatalError("Unsupported kernel cloning option");
359 }
360
362 << "Cloning "
363 << toString(Config::get().ProteusKernelClone) << " "
364 << T.elapsed() << " ms\n");
365 }
366
367 // Internalize and cleanup to simplify the module and prepare it for
368 // optimization.
369 internalize(*KernelModule, KernelName);
370 proteus::runCleanupPassPipeline(*KernelModule);
371
372 // If the module is not in the provided context due to cloning, roundtrip
373 // it using bitcode. Re-use the roundtrip bitcode to return it.
374 if (&KernelModule->getContext() != &Ctx) {
375 SmallVector<char> CloneBuffer;
376 raw_svector_ostream OS(CloneBuffer);
377 WriteBitcodeToFile(*KernelModule, OS);
378 StringRef CloneStr = StringRef(CloneBuffer.data(), CloneBuffer.size());
379 auto ExpectedKernelModule =
380 parseBitcodeFile(MemoryBufferRef{CloneStr, KernelName}, Ctx);
381 if (auto E = ExpectedKernelModule.takeError())
382 reportFatalError("Error parsing bitcode: " + toString(std::move(E)));
383
384 KernelModule = std::move(*ExpectedKernelModule);
385 Bitcode = MemoryBuffer::getMemBufferCopy(CloneStr);
386 } else {
387 // Parse the kernel module to create the bitcode since it has not been
388 // created by roundtripping.
389 SmallVector<char> BitcodeBuffer;
390 raw_svector_ostream OS(BitcodeBuffer);
391 WriteBitcodeToFile(*KernelModule, OS);
392 auto BitcodeStr = StringRef{BitcodeBuffer.data(), BitcodeBuffer.size()};
393 Bitcode = MemoryBuffer::getMemBufferCopy(BitcodeStr);
394 }
395
396 return std::make_pair(std::move(KernelModule), std::move(Bitcode));
397 }
398
401
402 if (KernelInfo.hasModule() && KernelInfo.hasBitcode())
403 return;
404
405 if (KernelInfo.hasModule())
406 reportFatalError("Unexpected KernelInfo has module but not bitcode");
407
408 if (KernelInfo.hasBitcode())
409 reportFatalError("Unexpected KernelInfo has bitcode but not module");
410
411 BinaryInfo &BinInfo = KernelInfo.getBinaryInfo();
412
413 Timer T(Config::get().ProteusEnableTimers);
414 auto [KernelModule, BitcodeBuffer] = extractKernelModule(
415 BinInfo, KernelInfo.getName(), *KernelInfo.getLLVMContext());
416
417 if (!KernelModule)
418 reportFatalError("Expected non-null kernel module");
419 if (!BitcodeBuffer)
420 reportFatalError("Expected non-null kernel bitcode");
421
422 KernelInfo.setModule(std::move(KernelModule));
423 KernelInfo.setBitcode(std::move(BitcodeBuffer));
425 << "Extract kernel module " << T.elapsed() << " ms\n");
426 }
427
428 Module &getModule(JITKernelInfo &KernelInfo) {
429 if (!KernelInfo.hasModule())
430 extractModuleAndBitcode(KernelInfo);
431
432 if (!KernelInfo.hasModule())
433 reportFatalError("Expected module in KernelInfo");
434
435 return KernelInfo.getModule();
436 }
437
438 MemoryBufferRef getBitcode(JITKernelInfo &KernelInfo) {
439 if (!KernelInfo.hasBitcode())
440 extractModuleAndBitcode(KernelInfo);
441
442 if (!KernelInfo.hasBitcode())
443 reportFatalError("Expected bitcode in KernelInfo");
444
445 return KernelInfo.getBitcode();
446 }
447
448 void
449 getLambdaJitValues(SmallVector<uint64_t> &LambdaCalleeInfo,
450 LambdaCallsiteRuntimeConstantsMap &LambdaJitValuesMap) {
453 auto LaunchInfo = LR.takeDeviceLaunchInfo();
454 if (LaunchInfo.CallsiteRuntimeConstants.empty()) {
455 return;
456 }
457 LambdaCalleeInfo = std::move(LaunchInfo.LambdaCalleeInfo);
458 LambdaJitValuesMap = std::move(LaunchInfo.CallsiteRuntimeConstants);
459 }
460
461 void registerVar(void *Handle, const char *VarName, const void *HostAddr,
462 uint64_t VarSize) {
463 if (!HandleToBinaryInfo.count(Handle))
464 reportFatalError("Expected Handle in map");
465 BinaryInfo &BinInfo = HandleToBinaryInfo[Handle];
466
467 void *DeviceAddr = resolveDeviceGlobalAddr(HostAddr);
468 assert(DeviceAddr &&
469 "Expected non-null device address for global variable");
470
471 BinInfo.insertGlobalVar(VarName, HostAddr, DeviceAddr, VarSize);
472 }
473
475 const char *ModuleId);
477 const char *ModuleId);
479 void registerFunction(void *Handle, void *Kernel, char *KernelName,
480 ArrayRef<RuntimeConstantInfo *> RCInfoArray);
482 uint32_t CallsiteIndex,
483 uint32_t KernelArgIndex, int64_t Offset,
485
486 std::unordered_map<std::string, FatbinWrapperT *> ModuleIdToFatBinary;
487 std::unordered_map<const void *, BinaryInfo> HandleToBinaryInfo;
488 SmallPtrSet<void *, 8> GlobalLinkedBinaries;
489
490 bool containsJITKernelInfo(const void *Func) {
491 return JITKernelInfoMap.contains(Func);
492 }
493
494 std::optional<std::reference_wrapper<JITKernelInfo>>
495 getJITKernelInfo(const void *Func) {
497 return std::nullopt;
498 }
499 return JITKernelInfoMap[Func];
500 }
501
503 if (KernelInfo.hasStaticHash())
504 return KernelInfo.getStaticHash();
505
506 BinaryInfo &BinInfo = KernelInfo.getBinaryInfo();
507
508 if (BinInfo.hasModuleHash()) {
509 KernelInfo.createStaticHash(BinInfo.getModuleHash());
510 return KernelInfo.getStaticHash();
511 }
512
513 HashT ModuleHash = static_cast<ImplT &>(*this).getModuleHash(BinInfo);
514
515 KernelInfo.createStaticHash(BinInfo.getModuleHash());
516 return KernelInfo.getStaticHash();
517 }
518
519public:
520 StringRef getDeviceArch() const { return DeviceArch; }
521
522protected:
526
527 for (auto &[Handle, FatbinInfo] : JitEngineInfo.FatbinaryMap) {
529 Handle, reinterpret_cast<FatbinWrapperT *>(FatbinInfo.FatbinWrapper),
530 FatbinInfo.ModuleId);
531
532 for (auto &LinkedBin : FatbinInfo.LinkedBinaries)
534 Handle, reinterpret_cast<FatbinWrapperT *>(LinkedBin.FatbinWrapper),
535 LinkedBin.ModuleId);
536
537 for (auto &Func : FatbinInfo.Functions)
538 registerFunction(Handle, Func.Kernel, Func.KernelName,
539 Func.RCInfoArray);
540
541 for (auto &Var : FatbinInfo.Vars)
542 registerVar(Var.Handle, Var.VarName, Var.HostAddr, Var.VarSize);
543 }
544
546
547 if (Config::get().ProteusAsyncCompilation)
549 std::make_unique<CompilerAsync>(Config::get().ProteusAsyncThreads);
550 }
551
553 // Thread joining is handled by CompilerAsync's shutdown guard to ensure it
554 // happens before static objects are destroyed. If this destructor does run,
555 // joinAllThreads() is idempotent.
556 if (AsyncCompiler)
557 AsyncCompiler->joinAllThreads();
558 }
559
561 if (!Dispatch)
562 reportFatalError("Dispatcher has not been created by the engine");
563 return *Dispatch;
564 }
565 std::unique_ptr<Dispatcher> Dispatch;
566 std::string DeviceArch;
567
568 DenseMap<const void *, JITKernelInfo> JITKernelInfoMap;
569 DenseMap<const void *, LambdaCallsiteLocationMap>
571 std::unique_ptr<CompilerAsync> AsyncCompiler;
572};
573
574template <typename ImplT>
577 JITKernelInfo &KernelInfo, dim3 GridDim, dim3 BlockDim, void **KernelArgs,
578 uint64_t ShmemSize, typename DeviceTraits<ImplT>::DeviceStream_t Stream) {
579 TIMESCOPE(JitEngineDevice, compileAndRun);
580
581 auto &BinInfo = KernelInfo.getBinaryInfo();
582
583 SmallVector<RuntimeConstant> RCVec =
584 getRuntimeConstantValues(KernelArgs, KernelInfo.getRCInfoArray());
585
586 SmallVector<uint64_t> LambdaCalleeInfoToSpecialize;
587 LambdaCallsiteRuntimeConstantsMap LambdaJitValuesMap;
588 getLambdaJitValues(LambdaCalleeInfoToSpecialize, LambdaJitValuesMap);
589 const auto &CGConfig = Config::get().getCGConfig(KernelInfo.getName());
590 // Determine the hash based on dimension specialization. If we do not
591 // specialize IR based on grid dimensions, avoid hashing on those to
592 // eliminate repeated compilation overhead.
593 HashT HashValue =
594 hash(getStaticHash(KernelInfo), RCVec, LambdaJitValuesMap, BlockDim.x,
595 BlockDim.y, BlockDim.z, hashCodeGenConfig(CGConfig),
597 if (CGConfig.specializeDims() || CGConfig.specializeDimsRange())
598 HashValue = hash(HashValue, GridDim.x, GridDim.y, GridDim.z);
599
600 Dispatcher &Dispatch = getDispatcher();
601
602 if (void *KernelFunc =
603 Dispatch.lookupFunction(KernelInfo.getName(), HashValue))
604 return static_cast<DeviceError_t>(Dispatch
605 .launch(KernelFunc, GridDim, BlockDim,
606 KernelArgs, ShmemSize, Stream)
607 .Ret);
608
609 // NOTE: we don't need a suffix to differentiate kernels, each
610 // specialization will be in its own module uniquely identify by HashValue.
611 // It exists only for debugging purposes to verify that the jitted kernel
612 // executes.
613 KernelName Name{KernelInfo.getName(), HashValue};
614
615 auto LoadKernel = [&](CompiledLibrary &Library) {
616 Library.VarNameToGlobalInfo = &BinInfo.getVarNameToGlobalInfo();
617 Library.RelinkGlobalsByCopy = Config::get().ProteusRelinkGlobalsByCopy;
618 return Dispatch.insertFunction(Name, HashValue, Library);
619 };
620
621 if (auto CompiledLib = Dispatch.lookupCompiledLibrary(HashValue))
622 return static_cast<DeviceError_t>(Dispatch
623 .launch(LoadKernel(*CompiledLib),
624 GridDim, BlockDim, KernelArgs,
625 ShmemSize, Stream)
626 .Ret);
627
628 MemoryBufferRef KernelBitcode = getBitcode(KernelInfo);
629 std::unique_ptr<MemoryBuffer> ObjBuf = nullptr;
630 const LambdaCallsiteRuntimeConstantsMap EmptyLambdaCallsiteRuntimeConstants;
632 LambdaJitValuesMap.empty() ? EmptyLambdaCallsiteRuntimeConstants
633 : LambdaJitValuesMap;
634
635 auto CreateTask = [&]() {
636 return CompilationTask{
637 Dispatch,
638 KernelBitcode,
639 HashValue,
640 Name,
641 BlockDim,
642 GridDim,
643 RCVec,
644 LambdaCalleeInfoToSpecialize,
646 BinInfo.getVarNameToGlobalInfo(),
647 /*CodeGenConfig */ CGConfig,
648 /*DumpIR*/ Config::get().ProteusDumpLLVMIR,
649 /*RelinkGlobalsByCopy*/ Config::get().ProteusRelinkGlobalsByCopy};
650 };
651
652 if (Config::get().ProteusAsyncCompilation) {
653 // If there is no compilation pending for the specialization, post the
654 // compilation task to the compiler.
655 if (!AsyncCompiler->isCompilationPending(HashValue)) {
656 PROTEUS_DBG(Logger::logs("proteus") << "Compile async for HashValue "
657 << HashValue.toString() << "\n");
658
659 AsyncCompiler->compile(CreateTask());
660 }
661
662 // Compilation is pending, try to get the compilation result buffer. If
663 // buffer is null, compilation is not done, so execute the AOT version
664 // directly.
665 ObjBuf = AsyncCompiler->takeCompilationResult(
666 HashValue, Config::get().ProteusAsyncTestBlocking);
667 if (!ObjBuf) {
668 return launchKernelDirect(KernelInfo.getKernel(), GridDim, BlockDim,
669 KernelArgs, ShmemSize, Stream);
670 }
671 } else {
672 // Process through synchronous compilation.
673 ObjBuf = CompilerSync::instance().compile(CreateTask());
674 }
675
676 if (!ObjBuf)
677 reportFatalError("Expected non-null object");
678
679 Dispatch.registerObject(HashValue, ObjBuf->getMemBufferRef());
680
681 CompiledLibrary Library{std::move(ObjBuf)};
682 Library.GlobalsRelinked = true;
683
684 return static_cast<DeviceError_t>(Dispatch
685 .launch(LoadKernel(Library), GridDim,
686 BlockDim, KernelArgs, ShmemSize,
687 Stream)
688 .Ret);
689}
690
691template <typename ImplT>
694 const char *ModuleId) {
696 PROTEUS_DBG(Logger::logs("proteus")
697 << "Register fatbinary Handle " << Handle << " FatbinWrapper "
698 << FatbinWrapper << " Binary " << (void *)FatbinWrapper->Binary
699 << " ModuleId " << ModuleId << "\n");
700 if (FatbinWrapper->PrelinkedFatbins) {
701 // This is RDC compilation, just insert the FatbinWrapper and ignore the
702 // ModuleId coming from the link.stub.
703 HandleToBinaryInfo.try_emplace(Handle, FatbinWrapper,
704 SmallVector<std::string>{});
705
706 // Initialize GlobalLinkedBinaries with prelinked fatbins.
707 void *Ptr = FatbinWrapper->PrelinkedFatbins[0];
708 for (int I = 0; Ptr != nullptr;
709 ++I, Ptr = FatbinWrapper->PrelinkedFatbins[I]) {
710 PROTEUS_DBG(Logger::logs("proteus")
711 << "I " << I << " PrelinkedFatbin " << Ptr << "\n");
712 GlobalLinkedBinaries.insert(Ptr);
713 }
714 } else {
715 // This is non-RDC compilation, associate the ModuleId of the JIT bitcode
716 // in the module with the FatbinWrapper.
717 ModuleIdToFatBinary[ModuleId] = FatbinWrapper;
718 HandleToBinaryInfo.try_emplace(Handle, FatbinWrapper,
719 SmallVector<std::string>{ModuleId});
720 }
721}
722
723template <typename ImplT> void JitEngineDevice<ImplT>::finalizeRegistration() {
724 TIMESCOPE(JitEngineDevice, finalizeRegistration);
725 PROTEUS_DBG(Logger::logs("proteus") << "Finalize registration\n");
726 // Erase linked binaries for which we have LLVM IR code, those binaries are
727 // stored in the ModuleIdToFatBinary map.
728 for (auto &[ModuleId, FatbinWrapper] : ModuleIdToFatBinary)
729 GlobalLinkedBinaries.erase((void *)FatbinWrapper->Binary);
730}
731
732template <typename ImplT>
734 void *Handle, void *Kernel, char *KernelName,
735 ArrayRef<RuntimeConstantInfo *> RCInfoArray) {
736 PROTEUS_DBG(Logger::logs("proteus") << "Register function " << Kernel
737 << " To Handle " << Handle << "\n");
738 // NOTE: HIP RDC might call multiple times the registerFunction for the same
739 // kernel, which has weak linkage, when it comes from different translation
740 // units. Either the first or the second call can prevail and should be
741 // equivalent. We let the first one prevail.
742 if (JITKernelInfoMap.contains(Kernel)) {
743 PROTEUS_DBG(Logger::logs("proteus")
744 << "Warning: duplicate register function for kernel " +
745 std::string(KernelName)
746 << "\n");
747 return;
748 }
749
750 if (!HandleToBinaryInfo.count(Handle))
751 reportFatalError("Expected Handle in map");
752 BinaryInfo &BinInfo = HandleToBinaryInfo[Handle];
753
754 PROTEUS_DBG(Logger::logs("proteus")
755 << "Register function " << KernelName << " with binary handle "
756 << Handle << "\n");
757
758 JITKernelInfoMap[Kernel] =
759 JITKernelInfo{Kernel, BinInfo, KernelName, RCInfoArray};
760 auto PendingIt = PendingLambdaCallsiteLocationInfo.find(Kernel);
761 if (PendingIt != PendingLambdaCallsiteLocationInfo.end()) {
762 JITKernelInfoMap[Kernel].setLambdaCallsiteLocationInfo(
763 std::move(PendingIt->second));
764 PendingLambdaCallsiteLocationInfo.erase(PendingIt);
765 }
766}
767
768template <typename ImplT>
770 void *Kernel, uint64_t LambdaID, uint32_t CallsiteIndex,
773 auto It = JITKernelInfoMap.find(Kernel);
774 if (It == JITKernelInfoMap.end()) {
775 PendingLambdaCallsiteLocationInfo[Kernel][LambdaID][CallsiteIndex] =
776 Location;
777 return;
778 }
779
780 It->second.addLambdaCallsiteLocation(LambdaID, CallsiteIndex, Location);
781}
782
783template <typename ImplT>
786 const char *ModuleId) {
788 PROTEUS_DBG(Logger::logs("proteus")
789 << "Register linked binary FatBinary " << FatbinWrapper
790 << " Binary " << (void *)FatbinWrapper->Binary << " ModuleId "
791 << ModuleId << "\n");
792 if (!HandleToBinaryInfo.count(Handle))
793 reportFatalError("Expected Handle in map");
794
795 HandleToBinaryInfo[Handle].addModuleId(ModuleId);
796 ModuleIdToFatBinary[ModuleId] = FatbinWrapper;
797}
798
799} // namespace proteus
800
801#endif
void const char * ModuleId
Definition CompilerInterfaceDevice.cpp:44
void * FatbinWrapper
Definition CompilerInterfaceDevice.cpp:43
uint64_t uint32_t CallsiteIndex
Definition CompilerInterfaceDevice.cpp:73
const void const char * VarName
Definition CompilerInterfaceDevice.cpp:32
auto & JitEngineInfo
Definition CompilerInterfaceDevice.cpp:67
uint64_t LambdaID
Definition CompilerInterfaceDevice.cpp:72
uint64_t uint32_t uint32_t int64_t Offset
Definition CompilerInterfaceDevice.cpp:75
void * Kernel
Definition CompilerInterfaceDevice.cpp:62
JitEngineInfo registerFatBinary(Handle, FatbinWrapper, ModuleId)
uint64_t uint32_t uint32_t KernelArgIndex
Definition CompilerInterfaceDevice.cpp:74
const void const char uint64_t VarSize
Definition CompilerInterfaceDevice.cpp:33
const void * HostAddr
Definition CompilerInterfaceDevice.cpp:31
auto & LR
Definition CompilerInterfaceDevice.cpp:129
uint64_t uint32_t uint32_t int64_t int32_t StorageType
Definition CompilerInterfaceDevice.cpp:76
JitEngineInfo registerLinkedBinary(FatbinWrapper, ModuleId)
#define PROTEUS_TIMER_OUTPUT(x)
Definition Config.h:483
#define PROTEUS_DBG(x)
Definition Debug.h:9
#define TIMESCOPE(...)
Definition TimeTracing.h:66
Definition JitEngineDevice.h:84
FatbinWrapperT * getFatbinWrapper() const
Definition JitEngineDevice.h:106
void setExtractedModules(SmallVector< std::unique_ptr< Module > > &Modules)
Definition JitEngineDevice.h:156
std::unordered_map< std::string, GlobalVarInfo > & getVarNameToGlobalInfo()
Definition JitEngineDevice.h:219
MemoryBufferRef getDeviceBinary()
Definition JitEngineDevice.h:185
bool hasModuleHash() const
Definition JitEngineDevice.h:160
std::unique_ptr< LLVMContext > & getLLVMContext()
Definition JitEngineDevice.h:108
Module & getLinkedModule()
Definition JitEngineDevice.h:111
auto & getModuleIds()
Definition JitEngineDevice.h:223
bool hasLinkedModule() const
Definition JitEngineDevice.h:110
bool hasDeviceBinary()
Definition JitEngineDevice.h:184
const SmallVector< std::reference_wrapper< Module > > getExtractedModules() const
Definition JitEngineDevice.h:147
void updateModuleHash(HashT HashValue)
Definition JitEngineDevice.h:168
HashT getModuleHash() const
Definition JitEngineDevice.h:161
void insertGlobalVar(const char *VarName, const void *HostAddr, const void *DeviceAddr, uint64_t VarSize)
Definition JitEngineDevice.h:198
bool hasExtractedModules() const
Definition JitEngineDevice.h:145
void addModuleId(const char *ModuleId)
Definition JitEngineDevice.h:194
CallGraph & getCallGraph()
Definition JitEngineDevice.h:175
BinaryInfo(FatbinWrapperT *FatbinWrapper, SmallVector< std::string > &&LinkedModuleIds)
Definition JitEngineDevice.h:99
void setModuleHash(HashT HashValue)
Definition JitEngineDevice.h:167
void setDeviceBinary(std::unique_ptr< MemoryBuffer > DeviceBinaryBuffer)
Definition JitEngineDevice.h:190
Definition CompilationTask.h:21
std::unique_ptr< MemoryBuffer > compile(CompilationTask &&CT)
Definition CompilerSync.h:22
static CompilerSync & instance()
Definition CompilerSync.h:17
static Config & get()
Definition Config.h:371
bool ProteusRelinkGlobalsByCopy
Definition Config.h:380
bool ProteusDumpLLVMIR
Definition Config.h:379
const CodeGenerationConfig & getCGConfig(llvm::StringRef KName="") const
Definition Config.h:397
Definition Dispatcher.h:79
virtual void * lookupFunction(llvm::StringRef BaseName, const HashT &ModuleHash)=0
std::unique_ptr< CompiledLibrary > lookupCompiledLibrary(const HashT &ModuleHash)
Definition Dispatcher.cpp:124
void registerObject(const HashT &HashValue, const llvm::MemoryBufferRef &Obj)
Definition Dispatcher.cpp:141
virtual void * insertFunction(const KernelName &Name, const HashT &ModuleHash, CompiledLibrary &Library)=0
virtual DispatchResult launch(void *KernelFunc, LaunchDims GridDim, LaunchDims BlockDim, void *KernelArgs[], uint64_t ShmemSize, void *Stream)=0
Definition Func.h:296
Definition Hashing.h:27
std::string toString() const
Definition Hashing.h:35
Definition JitEngineDevice.h:226
bool hasBitcode()
Definition JitEngineDevice.h:263
bool hasLambdaCallsiteLocationInfo() const
Definition JitEngineDevice.h:282
const std::string & getName() const
Definition JitEngineDevice.h:254
void createStaticHash(HashT ModuleHash)
Definition JitEngineDevice.h:271
void setLambdaCallsiteLocationInfo(LambdaCallsiteLocationMap &&Info)
Definition JitEngineDevice.h:288
void addLambdaCallsiteLocation(uint64_t LambdaID, uint32_t CallsiteIndex, LambdaKernelArgLocation Location)
Definition JitEngineDevice.h:291
JITKernelInfo(void *Kernel, BinaryInfo &BinInfo, char const *Name, ArrayRef< RuntimeConstantInfo * > RCInfoArray)
Definition JitEngineDevice.h:241
void * getKernel() const
Definition JitEngineDevice.h:249
bool hasModule() const
Definition JitEngineDevice.h:256
const HashT getStaticHash() const
Definition JitEngineDevice.h:270
void setLambdaCalleeInfo(SmallVector< uint64_t > &&LambdaInfo)
Definition JitEngineDevice.h:278
BinaryInfo & getBinaryInfo() const
Definition JitEngineDevice.h:258
ArrayRef< RuntimeConstantInfo * > getRCInfoArray() const
Definition JitEngineDevice.h:255
Module & getModule() const
Definition JitEngineDevice.h:257
void setModule(std::unique_ptr< llvm::Module > Mod)
Definition JitEngineDevice.h:259
const auto & getLambdaCalleeInfo()
Definition JitEngineDevice.h:277
bool hasLambdaCalleeInfo()
Definition JitEngineDevice.h:276
MemoryBufferRef getBitcode()
Definition JitEngineDevice.h:267
bool hasStaticHash() const
Definition JitEngineDevice.h:269
std::unique_ptr< LLVMContext > & getLLVMContext()
Definition JitEngineDevice.h:253
const auto & getLambdaCallsiteLocationInfo() const
Definition JitEngineDevice.h:285
void setBitcode(std::unique_ptr< MemoryBuffer > ExtractedBitcode)
Definition JitEngineDevice.h:264
Definition JitEngineDevice.h:312
MemoryBufferRef getBitcode(JITKernelInfo &KernelInfo)
Definition JitEngineDevice.h:438
void getLambdaJitValues(SmallVector< uint64_t > &LambdaCalleeInfo, LambdaCallsiteRuntimeConstantsMap &LambdaJitValuesMap)
Definition JitEngineDevice.h:449
~JitEngineDevice()
Definition JitEngineDevice.h:552
typename DeviceTraits< ImplT >::DeviceError_t DeviceError_t
Definition JitEngineDevice.h:314
void extractModuleAndBitcode(JITKernelInfo &KernelInfo)
Definition JitEngineDevice.h:399
JitEngineDevice()
Definition JitEngineDevice.h:523
std::unique_ptr< CompilerAsync > AsyncCompiler
Definition JitEngineDevice.h:571
std::unordered_map< const void *, BinaryInfo > HandleToBinaryInfo
Definition JitEngineDevice.h:487
DenseMap< const void *, JITKernelInfo > JITKernelInfoMap
Definition JitEngineDevice.h:568
void registerLambdaCallsiteLocation(void *Kernel, uint64_t LambdaID, uint32_t CallsiteIndex, uint32_t KernelArgIndex, int64_t Offset, RuntimeConstantType StorageType)
Definition JitEngineDevice.h:769
typename DeviceTraits< ImplT >::DeviceStream_t DeviceStream_t
Definition JitEngineDevice.h:315
void finalizeRegistration()
Definition JitEngineDevice.h:723
DenseMap< const void *, LambdaCallsiteLocationMap > PendingLambdaCallsiteLocationInfo
Definition JitEngineDevice.h:570
std::unique_ptr< Dispatcher > Dispatch
Definition JitEngineDevice.h:565
Module & getModule(JITKernelInfo &KernelInfo)
Definition JitEngineDevice.h:428
bool containsJITKernelInfo(const void *Func)
Definition JitEngineDevice.h:490
std::optional< std::reference_wrapper< JITKernelInfo > > getJITKernelInfo(const void *Func)
Definition JitEngineDevice.h:495
std::pair< std::unique_ptr< Module >, std::unique_ptr< MemoryBuffer > > extractKernelModule(BinaryInfo &BinInfo, StringRef KernelName, LLVMContext &Ctx)
Definition JitEngineDevice.h:324
SmallPtrSet< void *, 8 > GlobalLinkedBinaries
Definition JitEngineDevice.h:488
void registerFunction(void *Handle, void *Kernel, char *KernelName, ArrayRef< RuntimeConstantInfo * > RCInfoArray)
Definition JitEngineDevice.h:733
Dispatcher & getDispatcher()
Definition JitEngineDevice.h:560
void registerFatBinary(void *Handle, FatbinWrapperT *FatbinWrapper, const char *ModuleId)
Definition JitEngineDevice.h:692
DeviceError_t compileAndRun(JITKernelInfo &KernelInfo, dim3 GridDim, dim3 BlockDim, void **KernelArgs, uint64_t ShmemSize, typename DeviceTraits< ImplT >::DeviceStream_t Stream)
Definition JitEngineDevice.h:576
void registerLinkedBinary(void *Handle, FatbinWrapperT *FatbinWrapper, const char *ModuleId)
Definition JitEngineDevice.h:784
std::unordered_map< std::string, FatbinWrapperT * > ModuleIdToFatBinary
Definition JitEngineDevice.h:486
typename DeviceTraits< ImplT >::KernelFunction_t KernelFunction_t
Definition JitEngineDevice.h:316
StringRef getDeviceArch() const
Definition JitEngineDevice.h:520
void registerVar(void *Handle, const char *VarName, const void *HostAddr, uint64_t VarSize)
Definition JitEngineDevice.h:461
std::string DeviceArch
Definition JitEngineDevice.h:566
HashT getStaticHash(JITKernelInfo &KernelInfo)
Definition JitEngineDevice.h:502
static JitEngineInfoRegistry & instance()
Definition JitEngineInfoRegistry.h:58
Definition JitEngine.h:32
Definition KernelName.h:18
Definition LambdaRegistry.h:25
static LambdaRegistry & instance()
Definition LambdaRegistry.h:27
static llvm::raw_ostream & outs(const std::string &Name)
Definition Logger.h:25
static void trace(llvm::StringRef Msg)
Definition Logger.h:30
static llvm::raw_ostream & logs(const std::string &Name)
Definition Logger.h:19
Definition TimeTracing.h:33
uint64_t elapsed()
Definition TimeTracing.cpp:66
Definition CompiledLibrary.h:8
Definition MemoryCache.h:27
DenseMap< uint64_t, DenseMap< uint32_t, LambdaKernelArgLocation > > LambdaCallsiteLocationMap
Definition LambdaCallsite.h:29
RuntimeConstantType
Definition CompilerInterfaceTypes.h:20
std::unique_ptr< Module > cloneKernelFromModules(ArrayRef< std::reference_wrapper< Module > > Mods, StringRef EntryName, function_ref< bool(const GlobalValue *)> ShouldCloneDefinition=nullptr)
Definition Cloning.h:513
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 launchKernelDirect(void *KernelFunc, dim3 GridDim, dim3 BlockDim, void **KernelArgs, uint64_t ShmemSize, CUstream Stream)
Definition CoreDeviceCUDA.h:41
HashT hashRuntimeSpecializationConfig(const CodeGenerationConfig &CGConfig)
Definition Hashing.h:252
HashT hashCodeGenConfig(const CodeGenerationConfig &CGConfig)
Definition Hashing.h:232
HashT hashCombine(HashT A, HashT B)
Definition Hashing.h:228
std::string toString(CodegenOption Option)
Definition Config.h:31
void * resolveDeviceGlobalAddr(const void *Addr)
Definition CoreDeviceCUDA.h:29
void internalize(Module &M, StringRef PreserveFunctionName)
Definition CoreLLVM.h:503
DenseMap< uint64_t, LambdaCallsiteRuntimeConstants > LambdaCallsiteRuntimeConstantsMap
Definition LambdaCallsite.h:32
SmallVector< RuntimeConstant, 8 > LambdaCallsiteRuntimeConstants
Definition LambdaCallsite.h:30
void runCleanupPassPipeline(Module &M)
Definition CoreLLVM.h:398
std::unique_ptr< Module > linkModules(LLVMContext &Ctx, SmallVector< std::unique_ptr< Module > > LinkedModules)
Definition CoreLLVM.h:351
Definition Hashing.h:284
Definition CompiledLibrary.h:21
bool GlobalsRelinked
Definition CompiledLibrary.h:36
Definition JitEngineDevice.h:310
int Ret
Definition Dispatcher.h:58
Definition JitEngineDevice.h:77
const char * Binary
Definition JitEngineDevice.h:80
void ** PrelinkedFatbins
Definition JitEngineDevice.h:81
int32_t Magic
Definition JitEngineDevice.h:78
int32_t Version
Definition JitEngineDevice.h:79
Definition GlobalVarInfo.h:5
Definition LambdaCallsite.h:19
int64_t Offset
Definition LambdaCallsite.h:21
RuntimeConstantType StorageType
Definition LambdaCallsite.h:22
uint32_t KernelArgIndex
Definition LambdaCallsite.h:20
Definition Var.h:15