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 fn new(
103 evidence: CudaAbiEvidence,
104 driver: Library,
105 runtime: Library,
106 cublas: Library,
107 cublaslt: Library,
108 ) -> Self {
109 Self {
110 evidence,
111 driver,
112 runtime,
113 cublas,
114 cublaslt,
115 }
116 }
117
118 pub fn evidence(&self) -> &CudaAbiEvidence {
120 &self.evidence
121 }
122
123 pub fn handles(&self) -> (&Library, &Library, &Library) {
125 (&self.driver, &self.cublas, &self.cublaslt)
126 }
127
128 pub(crate) fn execution_handles(&self) -> (&Library, &Library) {
129 (&self.runtime, &self.cublas)
130 }
131}
132
133impl fmt::Debug for CudaLibrarySet {
134 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
135 formatter
136 .debug_struct("CudaLibrarySet")
137 .field("evidence", &self.evidence)
138 .finish_non_exhaustive()
139 }
140}
141
142#[derive(Clone, Debug)]
144pub struct CudaRuntimeProbe {
145 pub runtime: Option<Arc<CudaLibrarySet>>,
147 pub evidence: Option<CudaAbiEvidence>,
149 pub diagnostics: Vec<String>,
151}
152
153impl CudaRuntimeProbe {
154 pub fn fake_present(evidence: CudaAbiEvidence) -> Self {
157 Self {
158 runtime: None,
159 evidence: Some(evidence),
160 diagnostics: Vec::new(),
161 }
162 }
163
164 pub fn is_available(&self) -> bool {
166 self.evidence
167 .as_ref()
168 .is_some_and(CudaAbiEvidence::is_complete)
169 }
170}
171
172#[derive(Clone, Debug, PartialEq, Eq)]
174pub struct CudaLoadError {
175 pub message: String,
177}
178
179impl fmt::Display for CudaLoadError {
180 fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
181 formatter.write_str(&self.message)
182 }
183}
184
185impl std::error::Error for CudaLoadError {}
186
187pub trait DynamicCudaLoader {
189 fn discover(&self) -> Result<CudaRuntimeProbe, CudaLoadError>;
191}
192
193#[derive(Clone, Debug)]
195pub struct CudaRuntimeLoader {
196 search_dirs: Vec<PathBuf>,
197 search_system: bool,
198}
199
200impl Default for CudaRuntimeLoader {
201 fn default() -> Self {
202 Self {
203 search_dirs: Vec::new(),
204 search_system: true,
205 }
206 }
207}
208
209impl CudaRuntimeLoader {
210 pub fn new() -> Self {
212 Self::default()
213 }
214
215 pub fn with_search_dirs(search_dirs: Vec<PathBuf>) -> Self {
217 Self {
218 search_dirs,
219 search_system: true,
220 }
221 }
222
223 pub fn with_search_dirs_only(search_dirs: Vec<PathBuf>) -> Self {
226 Self {
227 search_dirs,
228 search_system: false,
229 }
230 }
231}
232
233impl DynamicCudaLoader for CudaRuntimeLoader {
234 fn discover(&self) -> Result<CudaRuntimeProbe, CudaLoadError> {
235 let mut diagnostics = Vec::new();
236 let (driver_name, driver) = self.open_first(DRIVER_NAMES, &mut diagnostics)?;
237 let (runtime_name, runtime) = self.open_first(RUNTIME_NAMES, &mut diagnostics)?;
238 let (cublas_name, cublas) = self.open_first(CUBLAS_NAMES, &mut diagnostics)?;
239 let (cublaslt_name, cublaslt) = self.open_first(CUBLASLT_NAMES, &mut diagnostics)?;
240
241 let driver_symbols = symbol_evidence(&driver, DRIVER_SYMBOLS);
242 let runtime_symbols = symbol_evidence(&runtime, RUNTIME_SYMBOLS);
243 let cublas_symbols = symbol_evidence(&cublas, CUBLAS_SYMBOLS);
244 let cublaslt_symbols = symbol_evidence(&cublaslt, CUBLASLT_SYMBOLS);
245 let driver_version = driver_version(&driver).ok();
246 let evidence = CudaAbiEvidence {
247 driver_library: driver_name,
248 runtime_library: runtime_name,
249 cublas_library: cublas_name,
250 cublaslt_library: cublaslt_name,
251 driver_version,
252 driver_symbols,
253 runtime_symbols,
254 cublas_symbols,
255 cublaslt_symbols,
256 };
257 if !evidence.is_complete() {
258 return Ok(CudaRuntimeProbe {
259 runtime: None,
260 evidence: Some(evidence),
261 diagnostics,
262 });
263 }
264 let runtime = Arc::new(CudaLibrarySet::new(
265 evidence.clone(),
266 driver,
267 runtime,
268 cublas,
269 cublaslt,
270 ));
271 Ok(CudaRuntimeProbe {
272 runtime: Some(runtime),
273 evidence: Some(evidence),
274 diagnostics,
275 })
276 }
277}
278
279impl CudaRuntimeLoader {
280 fn open_first(
281 &self,
282 names: &[&str],
283 diagnostics: &mut Vec<String>,
284 ) -> Result<(String, Library), CudaLoadError> {
285 for name in candidate_paths(&self.search_dirs, names, self.search_system) {
286 match open_library(&name) {
287 Ok(library) => return Ok((name.display().to_string(), library)),
288 Err(error) => diagnostics.push(format!("{}: {error}", name.display())),
289 }
290 }
291 Err(CudaLoadError {
292 message: format!("CUDA library was not found; tried {}", names.join(", ")),
293 })
294 }
295}
296
297#[derive(Clone, Debug)]
299pub struct FakeCudaLoader {
300 probe: Result<CudaRuntimeProbe, CudaLoadError>,
301}
302
303impl FakeCudaLoader {
304 pub fn available() -> Self {
306 Self {
307 probe: Ok(CudaRuntimeProbe::fake_present(complete_fake_evidence())),
308 }
309 }
310
311 pub fn incomplete() -> Self {
313 let mut evidence = complete_fake_evidence();
314 if let Some(symbol) = evidence
315 .cublaslt_symbols
316 .iter_mut()
317 .find(|symbol| symbol.name == "cublasLtMatmul")
318 {
319 symbol.present = false;
320 }
321 Self {
322 probe: Ok(CudaRuntimeProbe {
323 runtime: None,
324 evidence: Some(evidence),
325 diagnostics: vec!["missing cublasLtMatmul".to_owned()],
326 }),
327 }
328 }
329
330 pub fn absent() -> Self {
332 Self {
333 probe: Err(CudaLoadError {
334 message: "CUDA runtime absent".to_owned(),
335 }),
336 }
337 }
338}
339
340impl DynamicCudaLoader for FakeCudaLoader {
341 fn discover(&self) -> Result<CudaRuntimeProbe, CudaLoadError> {
342 self.probe.clone()
343 }
344}
345
346pub fn discover_cuda_runtime() -> Result<CudaRuntimeProbe, CudaLoadError> {
348 CudaRuntimeLoader::new().discover()
349}
350
351fn complete_fake_evidence() -> CudaAbiEvidence {
352 CudaAbiEvidence {
353 driver_library: "fake-libcuda".to_owned(),
354 runtime_library: "fake-libcudart".to_owned(),
355 cublas_library: "fake-libcublas".to_owned(),
356 cublaslt_library: "fake-libcublasLt".to_owned(),
357 driver_version: Some(12_000),
358 driver_symbols: DRIVER_SYMBOLS
359 .iter()
360 .map(|name| CudaSymbolEvidence {
361 name: (*name).to_owned(),
362 present: true,
363 })
364 .collect(),
365 runtime_symbols: RUNTIME_SYMBOLS
366 .iter()
367 .map(|name| CudaSymbolEvidence {
368 name: (*name).to_owned(),
369 present: true,
370 })
371 .collect(),
372 cublas_symbols: CUBLAS_SYMBOLS
373 .iter()
374 .map(|name| CudaSymbolEvidence {
375 name: (*name).to_owned(),
376 present: true,
377 })
378 .collect(),
379 cublaslt_symbols: CUBLASLT_SYMBOLS
380 .iter()
381 .map(|name| CudaSymbolEvidence {
382 name: (*name).to_owned(),
383 present: true,
384 })
385 .collect(),
386 }
387}
388
389fn candidate_paths(search_dirs: &[PathBuf], names: &[&str], search_system: bool) -> Vec<PathBuf> {
390 let mut candidates = Vec::new();
391 for directory in search_dirs {
392 for name in names {
393 candidates.push(directory.join(name));
394 }
395 }
396 if search_system {
397 candidates.extend(names.iter().map(PathBuf::from));
398 }
399 candidates
400}
401
402fn symbol_evidence(library: &Library, names: &[&str]) -> Vec<CudaSymbolEvidence> {
403 names
404 .iter()
405 .map(|name| CudaSymbolEvidence {
406 name: (*name).to_owned(),
407 present: symbol_present(library, name),
408 })
409 .collect()
410}
411
412fn open_library(path: &Path) -> Result<Library, libloading::Error> {
413 unsafe { Library::new(path) }
417}
418
419fn symbol_present(library: &Library, name: &str) -> bool {
420 let mut bytes = name.as_bytes().to_vec();
421 bytes.push(0);
422 unsafe { library.get::<*mut c_void>(&bytes).is_ok() }
425}
426
427fn driver_version(library: &Library) -> Result<i32, CudaLoadError> {
428 type CuInit = unsafe extern "C" fn(u32) -> i32;
429 type CuDriverGetVersion = unsafe extern "C" fn(*mut i32) -> i32;
430 unsafe {
434 let cu_init = library
435 .get::<CuInit>(b"cuInit\0")
436 .map_err(|error| CudaLoadError {
437 message: error.to_string(),
438 })?;
439 let get_version = library
440 .get::<CuDriverGetVersion>(b"cuDriverGetVersion\0")
441 .map_err(|error| CudaLoadError {
442 message: error.to_string(),
443 })?;
444 let init_status = cu_init(0);
445 if init_status != 0 {
446 return Err(CudaLoadError {
447 message: format!("cuInit failed with status {init_status}"),
448 });
449 }
450 let mut version = 0;
451 let version_status = get_version(&mut version);
452 if version_status != 0 {
453 return Err(CudaLoadError {
454 message: format!("cuDriverGetVersion failed with status {version_status}"),
455 });
456 }
457 Ok(version)
458 }
459}