Proteus
Programmable JIT compilation and optimization for C/C++ using LLVM
Loading...
Searching...
No Matches
CoreLLVM.h
Go to the documentation of this file.
1#ifndef PROTEUS_CORE_LLVM_H
2#define PROTEUS_CORE_LLVM_H
3
4static_assert(__cplusplus >= 201703L,
5 "This header requires C++17 or later due to LLVM.");
6
7#include "proteus/Error.h"
10#include "proteus/impl/Debug.h"
12#include "proteus/impl/Logger.h"
13
14#include <llvm/Analysis/ValueTracking.h>
15#include <llvm/CodeGen/CommandFlags.h>
16#include <llvm/IR/Constants.h>
17#include <llvm/IR/DebugInfo.h>
18#include <llvm/IR/Metadata.h>
19#include <llvm/IR/Module.h>
20#include <llvm/IR/PassManager.h>
21#include <llvm/Linker/Linker.h>
22#include <llvm/MC/TargetRegistry.h>
23#include <llvm/Passes/PassBuilder.h>
24#if __has_include(<llvm/Plugins/PassPlugin.h>)
25#include <llvm/Plugins/PassPlugin.h>
26#elif __has_include(<llvm/Passes/PassPlugin.h>)
27#include <llvm/Passes/PassPlugin.h>
28#else
29#error "Cannot find LLVM PassPlugin.h"
30#endif
31#include <llvm/Support/TargetSelect.h>
32#include <llvm/Target/TargetMachine.h>
33#include <llvm/Transforms/IPO/MergeFunctions.h>
34
35#if LLVM_VERSION_MAJOR >= 18
36#include <llvm/TargetParser/SubtargetFeature.h>
37// This convoluted logic below is because AMD ROCm 5.7.1 identifies as LLVM 17
38// but includes the header SubtargetFeature.h to a different directory than
39// upstream LLVM 17. We basically detect if it's the HIP version and include it
40// from the expected MC directory, otherwise from TargetParser.
41#elif LLVM_VERSION_MAJOR == 17
42#if defined(__HIP_PLATFORM_HCC__) || defined(HIP_VERSION_MAJOR)
43#include <llvm/MC/SubtargetFeature.h>
44#else
45#include <llvm/TargetParser/SubtargetFeature.h>
46#endif
47#else
48#define STRINGIFY_HELPER(x) #x
49#define STRINGIFY(x) STRINGIFY_HELPER(x)
50#error "Unsupported LLVM version " STRINGIFY(LLVM_VERSION_MAJOR)
51#endif
52#include <llvm/Transforms/IPO/GlobalDCE.h>
53#include <llvm/Transforms/IPO/Internalize.h>
54#include <llvm/Transforms/IPO/StripDeadPrototypes.h>
55#include <llvm/Transforms/IPO/StripSymbols.h>
56#include <llvm/Transforms/Utils/ModuleUtils.h>
57
58#include <algorithm>
59#include <optional>
60#include <string>
61#include <utility>
62#include <vector>
63
64namespace proteus {
65using namespace llvm;
66
67namespace detail {
68
69inline Expected<std::unique_ptr<TargetMachine>>
70createTargetMachine(Module &M, StringRef Arch, unsigned OptLevel = 3) {
71 Triple TT(M.getTargetTriple());
72 auto CGOptLevel = CodeGenOpt::getLevel(OptLevel);
73 if (CGOptLevel == std::nullopt)
74 reportFatalError("Invalid opt level");
75
76 std::string Msg;
77 const Target *T = TargetRegistry::lookupTarget(M.getTargetTriple(), Msg);
78 if (!T)
79 return make_error<StringError>(Msg, inconvertibleErrorCode());
80
81 SubtargetFeatures Features;
82 Features.getDefaultSubtargetFeatures(TT);
83
84 std::optional<Reloc::Model> RelocModel;
85 if (M.getModuleFlag("PIC Level"))
86 RelocModel =
87 M.getPICLevel() == PICLevel::NotPIC ? Reloc::Static : Reloc::PIC_;
88
89 std::optional<CodeModel::Model> CodeModel = M.getCodeModel();
90
91 // Use default target options.
92 // TODO: Customize based on AOT compilation flags or by creating a
93 // constructor that sets target options based on the triple.
94 TargetOptions Options;
95 std::unique_ptr<TargetMachine> TM(T->createTargetMachine(
96 M.getTargetTriple(), Arch, Features.getString(), Options, RelocModel,
97 CodeModel, CGOptLevel.value()));
98 if (!TM)
99 return make_error<StringError>("Failed to create target machine",
100 inconvertibleErrorCode());
101 return TM;
102}
103
104inline std::string getDefaultOptimizationPipeline(char OptLevel) {
105 switch (OptLevel) {
106 case '0':
107 return "default<O0>";
108 case '1':
109 return "default<O1>";
110 case '2':
111 return "default<O2>";
112 case '3':
113 return "default<O3>";
114 case 's':
115 return "default<Os>";
116 case 'z':
117 return "default<Oz>";
118 default:
119 reportFatalError(std::string("Unsupported optimization level ") + OptLevel);
120 }
121 return "";
122}
123
125 std::optional<std::string> PassPipeline, char OptLevel,
126 const std::vector<JITPassPluginConfig> &Plugins) {
127 std::string Pipeline;
128 for (const auto &Plugin : Plugins) {
129 if (!Plugin.Insertion ||
130 Plugin.Insertion->Position != JITPassPluginPosition::Prepend)
131 continue;
132 if (!Pipeline.empty())
133 Pipeline += ",";
134 Pipeline += Plugin.Insertion->Pipeline;
135 }
136
137 if (!Pipeline.empty())
138 Pipeline += ",";
139 Pipeline += PassPipeline ? std::move(*PassPipeline)
141
142 for (const auto &Plugin : Plugins) {
143 if (!Plugin.Insertion ||
144 Plugin.Insertion->Position != JITPassPluginPosition::Append)
145 continue;
146 Pipeline += ",";
147 Pipeline += Plugin.Insertion->Pipeline;
148 }
149
150 return Pipeline;
151}
152
153inline bool
154hasJITPassPluginInsertion(const std::vector<JITPassPluginConfig> &Plugins) {
155 return std::any_of(Plugins.begin(), Plugins.end(),
156 [](const JITPassPluginConfig &Plugin) {
157 return Plugin.Insertion.has_value();
158 });
159}
160
161inline std::vector<std::string>
162getUniqueJITPassPluginPaths(const std::vector<JITPassPluginConfig> &Plugins) {
163 std::vector<std::string> Paths;
164 Paths.reserve(Plugins.size());
165 for (const auto &Plugin : Plugins) {
166 if (std::find(Paths.begin(), Paths.end(), Plugin.Path) == Paths.end())
167 Paths.push_back(Plugin.Path);
168 }
169 return Paths;
170}
171
172inline std::vector<PassPlugin>
173loadJITPassPlugins(const std::vector<JITPassPluginConfig> &Plugins) {
174 std::vector<PassPlugin> LoadedPlugins;
175 const auto PluginPaths = getUniqueJITPassPluginPaths(Plugins);
176 LoadedPlugins.reserve(PluginPaths.size());
177 for (const auto &PluginPath : PluginPaths) {
178 auto LoadedPlugin = PassPlugin::Load(PluginPath);
179 if (!LoadedPlugin)
180 reportFatalError("Failed to load JIT pass plugin '" + PluginPath +
181 "': " + toString(LoadedPlugin.takeError()));
182 LoadedPlugins.push_back(std::move(*LoadedPlugin));
183 }
184 return LoadedPlugins;
185}
186
188 Module &M, StringRef Arch, const std::string &PassPipeline,
189 unsigned CodegenOptLevel,
190 const std::vector<JITPassPluginConfig> &Plugins = {}) {
191 PipelineTuningOptions PTO;
192
193 std::optional<PGOOptions> PGOOpt;
194 auto TM = createTargetMachine(M, Arch, CodegenOptLevel);
195 if (auto Err = TM.takeError())
196 report_fatal_error(std::move(Err));
197 TargetLibraryInfoImpl TLII(Triple(M.getTargetTriple()));
198
199 auto LoadedPlugins = loadJITPassPlugins(Plugins);
200 PassBuilder PB(TM->get(), PTO, PGOOpt, nullptr);
201 for (const auto &Plugin : LoadedPlugins)
202 Plugin.registerPassBuilderCallbacks(PB);
203 LoopAnalysisManager LAM;
204 FunctionAnalysisManager FAM;
205 CGSCCAnalysisManager CGAM;
206 ModuleAnalysisManager MAM;
207
208 FAM.registerPass([&] { return TargetLibraryAnalysis(TLII); });
209
210 PB.registerModuleAnalyses(MAM);
211 PB.registerCGSCCAnalyses(CGAM);
212 PB.registerFunctionAnalyses(FAM);
213 PB.registerLoopAnalyses(LAM);
214 PB.crossRegisterProxies(LAM, FAM, CGAM, MAM);
215 ModulePassManager Passes;
216 if (auto E = PB.parsePassPipeline(Passes, PassPipeline))
217 reportFatalError("Error: " + toString(std::move(E)));
218
219 Passes.run(M, MAM);
220}
221
223 Module &M, StringRef Arch, char OptLevel = '3',
224 unsigned CodegenOptLevel = 3,
225 const std::vector<JITPassPluginConfig> &Plugins = {}) {
226 PipelineTuningOptions PTO;
227
228 std::optional<PGOOptions> PGOOpt;
229 auto TM = createTargetMachine(M, Arch, CodegenOptLevel);
230 if (auto Err = TM.takeError())
231 report_fatal_error(std::move(Err));
232 TargetLibraryInfoImpl TLII(Triple(M.getTargetTriple()));
233
234 auto LoadedPlugins = loadJITPassPlugins(Plugins);
235 PassBuilder PB(TM->get(), PTO, PGOOpt, nullptr);
236 for (const auto &Plugin : LoadedPlugins)
237 Plugin.registerPassBuilderCallbacks(PB);
238 LoopAnalysisManager LAM;
239 FunctionAnalysisManager FAM;
240 CGSCCAnalysisManager CGAM;
241 ModuleAnalysisManager MAM;
242
243 FAM.registerPass([&] { return TargetLibraryAnalysis(TLII); });
244
245 PB.registerModuleAnalyses(MAM);
246 PB.registerCGSCCAnalyses(CGAM);
247 PB.registerFunctionAnalyses(FAM);
248 PB.registerLoopAnalyses(LAM);
249 PB.crossRegisterProxies(LAM, FAM, CGAM, MAM);
250
251 OptimizationLevel OptSetting;
252 switch (OptLevel) {
253 case '0':
254 OptSetting = OptimizationLevel::O0;
255 break;
256 case '1':
257 OptSetting = OptimizationLevel::O1;
258 break;
259 case '2':
260 OptSetting = OptimizationLevel::O2;
261 break;
262 case '3':
263 OptSetting = OptimizationLevel::O3;
264 break;
265 case 's':
266 OptSetting = OptimizationLevel::Os;
267 break;
268 case 'z':
269 OptSetting = OptimizationLevel::Oz;
270 break;
271 default:
272 reportFatalError(std::string("Unsupported optimization level ") + OptLevel);
273 };
274
275 ModulePassManager Passes = PB.buildPerModuleDefaultPipeline(OptSetting);
276 Passes.run(M, MAM);
277}
278
279} // namespace detail
280
283 InitializeAllTargetInfos();
284 InitializeAllTargets();
285 InitializeAllTargetMCs();
286 InitializeAllAsmParsers();
287 InitializeAllAsmPrinters();
288 }
289};
290
292 std::optional<std::string> PassPipeline;
295
297 : OptLevel(CGConfig.optLevel()),
298 CodegenOptLevel(CGConfig.codeGenOptLevel()) {
299 if (auto Pipeline = CGConfig.optPipeline())
300 PassPipeline = Pipeline.value();
301 }
302
307};
308
309inline void optimizeIR(Module &M, StringRef Arch,
310 const OptimizationPipelineConfig &OptConfig) {
311 TIMESCOPE("proteus::optimizeIR");
312 Timer T(Config::get().ProteusEnableTimers);
313
314 const auto Plugins = getJITPassPluginConfigs();
315 const bool UseTextualPipeline =
317 const std::string FinalPipeline =
318 UseTextualPipeline
320 OptConfig.OptLevel, Plugins)
321 : std::string();
322
323 if (UseTextualPipeline) {
324 auto TraceOut = [](const std::string &PassPipeline) {
325 SmallString<128> S;
326 raw_svector_ostream OS(S);
327 OS << "[CustomPipeline] " << PassPipeline << "\n";
328 return S;
329 };
330
331 if (Config::get().traceSpecializations())
332 Logger::trace(TraceOut(FinalPipeline));
333
334 detail::runOptimizationPassPipeline(M, Arch, FinalPipeline,
335 OptConfig.CodegenOptLevel, Plugins);
336 } else {
338 OptConfig.CodegenOptLevel, Plugins);
339 }
340
342 << "optimizeIR optlevel "
343 << (UseTextualPipeline
344 ? StringRef(FinalPipeline)
345 : StringRef(&OptConfig.OptLevel, 1))
346 << " codegenopt " << OptConfig.CodegenOptLevel << " "
347 << T.elapsed() << " ms\n");
348}
349
350inline std::unique_ptr<Module>
351linkModules(LLVMContext &Ctx,
352 SmallVector<std::unique_ptr<Module>> LinkedModules) {
353 TIMESCOPE("proteus::linkModules");
354 if (LinkedModules.empty())
355 reportFatalError("Expected jit module");
356
357 auto LinkedModule = std::make_unique<llvm::Module>("JitModule", Ctx);
358 Linker IRLinker(*LinkedModule);
359 // Link in all the proteus-enabled extracted modules.
360 for (auto &LinkedM : LinkedModules) {
361 // Returns true if linking failed.
362 if (IRLinker.linkInModule(std::move(LinkedM)))
363 reportFatalError("Linking failed");
364 }
365
366 return LinkedModule;
367}
368
369inline void pruneDanglingNVVMAnnotations(Module &M) {
370 NamedMDNode *Annotations = M.getNamedMetadata("nvvm.annotations");
371 if (!Annotations)
372 return;
373
374 SmallVector<MDNode *> LiveEntries;
375 for (MDNode *Entry : Annotations->operands()) {
376 // Global DCE nulls out the entry of a function it removes and stripping
377 // debug info rewrites such a node to an empty tuple, !{}.
378 // LLVM's nvvm.annotations which LLVM's later UpgradeNVVMAnnotations
379 // pass reads unconditionally, so we need to prune out such entries.
380 // Otherwise, UpgradeNVVMAnnotations runs out of bounds!
381 if (!Entry || Entry->getNumOperands() == 0)
382 continue;
383
384 if (!mdconst::dyn_extract_or_null<GlobalValue>(Entry->getOperand(0)))
385 continue;
386
387 LiveEntries.push_back(Entry);
388 }
389
390 if (LiveEntries.size() == Annotations->getNumOperands())
391 return;
392
393 Annotations->clearOperands();
394 for (MDNode *Entry : LiveEntries)
395 Annotations->addOperand(Entry);
396}
397
398inline void runCleanupPassPipeline(Module &M) {
399 TIMESCOPE("proteus::runCleanupPassPipeline");
400 PassBuilder PB;
401 LoopAnalysisManager LAM;
402 FunctionAnalysisManager FAM;
403 CGSCCAnalysisManager CGAM;
404 ModuleAnalysisManager MAM;
405
406 PB.registerModuleAnalyses(MAM);
407 PB.registerCGSCCAnalyses(CGAM);
408 PB.registerFunctionAnalyses(FAM);
409 PB.registerLoopAnalyses(LAM);
410 PB.crossRegisterProxies(LAM, FAM, CGAM, MAM);
411
412 ModulePassManager Passes;
413 Passes.addPass(GlobalDCEPass());
414 // Passes.addPass(StripDeadDebugInfoPass());
415 Passes.addPass(StripDeadPrototypesPass());
416
417 Passes.run(M, MAM);
418
420
421 // Keep line-table debug info so that recorded device modules retain the
422 // kernel's source file and line information.
423 stripNonLineTableDebugInfo(M);
424}
425
427 llvm::Module &M, llvm::StringRef Key,
428 llvm::SmallVectorImpl<std::pair<llvm::Function *, std::uint64_t>> &Out) {
429 llvm::SmallDenseMap<llvm::Function *, std::uint64_t, 32> Seen;
430 for (llvm::Function &F : M) {
431 llvm::MDNode *Node = F.getMetadata(Key);
432 if (!Node || Node->getNumOperands() < 1)
433 continue;
434 auto *CAM = llvm::dyn_cast<llvm::ConstantAsMetadata>(Node->getOperand(0));
435 auto *CI =
436 CAM ? llvm::dyn_cast<llvm::ConstantInt>(CAM->getValue()) : nullptr;
437 if (!CI)
438 continue;
439 std::uint64_t Id = CI->getZExtValue();
440 if (Seen.try_emplace(&F, Id).second)
441 Out.emplace_back(&F, Id);
442 }
443}
444
445inline void
446findFunctionsWithU64Metadata(llvm::Module &M, llvm::StringRef Key,
447 llvm::SmallVectorImpl<std::uint64_t> &Out) {
448 llvm::SmallDenseMap<llvm::Function *, std::uint64_t, 32> Seen;
449 for (llvm::Function &F : M) {
450 llvm::MDNode *Node = F.getMetadata(Key);
451 if (!Node || Node->getNumOperands() < 1)
452 continue;
453 auto *CAM = llvm::dyn_cast<llvm::ConstantAsMetadata>(Node->getOperand(0));
454 auto *CI =
455 CAM ? llvm::dyn_cast<llvm::ConstantInt>(CAM->getValue()) : nullptr;
456 if (!CI)
457 continue;
458 std::uint64_t Id = CI->getZExtValue();
459 if (Seen.try_emplace(&F, Id).second)
460 Out.emplace_back(Id);
461 }
462}
463
464inline bool hasU64Metadata(Function *F, StringRef Key) {
465 return F->getMetadata(Key) != nullptr;
466}
467
468inline void pruneIR(Module &M, bool UnsetExternallyInitialized = true) {
469 // Remove llvm.global.annotations now that we have read them.
470 if (auto *GlobalAnnotations = M.getGlobalVariable("llvm.global.annotations"))
471 M.eraseGlobalVariable(GlobalAnnotations);
472
473 // Remove llvm.compiler.used
474 if (auto *CompilerUsed = M.getGlobalVariable("llvm.compiler.used"))
475 M.eraseGlobalVariable(CompilerUsed);
476
477 // Remove the __clang_gpu_used_external used in HIP RDC compilation and its
478 // uses in llvm.used, llvm.compiler.used.
479 SmallVector<GlobalVariable *> GlobalsToErase;
480 for (auto &GV : M.globals()) {
481 auto Name = GV.getName();
482 if (Name.starts_with("__clang_gpu_used_external") ||
483 Name.starts_with("_jit_bitcode") || Name.starts_with("__hip_cuid")) {
484 GlobalsToErase.push_back(&GV);
485 removeFromUsedLists(M, [&GV](Constant *C) {
486 if (auto *Global = dyn_cast<GlobalVariable>(C))
487 return Global == &GV;
488 return false;
489 });
490 }
491 }
492 for (auto *GV : GlobalsToErase) {
493 M.eraseGlobalVariable(GV);
494 }
495
496 // Remove externaly_initialized attributes.
497 if (UnsetExternallyInitialized)
498 for (auto &GV : M.globals())
499 if (GV.isExternallyInitialized())
500 GV.setExternallyInitialized(false);
501}
502
503inline void internalize(Module &M, StringRef PreserveFunctionName) {
504 auto *F = M.getFunction(PreserveFunctionName);
505 // Internalize others besides the kernel function.
506 internalizeModule(M, [&F](const GlobalValue &GV) {
507 // Do not internalize the kernel function.
508 if (&GV == F)
509 return true;
510
511 // Internalize everything else.
512 return false;
513 });
514}
515
516inline std::optional<uint64_t> getFunctionU64Metadata(Function &F,
517 StringRef Key) {
518 MDNode *Node = F.getMetadata(Key);
519 if (!Node || Node->getNumOperands() < 1)
520 return std::nullopt;
521
522 auto *CAM = dyn_cast<ConstantAsMetadata>(Node->getOperand(0));
523 auto *CI = CAM ? dyn_cast<ConstantInt>(CAM->getValue()) : nullptr;
524 if (!CI)
525 return std::nullopt;
526
527 return CI->getZExtValue();
528}
529
530} // namespace proteus
531
532#endif
#define PROTEUS_TIMER_OUTPUT(x)
Definition Config.h:483
#define TIMESCOPE(...)
Definition TimeTracing.h:66
Definition Config.h:167
std::optional< const std::string > optPipeline() const
Definition Config.h:295
static Config & get()
Definition Config.h:371
static llvm::raw_ostream & outs(const std::string &Name)
Definition Logger.h:25
static void trace(llvm::StringRef Msg)
Definition Logger.h:30
Definition TimeTracing.h:33
uint64_t elapsed()
Definition TimeTracing.cpp:66
Definition CompiledLibrary.h:8
std::vector< std::string > getUniqueJITPassPluginPaths(const std::vector< JITPassPluginConfig > &Plugins)
Definition CoreLLVM.h:162
std::vector< PassPlugin > loadJITPassPlugins(const std::vector< JITPassPluginConfig > &Plugins)
Definition CoreLLVM.h:173
std::string composeOptimizationPassPipeline(std::optional< std::string > PassPipeline, char OptLevel, const std::vector< JITPassPluginConfig > &Plugins)
Definition CoreLLVM.h:124
Expected< std::unique_ptr< TargetMachine > > createTargetMachine(Module &M, StringRef Arch, unsigned OptLevel=3)
Definition CoreLLVM.h:70
bool hasJITPassPluginInsertion(const std::vector< JITPassPluginConfig > &Plugins)
Definition CoreLLVM.h:154
void runOptimizationPassPipeline(Module &M, StringRef Arch, const std::string &PassPipeline, unsigned CodegenOptLevel, const std::vector< JITPassPluginConfig > &Plugins={})
Definition CoreLLVM.h:187
std::string getDefaultOptimizationPipeline(char OptLevel)
Definition CoreLLVM.h:104
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
void findFunctionsWithU64Metadata(llvm::Module &M, llvm::StringRef Key, llvm::SmallVectorImpl< std::pair< llvm::Function *, std::uint64_t > > &Out)
Definition CoreLLVM.h:426
void reportFatalError(const llvm::Twine &Reason, const char *FILE, unsigned Line)
Definition Error.cpp:14
void pruneIR(Module &M, bool UnsetExternallyInitialized=true)
Definition CoreLLVM.h:468
std::string toString(CodegenOption Option)
Definition Config.h:31
void internalize(Module &M, StringRef PreserveFunctionName)
Definition CoreLLVM.h:503
void pruneDanglingNVVMAnnotations(Module &M)
Definition CoreLLVM.h:369
bool hasU64Metadata(Function *F, StringRef Key)
Definition CoreLLVM.h:464
std::optional< uint64_t > getFunctionU64Metadata(Function &F, StringRef Key)
Definition CoreLLVM.h:516
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 CoreLLVM.h:281
InitLLVMTargets()
Definition CoreLLVM.h:282
Definition JITPassPluginRegistry.h:21
Definition CoreLLVM.h:291
std::optional< std::string > PassPipeline
Definition CoreLLVM.h:292
OptimizationPipelineConfig(const CodeGenerationConfig &CGConfig)
Definition CoreLLVM.h:296
unsigned CodegenOptLevel
Definition CoreLLVM.h:294
OptimizationPipelineConfig(std::optional< std::string > PassPipeline, char OptLevel, unsigned CodegenOptLevel)
Definition CoreLLVM.h:303
char OptLevel
Definition CoreLLVM.h:293