1#ifndef PROTEUS_CORE_LLVM_H
2#define PROTEUS_CORE_LLVM_H
4static_assert(__cplusplus >= 201703L,
5 "This header requires C++17 or later due to LLVM.");
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>
29#error "Cannot find LLVM PassPlugin.h"
31#include <llvm/Support/TargetSelect.h>
32#include <llvm/Target/TargetMachine.h>
33#include <llvm/Transforms/IPO/MergeFunctions.h>
35#if LLVM_VERSION_MAJOR >= 18
36#include <llvm/TargetParser/SubtargetFeature.h>
41#elif LLVM_VERSION_MAJOR == 17
42#if defined(__HIP_PLATFORM_HCC__) || defined(HIP_VERSION_MAJOR)
43#include <llvm/MC/SubtargetFeature.h>
45#include <llvm/TargetParser/SubtargetFeature.h>
48#define STRINGIFY_HELPER(x) #x
49#define STRINGIFY(x) STRINGIFY_HELPER(x)
50#error "Unsupported LLVM version " STRINGIFY(LLVM_VERSION_MAJOR)
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>
69inline Expected<std::unique_ptr<TargetMachine>>
71 Triple TT(M.getTargetTriple());
72 auto CGOptLevel = CodeGenOpt::getLevel(OptLevel);
73 if (CGOptLevel == std::nullopt)
77 const Target *T = TargetRegistry::lookupTarget(M.getTargetTriple(), Msg);
79 return make_error<StringError>(Msg, inconvertibleErrorCode());
81 SubtargetFeatures Features;
82 Features.getDefaultSubtargetFeatures(TT);
84 std::optional<Reloc::Model> RelocModel;
85 if (M.getModuleFlag(
"PIC Level"))
87 M.getPICLevel() == PICLevel::NotPIC ? Reloc::Static : Reloc::PIC_;
89 std::optional<CodeModel::Model> CodeModel = M.getCodeModel();
94 TargetOptions Options;
95 std::unique_ptr<TargetMachine> TM(T->createTargetMachine(
96 M.getTargetTriple(), Arch, Features.getString(), Options, RelocModel,
97 CodeModel, CGOptLevel.value()));
99 return make_error<StringError>(
"Failed to create target machine",
100 inconvertibleErrorCode());
107 return "default<O0>";
109 return "default<O1>";
111 return "default<O2>";
113 return "default<O3>";
115 return "default<Os>";
117 return "default<Oz>";
119 reportFatalError(std::string(
"Unsupported optimization level ") + OptLevel);
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 ||
132 if (!Pipeline.empty())
134 Pipeline += Plugin.Insertion->Pipeline;
137 if (!Pipeline.empty())
139 Pipeline += PassPipeline ? std::move(*PassPipeline)
142 for (
const auto &Plugin : Plugins) {
143 if (!Plugin.Insertion ||
147 Pipeline += Plugin.Insertion->Pipeline;
155 return std::any_of(Plugins.begin(), Plugins.end(),
157 return Plugin.Insertion.has_value();
161inline std::vector<std::string>
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);
172inline std::vector<PassPlugin>
174 std::vector<PassPlugin> LoadedPlugins;
176 LoadedPlugins.reserve(PluginPaths.size());
177 for (
const auto &PluginPath : PluginPaths) {
178 auto LoadedPlugin = PassPlugin::Load(PluginPath);
181 "': " +
toString(LoadedPlugin.takeError()));
182 LoadedPlugins.push_back(std::move(*LoadedPlugin));
184 return LoadedPlugins;
188 Module &M, StringRef Arch,
const std::string &PassPipeline,
189 unsigned CodegenOptLevel,
190 const std::vector<JITPassPluginConfig> &Plugins = {}) {
191 PipelineTuningOptions PTO;
193 std::optional<PGOOptions> PGOOpt;
195 if (
auto Err = TM.takeError())
196 report_fatal_error(std::move(Err));
197 TargetLibraryInfoImpl TLII(Triple(M.getTargetTriple()));
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;
208 FAM.registerPass([&] {
return TargetLibraryAnalysis(TLII); });
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))
223 Module &M, StringRef Arch,
char OptLevel =
'3',
224 unsigned CodegenOptLevel = 3,
225 const std::vector<JITPassPluginConfig> &Plugins = {}) {
226 PipelineTuningOptions PTO;
228 std::optional<PGOOptions> PGOOpt;
230 if (
auto Err = TM.takeError())
231 report_fatal_error(std::move(Err));
232 TargetLibraryInfoImpl TLII(Triple(M.getTargetTriple()));
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;
243 FAM.registerPass([&] {
return TargetLibraryAnalysis(TLII); });
245 PB.registerModuleAnalyses(MAM);
246 PB.registerCGSCCAnalyses(CGAM);
247 PB.registerFunctionAnalyses(FAM);
248 PB.registerLoopAnalyses(LAM);
249 PB.crossRegisterProxies(LAM, FAM, CGAM, MAM);
251 OptimizationLevel OptSetting;
254 OptSetting = OptimizationLevel::O0;
257 OptSetting = OptimizationLevel::O1;
260 OptSetting = OptimizationLevel::O2;
263 OptSetting = OptimizationLevel::O3;
266 OptSetting = OptimizationLevel::Os;
269 OptSetting = OptimizationLevel::Oz;
272 reportFatalError(std::string(
"Unsupported optimization level ") + OptLevel);
275 ModulePassManager Passes = PB.buildPerModuleDefaultPipeline(OptSetting);
283 InitializeAllTargetInfos();
284 InitializeAllTargets();
285 InitializeAllTargetMCs();
286 InitializeAllAsmParsers();
287 InitializeAllAsmPrinters();
315 const bool UseTextualPipeline =
317 const std::string FinalPipeline =
323 if (UseTextualPipeline) {
324 auto TraceOut = [](
const std::string &PassPipeline) {
326 raw_svector_ostream OS(S);
327 OS <<
"[CustomPipeline] " << PassPipeline <<
"\n";
342 <<
"optimizeIR optlevel "
343 << (UseTextualPipeline
344 ? StringRef(FinalPipeline)
345 : StringRef(&OptConfig.
OptLevel, 1))
350inline std::unique_ptr<Module>
352 SmallVector<std::unique_ptr<Module>> LinkedModules) {
354 if (LinkedModules.empty())
357 auto LinkedModule = std::make_unique<llvm::Module>(
"JitModule", Ctx);
358 Linker IRLinker(*LinkedModule);
360 for (
auto &LinkedM : LinkedModules) {
362 if (IRLinker.linkInModule(std::move(LinkedM)))
370 NamedMDNode *Annotations = M.getNamedMetadata(
"nvvm.annotations");
374 SmallVector<MDNode *> LiveEntries;
375 for (MDNode *Entry : Annotations->operands()) {
381 if (!Entry || Entry->getNumOperands() == 0)
384 if (!mdconst::dyn_extract_or_null<GlobalValue>(Entry->getOperand(0)))
387 LiveEntries.push_back(Entry);
390 if (LiveEntries.size() == Annotations->getNumOperands())
393 Annotations->clearOperands();
394 for (MDNode *Entry : LiveEntries)
395 Annotations->addOperand(Entry);
399 TIMESCOPE(
"proteus::runCleanupPassPipeline");
401 LoopAnalysisManager LAM;
402 FunctionAnalysisManager FAM;
403 CGSCCAnalysisManager CGAM;
404 ModuleAnalysisManager MAM;
406 PB.registerModuleAnalyses(MAM);
407 PB.registerCGSCCAnalyses(CGAM);
408 PB.registerFunctionAnalyses(FAM);
409 PB.registerLoopAnalyses(LAM);
410 PB.crossRegisterProxies(LAM, FAM, CGAM, MAM);
412 ModulePassManager Passes;
413 Passes.addPass(GlobalDCEPass());
415 Passes.addPass(StripDeadPrototypesPass());
423 stripNonLineTableDebugInfo(M);
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)
434 auto *CAM = llvm::dyn_cast<llvm::ConstantAsMetadata>(Node->getOperand(0));
436 CAM ? llvm::dyn_cast<llvm::ConstantInt>(CAM->getValue()) :
nullptr;
439 std::uint64_t Id = CI->getZExtValue();
440 if (Seen.try_emplace(&F, Id).second)
441 Out.emplace_back(&F, Id);
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)
453 auto *CAM = llvm::dyn_cast<llvm::ConstantAsMetadata>(Node->getOperand(0));
455 CAM ? llvm::dyn_cast<llvm::ConstantInt>(CAM->getValue()) :
nullptr;
458 std::uint64_t Id = CI->getZExtValue();
459 if (Seen.try_emplace(&F, Id).second)
460 Out.emplace_back(Id);
465 return F->getMetadata(Key) !=
nullptr;
468inline void pruneIR(Module &M,
bool UnsetExternallyInitialized =
true) {
470 if (
auto *GlobalAnnotations = M.getGlobalVariable(
"llvm.global.annotations"))
471 M.eraseGlobalVariable(GlobalAnnotations);
474 if (
auto *CompilerUsed = M.getGlobalVariable(
"llvm.compiler.used"))
475 M.eraseGlobalVariable(CompilerUsed);
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;
492 for (
auto *GV : GlobalsToErase) {
493 M.eraseGlobalVariable(GV);
497 if (UnsetExternallyInitialized)
498 for (
auto &GV : M.globals())
499 if (GV.isExternallyInitialized())
500 GV.setExternallyInitialized(
false);
503inline void internalize(Module &M, StringRef PreserveFunctionName) {
504 auto *F = M.getFunction(PreserveFunctionName);
506 internalizeModule(M, [&F](
const GlobalValue &GV) {
518 MDNode *Node = F.getMetadata(Key);
519 if (!Node || Node->getNumOperands() < 1)
522 auto *CAM = dyn_cast<ConstantAsMetadata>(Node->getOperand(0));
523 auto *CI = CAM ? dyn_cast<ConstantInt>(CAM->getValue()) :
nullptr;
527 return CI->getZExtValue();
#define PROTEUS_TIMER_OUTPUT(x)
Definition Config.h:483
#define TIMESCOPE(...)
Definition TimeTracing.h:66
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 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