Proteus
Programmable JIT compilation and optimization for C/C++ using LLVM
Loading...
Searching...
No Matches
CompilationTask.h
Go to the documentation of this file.
1#ifndef PROTEUS_COMPILATION_TASK_H
2#define PROTEUS_COMPILATION_TASK_H
3
12#include "proteus/impl/Utils.h"
13
14#include <llvm/Bitcode/BitcodeReader.h>
15#include <llvm/Bitcode/BitcodeWriter.h>
16
17namespace proteus {
18
19using namespace llvm;
20
22private:
23 Dispatcher *Dispatch;
24 MemoryBufferRef Bitcode;
25 HashT HashValue;
26 KernelName Name;
27 dim3 BlockDim;
28 dim3 GridDim;
29 SmallVector<RuntimeConstant> RCVec;
30 SmallVector<uint64_t> LambdaCalleeInfo;
32 std::unordered_map<std::string, GlobalVarInfo> VarNameToGlobalInfo;
33 const CodeGenerationConfig *CGConfig;
34 bool DumpIR;
35 bool RelinkGlobalsByCopy;
36 int MinBlocksPerSM;
37 bool SpecializeArgs;
38 bool SpecializeDims;
39 bool SpecializeDimsRange;
40 bool SpecializeLaunchBounds;
41
42 std::unique_ptr<Module> cloneKernelModule(LLVMContext &Ctx) {
43 TIMESCOPE(CompilationTask, cloneKernelModule);
44 auto ClonedModule = parseBitcodeFile(Bitcode, Ctx);
45 if (auto E = ClonedModule.takeError()) {
46 reportFatalError("Failed to parse bitcode" + toString(std::move(E)));
47 }
48
49 return std::move(*ClonedModule);
50 }
51
52 void dumpOptimizedIR(Module &M) {
53 if (Config::get().traceIRDump()) {
54 llvm::outs() << "LLVM IR module post optimization " << M << "\n";
55 }
56 if (DumpIR) {
57 const auto CreateDumpDirectory = []() {
58 const std::string DumpDirectory = ".proteus-dump";
59 std::filesystem::create_directory(DumpDirectory);
60 return DumpDirectory;
61 };
62
63 static const std::string DumpDirectory = CreateDumpDirectory();
64
65 saveToFile(DumpDirectory + "/device-jit-" + HashValue.toString() + ".ll",
66 M);
67 }
68 }
69
70public:
72 Dispatcher &Dispatch, MemoryBufferRef Bitcode, HashT HashValue,
73 KernelName Name, dim3 BlockDim, dim3 GridDim,
74 const SmallVector<RuntimeConstant> &RCVec,
75 const SmallVector<uint64_t> &LambdaCalleeInfo,
77 const std::unordered_map<std::string, GlobalVarInfo> &VarNameToGlobalInfo,
78 const CodeGenerationConfig &CGConfig, bool DumpIR,
79 bool RelinkGlobalsByCopy)
80 : Dispatch(&Dispatch), Bitcode(Bitcode), HashValue(HashValue),
81 Name(std::move(Name)), BlockDim(BlockDim), GridDim(GridDim),
82 RCVec(RCVec), LambdaCalleeInfo(LambdaCalleeInfo),
84 VarNameToGlobalInfo(VarNameToGlobalInfo), CGConfig(&CGConfig),
85 DumpIR(DumpIR), RelinkGlobalsByCopy(RelinkGlobalsByCopy),
86 MinBlocksPerSM(
87 CGConfig.minBlocksPerSM(BlockDim.x * BlockDim.y * BlockDim.z)),
88 SpecializeArgs(CGConfig.specializeArgs()),
89 SpecializeDims(CGConfig.specializeDims()),
90 SpecializeDimsRange(CGConfig.specializeDimsRange()),
91 SpecializeLaunchBounds(CGConfig.specializeLaunchBounds()) {
92 if (Config::get().traceSpecializations()) {
93 llvm::SmallString<128> S;
94 llvm::raw_svector_ostream OS(S);
95 OS << "[KernelConfig] ID:" << this->Name.base() << " ";
96 CGConfig.dump(OS);
97 OS << "\n";
98 Logger::trace(OS.str());
99 }
100 }
101
102 // Delete copy operations.
105
106 // Use default move operations.
107 CompilationTask(CompilationTask &&) noexcept = default;
108 CompilationTask &operator=(CompilationTask &&) noexcept = default;
109
110 HashT getHashValue() const { return HashValue; }
111
112 std::unique_ptr<MemoryBuffer> compile() {
114 struct TimerRAII {
115 std::chrono::high_resolution_clock::time_point Start, End;
116 HashT HashValue;
117 TimerRAII(HashT HashValue) : HashValue(HashValue) {
118 if (Config::get().ProteusDebugOutput) {
119 Start = std::chrono::high_resolution_clock::now();
120 }
121 }
122
123 ~TimerRAII() {
124 if (Config::get().ProteusDebugOutput) {
125 auto End = std::chrono::high_resolution_clock::now();
126 auto Duration = End - Start;
127 auto Milliseconds =
128 std::chrono::duration_cast<std::chrono::milliseconds>(Duration)
129 .count();
130 Logger::logs("proteus")
131 << "Compiled HashValue " << HashValue.toString() << " for "
132 << Milliseconds << "ms\n";
133 }
134 }
135 } Timer{HashValue};
136
137 LLVMContext Ctx;
138 std::unique_ptr<Module> M = cloneKernelModule(Ctx);
139
140 PROTEUS_DBG(Logger::logfile(HashValue.toString() + ".input.ll", *M));
141
142 proteus::specializeIR(*M, Name.base(), Name.suffix(), BlockDim, GridDim,
143 RCVec, LambdaCalleeInfo,
144 LambdaCallsiteRuntimeConstants, SpecializeArgs,
145 SpecializeDims, SpecializeDimsRange,
146 SpecializeLaunchBounds, MinBlocksPerSM);
147
148 PROTEUS_DBG(Logger::logfile(HashValue.toString() + ".specialized.ll", *M));
149
150 replaceGlobalVariablesWithPointers(*M, VarNameToGlobalInfo);
151
152 Dispatch->optimizeModule(*M, *CGConfig);
153 dumpOptimizedIR(*M);
154
155 auto ObjBuf = Dispatch->codegenModule(*M, *CGConfig);
156 if (!RelinkGlobalsByCopy)
157 proteus::relinkGlobalsObject(ObjBuf->getMemBufferRef(),
158 VarNameToGlobalInfo);
159
160 return ObjBuf;
161 }
162};
163
164} // namespace proteus
165
166#endif
#define PROTEUS_DBG(x)
Definition Debug.h:9
#define TIMESCOPE(...)
Definition TimeTracing.h:66
void saveToFile(llvm::StringRef Filepath, T &&Data)
Definition Utils.h:26
Definition Config.h:167
void dump(T &OS) const
Definition Config.h:311
Definition CompilationTask.h:21
CompilationTask & operator=(const CompilationTask &)=delete
CompilationTask(Dispatcher &Dispatch, MemoryBufferRef Bitcode, HashT HashValue, KernelName Name, dim3 BlockDim, dim3 GridDim, const SmallVector< RuntimeConstant > &RCVec, const SmallVector< uint64_t > &LambdaCalleeInfo, const LambdaCallsiteRuntimeConstantsMap &LambdaCallsiteRuntimeConstants, const std::unordered_map< std::string, GlobalVarInfo > &VarNameToGlobalInfo, const CodeGenerationConfig &CGConfig, bool DumpIR, bool RelinkGlobalsByCopy)
Definition CompilationTask.h:71
HashT getHashValue() const
Definition CompilationTask.h:110
CompilationTask(CompilationTask &&) noexcept=default
CompilationTask(const CompilationTask &)=delete
std::unique_ptr< MemoryBuffer > compile()
Definition CompilationTask.h:112
static Config & get()
Definition Config.h:371
Definition Dispatcher.h:79
virtual void optimizeModule(llvm::Module &M, const CodeGenerationConfig &CGConfig)
Definition Dispatcher.cpp:88
virtual std::unique_ptr< llvm::MemoryBuffer > codegenModule(llvm::Module &M, const CodeGenerationConfig &CGConfig)
Definition Dispatcher.cpp:93
Definition Hashing.h:27
std::string toString() const
Definition Hashing.h:35
Definition KernelName.h:18
std::string suffix() const
Definition KernelName.h:32
const std::string & base() const
Definition KernelName.h:30
static void trace(llvm::StringRef Msg)
Definition Logger.h:30
static void logfile(const std::string &Filename, T &&Data)
Definition Logger.h:33
Definition TimeTracing.h:33
Definition CompiledLibrary.h:8
Definition MemoryCache.h:27
void reportFatalError(const llvm::Twine &Reason, const char *FILE, unsigned Line)
Definition Error.cpp:14
std::string toString(CodegenOption Option)
Definition Config.h:31
DenseMap< uint64_t, LambdaCallsiteRuntimeConstants > LambdaCallsiteRuntimeConstantsMap
Definition LambdaCallsite.h:32
SmallVector< RuntimeConstant, 8 > LambdaCallsiteRuntimeConstants
Definition LambdaCallsite.h:30
Definition Hashing.h:284