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 RUNTIME_NAMES: &[&str] = &[
14 "libcudart.so.13",
15 "libcudart.so.12",
16 "libcudart.so.11.0",
17 "libcudart.so",
18 "cudart64_130.dll",
19 "cudart64_12.dll",
20];
21const CUBLAS_NAMES: &[&str] = &["libcublas.so.13", "libcublas.so.12", "libcublas.so"];
22const CUBLASLT_NAMES: &[&str] = &["libcublasLt.so.13", "libcublasLt.so.12", "libcublasLt.so"];
23
24const DRIVER_SYMBOLS: &[&str] = &["cuInit", "cuDriverGetVersion"];
25const RUNTIME_SYMBOLS: &[&str] = &[
26 "cudaMalloc",
27 "cudaFree",
28 "cudaMemcpy",
29 "cudaDeviceSynchronize",
30];
31const CUBLAS_SYMBOLS: &[&str] = &[
32 "cublasCreate_v2",
33 "cublasDestroy_v2",
34 "cublasSgemm_v2",
35 "cublasGemmEx",
36];
37const CUBLASLT_SYMBOLS: &[&str] = &["cublasLtCreate", "cublasLtDestroy", "cublasLtMatmul"];
38
39#[derive(Clone, Debug, PartialEq, Eq)]
41pub struct CudaSymbolEvidence {
42 pub name: String,
44 pub present: bool,
46}
47
48#[derive(Clone, Debug, PartialEq, Eq)]
50pub struct CudaAbiEvidence {
51 pub driver_library: String,
53 pub runtime_library: String,
55 pub cublas_library: String,
57 pub cublaslt_library: String,
59 pub driver_version: Option<i32>,
61 pub driver_symbols: Vec<CudaSymbolEvidence>,
63 pub runtime_symbols: Vec<CudaSymbolEvidence>,
65 pub cublas_symbols: Vec<CudaSymbolEvidence>,
67 pub cublaslt_symbols: Vec<CudaSymbolEvidence>,
69}
70
71impl CudaAbiEvidence {
72 pub fn is_complete(&self) -> bool {
74 self.driver_symbols.iter().all(|symbol| symbol.present)
75 && self.runtime_symbols.iter().all(|symbol| symbol.present)
76 && self.cublas_symbols.iter().all(|symbol| symbol.present)
77 && self.cublaslt_symbols.iter().all(|symbol| symbol.present)
78 }
79
80 pub fn supports_half_matmul(&self) -> bool {
82 self.cublas_symbols
83 .iter()
84 .any(|symbol| symbol.name == "cublasGemmEx" && symbol.present)
85 && self
86 .cublaslt_symbols
87 .iter()
88 .any(|symbol| symbol.name == "cublasLtMatmul" && symbol.present)
89 }
90}
91
92pub struct CudaLibrarySet {
94 evidence: CudaAbiEvidence,
95 driver: Library,
96 runtime: Library,
97 cublas: Library,
98 cublaslt: Library,
99}
100
101impl CudaLibrarySet {
102 pub fn new(
104 evidence: CudaAbiEvidence,
105 driver: Library,
106 runtime: Library,
107 cublas: Library,
108 cublaslt: Library,
109 ) -> Self {
110 Self {
111 evidence,
112 driver,
113 runtime,
114 cublas,
115 cublaslt,
116 }
117 }
118
119 pub fn evidence(&self) -> &CudaAbiEvidence {
121 &self.evidence
122 }
123
124 pub fn handles(&self) -> (&Library, &Library, &Library) {
126 (&self.driver, &self.cublas, &self.cublaslt)
127 }
128
129 pub(crate) fn execution_handles(&self) -> (&Library, &Library) {
130 (&self.runtime, &self.cublas)
131 }
132}
133
134impl fmt::Debug for CudaLibrarySet {
135 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
136 formatter
137 .debug_struct("CudaLibrarySet")
138 .field("evidence", &self.evidence)
139 .finish_non_exhaustive()
140 }
141}
142
143#[derive(Clone, Debug)]
145pub struct CudaRuntimeProbe {
146 pub runtime: Option<Arc<CudaLibrarySet>>,
148 pub evidence: Option<CudaAbiEvidence>,
150 pub diagnostics: Vec<String>,
152}
153
154impl CudaRuntimeProbe {
155 pub fn fake_present(evidence: CudaAbiEvidence) -> Self {
158 Self {
159 runtime: None,
160 evidence: Some(evidence),
161 diagnostics: Vec::new(),
162 }
163 }
164
165 pub fn is_available(&self) -> bool {
167 self.evidence
168 .as_ref()
169 .is_some_and(CudaAbiEvidence::is_complete)
170 }
171}
172
173#[derive(Clone, Debug, PartialEq, Eq)]
175pub struct CudaLoadError {
176 pub message: String,
178}
179
180impl fmt::Display for CudaLoadError {
181 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
182 formatter.write_str(&self.message)
183 }
184}
185
186impl std::error::Error for CudaLoadError {}
187
188pub trait DynamicCudaLoader {
190 fn discover(&self) -> Result<CudaRuntimeProbe, CudaLoadError>;
192}
193
194pub trait CudaProbePort {
196 fn probe_cuda(&self) -> Result<CudaRuntimeProbe, CudaLoadError>;
198}
199
200impl<T: DynamicCudaLoader + ?Sized> CudaProbePort for T {
201 fn probe_cuda(&self) -> Result<CudaRuntimeProbe, CudaLoadError> {
202 self.discover()
203 }
204}
205
206#[derive(Clone, Debug)]
208pub struct CudaRuntimeLoader {
209 search_dirs: Vec<PathBuf>,
210 search_system: bool,
211}
212
213impl Default for CudaRuntimeLoader {
214 fn default() -> Self {
215 Self {
216 search_dirs: Vec::new(),
217 search_system: true,
218 }
219 }
220}
221
222impl CudaRuntimeLoader {
223 pub fn new() -> Self {
225 Self::default()
226 }
227
228 pub fn with_search_dirs(search_dirs: Vec<PathBuf>) -> Self {
230 Self {
231 search_dirs,
232 search_system: true,
233 }
234 }
235
236 pub fn with_search_dirs_only(search_dirs: Vec<PathBuf>) -> Self {
239 Self {
240 search_dirs,
241 search_system: false,
242 }
243 }
244}
245
246impl DynamicCudaLoader for CudaRuntimeLoader {
247 fn discover(&self) -> Result<CudaRuntimeProbe, CudaLoadError> {
248 let mut diagnostics = Vec::new();
249 let (driver_name, driver) = self.open_first(DRIVER_NAMES, &mut diagnostics)?;
250 let (runtime_name, runtime) = self.open_first(RUNTIME_NAMES, &mut diagnostics)?;
251 let (cublas_name, cublas) = self.open_first(CUBLAS_NAMES, &mut diagnostics)?;
252 let (cublaslt_name, cublaslt) = self.open_first(CUBLASLT_NAMES, &mut diagnostics)?;
253
254 let driver_symbols = symbol_evidence(&driver, DRIVER_SYMBOLS);
255 let runtime_symbols = symbol_evidence(&runtime, RUNTIME_SYMBOLS);
256 let cublas_symbols = symbol_evidence(&cublas, CUBLAS_SYMBOLS);
257 let cublaslt_symbols = symbol_evidence(&cublaslt, CUBLASLT_SYMBOLS);
258 let driver_version = driver_version(&driver).ok();
259 let evidence = CudaAbiEvidence {
260 driver_library: driver_name,
261 runtime_library: runtime_name,
262 cublas_library: cublas_name,
263 cublaslt_library: cublaslt_name,
264 driver_version,
265 driver_symbols,
266 runtime_symbols,
267 cublas_symbols,
268 cublaslt_symbols,
269 };
270 if !evidence.is_complete() {
271 return Ok(CudaRuntimeProbe {
272 runtime: None,
273 evidence: Some(evidence),
274 diagnostics,
275 });
276 }
277 let runtime = Arc::new(CudaLibrarySet::new(
278 evidence.clone(),
279 driver,
280 runtime,
281 cublas,
282 cublaslt,
283 ));
284 Ok(CudaRuntimeProbe {
285 runtime: Some(runtime),
286 evidence: Some(evidence),
287 diagnostics,
288 })
289 }
290}
291
292impl CudaRuntimeLoader {
293 fn open_first(
294 &self,
295 names: &[&str],
296 diagnostics: &mut Vec<String>,
297 ) -> Result<(String, Library), CudaLoadError> {
298 for name in candidate_paths(&self.search_dirs, names, self.search_system) {
299 match open_library(&name) {
300 Ok(library) => return Ok((name.display().to_string(), library)),
301 Err(error) => diagnostics.push(format!("{}: {error}", name.display())),
302 }
303 }
304 Err(CudaLoadError {
305 message: format!("CUDA library was not found; tried {}", names.join(", ")),
306 })
307 }
308}
309
310#[derive(Clone, Debug)]
312pub struct FakeCudaLoader {
313 probe: Result<CudaRuntimeProbe, CudaLoadError>,
314}
315
316impl FakeCudaLoader {
317 pub fn available() -> Self {
319 Self {
320 probe: Ok(CudaRuntimeProbe::fake_present(complete_fake_evidence())),
321 }
322 }
323
324 pub fn incomplete() -> Self {
326 let mut evidence = complete_fake_evidence();
327 if let Some(symbol) = evidence
328 .cublaslt_symbols
329 .iter_mut()
330 .find(|symbol| symbol.name == "cublasLtMatmul")
331 {
332 symbol.present = false;
333 }
334 Self {
335 probe: Ok(CudaRuntimeProbe {
336 runtime: None,
337 evidence: Some(evidence),
338 diagnostics: vec!["missing cublasLtMatmul".to_owned()],
339 }),
340 }
341 }
342
343 pub fn absent() -> Self {
345 Self {
346 probe: Err(CudaLoadError {
347 message: "CUDA runtime absent".to_owned(),
348 }),
349 }
350 }
351}
352
353impl DynamicCudaLoader for FakeCudaLoader {
354 fn discover(&self) -> Result<CudaRuntimeProbe, CudaLoadError> {
355 self.probe.clone()
356 }
357}
358
359pub fn discover_cuda_runtime() -> Result<CudaRuntimeProbe, CudaLoadError> {
361 CudaRuntimeLoader::new().discover()
362}
363
364fn complete_fake_evidence() -> CudaAbiEvidence {
365 CudaAbiEvidence {
366 driver_library: "fake-libcuda".to_owned(),
367 runtime_library: "fake-libcudart".to_owned(),
368 cublas_library: "fake-libcublas".to_owned(),
369 cublaslt_library: "fake-libcublasLt".to_owned(),
370 driver_version: Some(12_000),
371 driver_symbols: DRIVER_SYMBOLS
372 .iter()
373 .map(|name| CudaSymbolEvidence {
374 name: (*name).to_owned(),
375 present: true,
376 })
377 .collect(),
378 runtime_symbols: RUNTIME_SYMBOLS
379 .iter()
380 .map(|name| CudaSymbolEvidence {
381 name: (*name).to_owned(),
382 present: true,
383 })
384 .collect(),
385 cublas_symbols: CUBLAS_SYMBOLS
386 .iter()
387 .map(|name| CudaSymbolEvidence {
388 name: (*name).to_owned(),
389 present: true,
390 })
391 .collect(),
392 cublaslt_symbols: CUBLASLT_SYMBOLS
393 .iter()
394 .map(|name| CudaSymbolEvidence {
395 name: (*name).to_owned(),
396 present: true,
397 })
398 .collect(),
399 }
400}
401
402fn candidate_paths(search_dirs: &[PathBuf], names: &[&str], search_system: bool) -> Vec<PathBuf> {
403 let mut candidates = Vec::new();
404 for directory in search_dirs {
405 for name in names {
406 candidates.push(directory.join(name));
407 }
408 }
409 if search_system {
410 candidates.extend(names.iter().map(PathBuf::from));
411 }
412 candidates
413}
414
415fn symbol_evidence(library: &Library, names: &[&str]) -> Vec<CudaSymbolEvidence> {
416 names
417 .iter()
418 .map(|name| CudaSymbolEvidence {
419 name: (*name).to_owned(),
420 present: symbol_present(library, name),
421 })
422 .collect()
423}
424
425fn open_library(path: &Path) -> Result<Library, libloading::Error> {
426 unsafe { Library::new(path) }
430}
431
432fn symbol_present(library: &Library, name: &str) -> bool {
433 let mut bytes = name.as_bytes().to_vec();
434 bytes.push(0);
435 unsafe { library.get::<*mut c_void>(&bytes).is_ok() }
438}
439
440fn driver_version(library: &Library) -> Result<i32, CudaLoadError> {
441 type CuInit = unsafe extern "C" fn(u32) -> i32;
442 type CuDriverGetVersion = unsafe extern "C" fn(*mut i32) -> i32;
443 unsafe {
447 let cu_init = library
448 .get::<CuInit>(b"cuInit\0")
449 .map_err(|error| CudaLoadError {
450 message: error.to_string(),
451 })?;
452 let get_version = library
453 .get::<CuDriverGetVersion>(b"cuDriverGetVersion\0")
454 .map_err(|error| CudaLoadError {
455 message: error.to_string(),
456 })?;
457 let init_status = cu_init(0);
458 if init_status != 0 {
459 return Err(CudaLoadError {
460 message: format!("cuInit failed with status {init_status}"),
461 });
462 }
463 let mut version = 0;
464 let version_status = get_version(&mut version);
465 if version_status != 0 {
466 return Err(CudaLoadError {
467 message: format!("cuDriverGetVersion failed with status {version_status}"),
468 });
469 }
470 Ok(version)
471 }
472}