Proteus
Programmable JIT compilation and optimization for C/C++ using LLVM
Loading...
Searching...
No Matches
DispatcherHIP.h
Go to the documentation of this file.
1#ifndef PROTEUS_FRONTEND_DISPATCHER_HIP_H
2#define PROTEUS_FRONTEND_DISPATCHER_HIP_H
3
4#if PROTEUS_ENABLE_HIP
5
6#include "proteus/Error.h"
10
11#include <llvm/Bitcode/BitcodeReader.h>
12#include <llvm/Linker/Linker.h>
13#include <llvm/Support/FileSystem.h>
14#include <llvm/Support/MemoryBuffer.h>
15
16namespace proteus {
17
18class DispatcherHIP : public DispatcherDevice<JitEngineDeviceHIP> {
19public:
20 static DispatcherHIP &instance() {
21 static DispatcherHIP D{"DispatcherHIP", JitEngineDeviceHIP::instance()};
22 return D;
23 }
24
25 DispatcherHIP(const std::string &Label, JitEngineDeviceHIP &Jit)
26 : DispatcherDevice(Label, TargetModelType::HIP, Jit) {}
27
28protected:
29 // Link ROCm device libraries (ocml/ockl + oclc config) so HIPRTC can
30 // resolve __ocml_* calls produced by math lowering.
31 void linkDeviceLibraries(Module &M) override {
32 TIMESCOPE(DispatcherHIP, linkDeviceLibraries);
33 const auto &Toolchain = resolveHIPToolchain();
34
35 auto LoadBitcode = [&](const llvm::SmallString<256> &Path) {
36 auto BufferOrErr = llvm::MemoryBuffer::getFile(Path);
37 if (!BufferOrErr || !BufferOrErr.get())
38 reportFatalError("DispatchHIP: failed to read ROCm bitcode file: " +
39 Path.str().str() + " (" + Toolchain.Origin + ")");
40 auto Parsed = llvm::parseBitcodeFile(
41 BufferOrErr->get()->getMemBufferRef(), M.getContext());
42 if (!Parsed)
43 reportFatalError("DispatchHIP: failed to parse ROCm bitcode file: " +
44 Path.str().str() + " (" + Toolchain.Origin + ")");
45 return std::move(Parsed.get());
46 };
47
48 auto AppendBitcodePath =
49 [&](llvm::SmallVectorImpl<llvm::SmallString<256>> &Paths,
50 llvm::StringRef Filename) {
51 llvm::SmallString<256> Path{Toolchain.DeviceLibDir};
52 llvm::sys::path::append(Path, Filename);
53 Paths.push_back(std::move(Path));
54 };
55
56 auto Exists = [&](llvm::StringRef Filename) -> bool {
57 llvm::SmallString<256> Path{Toolchain.DeviceLibDir};
58 llvm::sys::path::append(Path, Filename);
59 return llvm::sys::fs::exists(Path);
60 };
61
62 auto PickFirstExisting =
63 [&](std::initializer_list<llvm::StringRef> Candidates)
64 -> llvm::StringRef {
65 for (auto C : Candidates) {
66 if (Exists(C))
67 return C;
68 }
69 return {};
70 };
71
72 llvm::SmallVector<llvm::SmallString<256>, 8> LibsToLink;
73 AppendBitcodePath(LibsToLink, "ocml.bc");
74 AppendBitcodePath(LibsToLink, "ockl.bc");
75
76 // ABI: prefer the newest available.
77 if (auto Abi = PickFirstExisting({"oclc_abi_version_600.bc",
78 "oclc_abi_version_500.bc",
79 "oclc_abi_version_400.bc"});
80 !Abi.empty()) {
81 AppendBitcodePath(LibsToLink, Abi);
82 } else {
84 std::string("DispatchHIP: missing oclc ABI bitcode under ") +
85 Toolchain.DeviceLibDir + " (" + Toolchain.Origin +
86 "; expected oclc_abi_version_{600,500,400}.bc)");
87 }
88
89 // ISA: derived from device arch like "gfx90a" -> "90a".
90 const std::string DeviceArch = Jit.getDeviceArch().str();
91 if (!llvm::StringRef{DeviceArch}.starts_with("gfx"))
92 reportFatalError("DispatchHIP: unexpected HIP device arch: " +
93 DeviceArch);
94 const llvm::StringRef IsaSuffix = llvm::StringRef{DeviceArch}.drop_front(3);
95 const std::string IsaFile = ("oclc_isa_version_" + IsaSuffix + ".bc").str();
96 if (!Exists(IsaFile))
97 reportFatalError(std::string("DispatchHIP: missing ISA bitcode file ") +
98 IsaFile + " under " + Toolchain.DeviceLibDir + " (" +
99 Toolchain.Origin + "; DeviceArch=" + DeviceArch + ")");
100 AppendBitcodePath(LibsToLink, IsaFile);
101
102 // Math/FP mode defaults (safe defaults, can be revisited later).
103 AppendBitcodePath(LibsToLink, "oclc_unsafe_math_off.bc");
104 AppendBitcodePath(LibsToLink, "oclc_finite_only_off.bc");
105 AppendBitcodePath(LibsToLink, "oclc_daz_opt_off.bc");
106 AppendBitcodePath(LibsToLink, "oclc_correctly_rounded_sqrt_on.bc");
107
108 // Wavefront size selection: RDNA is typically wave32; CDNA/gfx9 wave64.
109 const bool IsWave32 = llvm::StringRef{DeviceArch}.starts_with("gfx10") ||
110 llvm::StringRef{DeviceArch}.starts_with("gfx11") ||
111 llvm::StringRef{DeviceArch}.starts_with("gfx12");
112 AppendBitcodePath(LibsToLink, IsWave32 ? "oclc_wavefrontsize64_off.bc"
113 : "oclc_wavefrontsize64_on.bc");
114
115 llvm::Linker Linker{M};
116 for (const auto &Path : LibsToLink) {
117 auto LibMod = LoadBitcode(Path);
118 Linker.linkInModule(std::move(LibMod),
119 llvm::Linker::Flags::LinkOnlyNeeded);
120 }
121 }
122};
123
124} // namespace proteus
125
126#endif
127
128#endif // PROTEUS_FRONTEND_DISPATCHER_HIP_H
#define TIMESCOPE(...)
Definition TimeTracing.h:66
Definition MemoryCache.h:27
TargetModelType
Definition TargetModel.h:8
void reportFatalError(const llvm::Twine &Reason, const char *FILE, unsigned Line)
Definition Error.cpp:14
const ResolvedHIPToolchain & resolveHIPToolchain()
Definition HIPToolchain.cpp:273