1use std::{
4 ffi::c_void,
5 fmt,
6 path::{Path, PathBuf},
7 sync::Arc,
8};
9
10use libloading::Library;
11
12const DRIVER_NAMES: &[&str] = &["libcuda.so.1", "libcuda.so", "nvcuda.dll"];
13const CUBLAS_NAMES: &[&str] = &["libcublas.so.13", "libcublas.so.12", "libcublas.so"];
14const CUBLASLT_NAMES: &[&str] = &["libcublasLt.so.13", "libcublasLt.so.12", "libcublasLt.so"];
15
16const DRIVER_SYMBOLS: &[&str] = &["cuInit", "cuDriverGetVersion"];
17const CUBLAS_SYMBOLS: &[&str] = &[
18 "cublasCreate_v2",
19 "cublasDestroy_v2",
20 "cublasSgemm_v2",
21 "cublasGemmEx",
22];
23const CUBLASLT_SYMBOLS: &[&str] = &["cublasLtCreate", "cublasLtDestroy", "cublasLtMatmul"];
24
25#[derive(Clone, Debug, PartialEq, Eq)]
27pub struct CudaSymbolEvidence {
28 pub name: String,
30 pub present: bool,
32}
33
34#[derive(Clone, Debug, PartialEq, Eq)]
36pub struct CudaAbiEvidence {
37 pub driver_library: String,
39 pub cublas_library: String,
41 pub cublaslt_library: String,
43 pub driver_version: Option<i32>,
45 pub driver_symbols: Vec<CudaSymbolEvidence>,
47 pub cublas_symbols: Vec<CudaSymbolEvidence>,
49 pub cublaslt_symbols: Vec<CudaSymbolEvidence>,
51}
52
53impl CudaAbiEvidence {
54 pub fn is_complete(&self) -> bool {
56 self.driver_symbols.iter().all(|symbol| symbol.present)
57 && self.cublas_symbols.iter().all(|symbol| symbol.present)
58 && self.cublaslt_symbols.iter().all(|symbol| symbol.present)
59 }
60
61 pub fn supports_half_matmul(&self) -> bool {
63 self.cublas_symbols
64 .iter()
65 .any(|symbol| symbol.name == "cublasGemmEx" && symbol.present)
66 && self
67 .cublaslt_symbols
68 .iter()
69 .any(|symbol| symbol.name == "cublasLtMatmul" && symbol.present)
70 }
71}
72
73pub struct CudaLibrarySet {
75 evidence: CudaAbiEvidence,
76 driver: Library,
77 cublas: Library,
78 cublaslt: Library,
79}
80
81impl CudaLibrarySet {
82 fn new(evidence: CudaAbiEvidence, driver: Library, cublas: Library, cublaslt: Library) -> Self {
83 Self {
84 evidence,
85 driver,
86 cublas,
87 cublaslt,
88 }
89 }
90
91 pub fn evidence(&self) -> &CudaAbiEvidence {
93 &self.evidence
94 }
95
96 pub fn handles(&self) -> (&Library, &Library, &Library) {
98 (&self.driver, &self.cublas, &self.cublaslt)
99 }
100}
101
102impl fmt::Debug for CudaLibrarySet {
103 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
104 formatter
105 .debug_struct("CudaLibrarySet")
106 .field("evidence", &self.evidence)
107 .finish_non_exhaustive()
108 }
109}
110
111#[derive(Clone, Debug)]
113pub struct CudaRuntimeProbe {
114 pub runtime: Option<Arc<CudaLibrarySet>>,
116 pub evidence: Option<CudaAbiEvidence>,
118 pub diagnostics: Vec<String>,
120}
121
122impl CudaRuntimeProbe {
123 pub fn fake_present(evidence: CudaAbiEvidence) -> Self {
126 Self {
127 runtime: None,
128 evidence: Some(evidence),
129 diagnostics: Vec::new(),
130 }
131 }
132
133 pub fn is_available(&self) -> bool {
135 self.evidence
136 .as_ref()
137 .is_some_and(CudaAbiEvidence::is_complete)
138 }
139}
140
141#[derive(Clone, Debug, PartialEq, Eq)]
143pub struct CudaLoadError {
144 pub message: String,
146}
147
148impl fmt::Display for CudaLoadError {
149 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
150 formatter.write_str(&self.message)
151 }
152}
153
154impl std::error::Error for CudaLoadError {}
155
156pub trait DynamicCudaLoader {
158 fn discover(&self) -> Result<CudaRuntimeProbe, CudaLoadError>;
160}
161
162#[derive(Clone, Debug, Default)]
164pub struct CudaRuntimeLoader {
165 search_dirs: Vec<PathBuf>,
166}
167
168impl CudaRuntimeLoader {
169 pub fn new() -> Self {
171 Self::default()
172 }
173
174 pub fn with_search_dirs(search_dirs: Vec<PathBuf>) -> Self {
176 Self { search_dirs }
177 }
178}
179
180impl DynamicCudaLoader for CudaRuntimeLoader {
181 fn discover(&self) -> Result<CudaRuntimeProbe, CudaLoadError> {
182 let mut diagnostics = Vec::new();
183 let (driver_name, driver) = self.open_first(DRIVER_NAMES, &mut diagnostics)?;
184 let (cublas_name, cublas) = self.open_first(CUBLAS_NAMES, &mut diagnostics)?;
185 let (cublaslt_name, cublaslt) = self.open_first(CUBLASLT_NAMES, &mut diagnostics)?;
186
187 let driver_symbols = symbol_evidence(&driver, DRIVER_SYMBOLS);
188 let cublas_symbols = symbol_evidence(&cublas, CUBLAS_SYMBOLS);
189 let cublaslt_symbols = symbol_evidence(&cublaslt, CUBLASLT_SYMBOLS);
190 let driver_version = driver_version(&driver).ok();
191 let evidence = CudaAbiEvidence {
192 driver_library: driver_name,
193 cublas_library: cublas_name,
194 cublaslt_library: cublaslt_name,
195 driver_version,
196 driver_symbols,
197 cublas_symbols,
198 cublaslt_symbols,
199 };
200 if !evidence.is_complete() {
201 return Ok(CudaRuntimeProbe {
202 runtime: None,
203 evidence: Some(evidence),
204 diagnostics,
205 });
206 }
207 let runtime = Arc::new(CudaLibrarySet::new(
208 evidence.clone(),
209 driver,
210 cublas,
211 cublaslt,
212 ));
213 Ok(CudaRuntimeProbe {
214 runtime: Some(runtime),
215 evidence: Some(evidence),
216 diagnostics,
217 })
218 }
219}
220
221impl CudaRuntimeLoader {
222 fn open_first(
223 &self,
224 names: &[&str],
225 diagnostics: &mut Vec<String>,
226 ) -> Result<(String, Library), CudaLoadError> {
227 for name in candidate_paths(&self.search_dirs, names) {
228 match open_library(&name) {
229 Ok(library) => return Ok((name.display().to_string(), library)),
230 Err(error) => diagnostics.push(format!("{}: {error}", name.display())),
231 }
232 }
233 Err(CudaLoadError {
234 message: format!("CUDA library was not found; tried {}", names.join(", ")),
235 })
236 }
237}
238
239#[derive(Clone, Debug)]
241pub struct FakeCudaLoader {
242 probe: Result<CudaRuntimeProbe, CudaLoadError>,
243}
244
245impl FakeCudaLoader {
246 pub fn available() -> Self {
248 Self {
249 probe: Ok(CudaRuntimeProbe::fake_present(complete_fake_evidence())),
250 }
251 }
252
253 pub fn incomplete() -> Self {
255 let mut evidence = complete_fake_evidence();
256 if let Some(symbol) = evidence
257 .cublaslt_symbols
258 .iter_mut()
259 .find(|symbol| symbol.name == "cublasLtMatmul")
260 {
261 symbol.present = false;
262 }
263 Self {
264 probe: Ok(CudaRuntimeProbe {
265 runtime: None,
266 evidence: Some(evidence),
267 diagnostics: vec!["missing cublasLtMatmul".to_owned()],
268 }),
269 }
270 }
271
272 pub fn absent() -> Self {
274 Self {
275 probe: Err(CudaLoadError {
276 message: "CUDA runtime absent".to_owned(),
277 }),
278 }
279 }
280}
281
282impl DynamicCudaLoader for FakeCudaLoader {
283 fn discover(&self) -> Result<CudaRuntimeProbe, CudaLoadError> {
284 self.probe.clone()
285 }
286}
287
288pub fn discover_cuda_runtime() -> Result<CudaRuntimeProbe, CudaLoadError> {
290 CudaRuntimeLoader::new().discover()
291}
292
293fn complete_fake_evidence() -> CudaAbiEvidence {
294 CudaAbiEvidence {
295 driver_library: "fake-libcuda".to_owned(),
296 cublas_library: "fake-libcublas".to_owned(),
297 cublaslt_library: "fake-libcublasLt".to_owned(),
298 driver_version: Some(12_000),
299 driver_symbols: DRIVER_SYMBOLS
300 .iter()
301 .map(|name| CudaSymbolEvidence {
302 name: (*name).to_owned(),
303 present: true,
304 })
305 .collect(),
306 cublas_symbols: CUBLAS_SYMBOLS
307 .iter()
308 .map(|name| CudaSymbolEvidence {
309 name: (*name).to_owned(),
310 present: true,
311 })
312 .collect(),
313 cublaslt_symbols: CUBLASLT_SYMBOLS
314 .iter()
315 .map(|name| CudaSymbolEvidence {
316 name: (*name).to_owned(),
317 present: true,
318 })
319 .collect(),
320 }
321}
322
323fn candidate_paths(search_dirs: &[PathBuf], names: &[&str]) -> Vec<PathBuf> {
324 let mut candidates = Vec::new();
325 for directory in search_dirs {
326 for name in names {
327 candidates.push(directory.join(name));
328 }
329 }
330 candidates.extend(names.iter().map(PathBuf::from));
331 candidates
332}
333
334fn symbol_evidence(library: &Library, names: &[&str]) -> Vec<CudaSymbolEvidence> {
335 names
336 .iter()
337 .map(|name| CudaSymbolEvidence {
338 name: (*name).to_owned(),
339 present: symbol_present(library, name),
340 })
341 .collect()
342}
343
344fn open_library(path: &Path) -> Result<Library, libloading::Error> {
345 unsafe { Library::new(path) }
349}
350
351fn symbol_present(library: &Library, name: &str) -> bool {
352 let mut bytes = name.as_bytes().to_vec();
353 bytes.push(0);
354 unsafe { library.get::<*mut c_void>(&bytes).is_ok() }
357}
358
359fn driver_version(library: &Library) -> Result<i32, CudaLoadError> {
360 type CuInit = unsafe extern "C" fn(u32) -> i32;
361 type CuDriverGetVersion = unsafe extern "C" fn(*mut i32) -> i32;
362 unsafe {
366 let cu_init = library
367 .get::<CuInit>(b"cuInit\0")
368 .map_err(|error| CudaLoadError {
369 message: error.to_string(),
370 })?;
371 let get_version = library
372 .get::<CuDriverGetVersion>(b"cuDriverGetVersion\0")
373 .map_err(|error| CudaLoadError {
374 message: error.to_string(),
375 })?;
376 let init_status = cu_init(0);
377 if init_status != 0 {
378 return Err(CudaLoadError {
379 message: format!("cuInit failed with status {init_status}"),
380 });
381 }
382 let mut version = 0;
383 let version_status = get_version(&mut version);
384 if version_status != 0 {
385 return Err(CudaLoadError {
386 message: format!("cuDriverGetVersion failed with status {version_status}"),
387 });
388 }
389 Ok(version)
390 }
391}