use std::{
ffi::c_void,
fmt,
path::{Path, PathBuf},
sync::Arc,
};
use libloading::Library;
const DRIVER_NAMES: &[&str] = &["libcuda.so.1", "libcuda.so", "nvcuda.dll"];
const CUBLAS_NAMES: &[&str] = &["libcublas.so.13", "libcublas.so.12", "libcublas.so"];
const CUBLASLT_NAMES: &[&str] = &["libcublasLt.so.13", "libcublasLt.so.12", "libcublasLt.so"];
const DRIVER_SYMBOLS: &[&str] = &["cuInit", "cuDriverGetVersion"];
const CUBLAS_SYMBOLS: &[&str] = &[
"cublasCreate_v2",
"cublasDestroy_v2",
"cublasSgemm_v2",
"cublasGemmEx",
];
const CUBLASLT_SYMBOLS: &[&str] = &["cublasLtCreate", "cublasLtDestroy", "cublasLtMatmul"];
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct CudaSymbolEvidence {
pub name: String,
pub present: bool,
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct CudaAbiEvidence {
pub driver_library: String,
pub cublas_library: String,
pub cublaslt_library: String,
pub driver_version: Option<i32>,
pub driver_symbols: Vec<CudaSymbolEvidence>,
pub cublas_symbols: Vec<CudaSymbolEvidence>,
pub cublaslt_symbols: Vec<CudaSymbolEvidence>,
}
impl CudaAbiEvidence {
pub fn is_complete(&self) -> bool {
self.driver_symbols.iter().all(|symbol| symbol.present)
&& self.cublas_symbols.iter().all(|symbol| symbol.present)
&& self.cublaslt_symbols.iter().all(|symbol| symbol.present)
}
pub fn supports_half_matmul(&self) -> bool {
self.cublas_symbols
.iter()
.any(|symbol| symbol.name == "cublasGemmEx" && symbol.present)
&& self
.cublaslt_symbols
.iter()
.any(|symbol| symbol.name == "cublasLtMatmul" && symbol.present)
}
}
pub struct CudaLibrarySet {
evidence: CudaAbiEvidence,
driver: Library,
cublas: Library,
cublaslt: Library,
}
impl CudaLibrarySet {
fn new(evidence: CudaAbiEvidence, driver: Library, cublas: Library, cublaslt: Library) -> Self {
Self {
evidence,
driver,
cublas,
cublaslt,
}
}
pub fn evidence(&self) -> &CudaAbiEvidence {
&self.evidence
}
pub fn handles(&self) -> (&Library, &Library, &Library) {
(&self.driver, &self.cublas, &self.cublaslt)
}
}
impl fmt::Debug for CudaLibrarySet {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter
.debug_struct("CudaLibrarySet")
.field("evidence", &self.evidence)
.finish_non_exhaustive()
}
}
#[derive(Clone, Debug)]
pub struct CudaRuntimeProbe {
pub runtime: Option<Arc<CudaLibrarySet>>,
pub evidence: Option<CudaAbiEvidence>,
pub diagnostics: Vec<String>,
}
impl CudaRuntimeProbe {
pub fn fake_present(evidence: CudaAbiEvidence) -> Self {
Self {
runtime: None,
evidence: Some(evidence),
diagnostics: Vec::new(),
}
}
pub fn is_available(&self) -> bool {
self.evidence
.as_ref()
.is_some_and(CudaAbiEvidence::is_complete)
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct CudaLoadError {
pub message: String,
}
impl fmt::Display for CudaLoadError {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str(&self.message)
}
}
impl std::error::Error for CudaLoadError {}
pub trait DynamicCudaLoader {
fn discover(&self) -> Result<CudaRuntimeProbe, CudaLoadError>;
}
#[derive(Clone, Debug, Default)]
pub struct CudaRuntimeLoader {
search_dirs: Vec<PathBuf>,
}
impl CudaRuntimeLoader {
pub fn new() -> Self {
Self::default()
}
pub fn with_search_dirs(search_dirs: Vec<PathBuf>) -> Self {
Self { search_dirs }
}
}
impl DynamicCudaLoader for CudaRuntimeLoader {
fn discover(&self) -> Result<CudaRuntimeProbe, CudaLoadError> {
let mut diagnostics = Vec::new();
let (driver_name, driver) = self.open_first(DRIVER_NAMES, &mut diagnostics)?;
let (cublas_name, cublas) = self.open_first(CUBLAS_NAMES, &mut diagnostics)?;
let (cublaslt_name, cublaslt) = self.open_first(CUBLASLT_NAMES, &mut diagnostics)?;
let driver_symbols = symbol_evidence(&driver, DRIVER_SYMBOLS);
let cublas_symbols = symbol_evidence(&cublas, CUBLAS_SYMBOLS);
let cublaslt_symbols = symbol_evidence(&cublaslt, CUBLASLT_SYMBOLS);
let driver_version = driver_version(&driver).ok();
let evidence = CudaAbiEvidence {
driver_library: driver_name,
cublas_library: cublas_name,
cublaslt_library: cublaslt_name,
driver_version,
driver_symbols,
cublas_symbols,
cublaslt_symbols,
};
if !evidence.is_complete() {
return Ok(CudaRuntimeProbe {
runtime: None,
evidence: Some(evidence),
diagnostics,
});
}
let runtime = Arc::new(CudaLibrarySet::new(
evidence.clone(),
driver,
cublas,
cublaslt,
));
Ok(CudaRuntimeProbe {
runtime: Some(runtime),
evidence: Some(evidence),
diagnostics,
})
}
}
impl CudaRuntimeLoader {
fn open_first(
&self,
names: &[&str],
diagnostics: &mut Vec<String>,
) -> Result<(String, Library), CudaLoadError> {
for name in candidate_paths(&self.search_dirs, names) {
match open_library(&name) {
Ok(library) => return Ok((name.display().to_string(), library)),
Err(error) => diagnostics.push(format!("{}: {error}", name.display())),
}
}
Err(CudaLoadError {
message: format!("CUDA library was not found; tried {}", names.join(", ")),
})
}
}
#[derive(Clone, Debug)]
pub struct FakeCudaLoader {
probe: Result<CudaRuntimeProbe, CudaLoadError>,
}
impl FakeCudaLoader {
pub fn available() -> Self {
Self {
probe: Ok(CudaRuntimeProbe::fake_present(complete_fake_evidence())),
}
}
pub fn incomplete() -> Self {
let mut evidence = complete_fake_evidence();
if let Some(symbol) = evidence
.cublaslt_symbols
.iter_mut()
.find(|symbol| symbol.name == "cublasLtMatmul")
{
symbol.present = false;
}
Self {
probe: Ok(CudaRuntimeProbe {
runtime: None,
evidence: Some(evidence),
diagnostics: vec!["missing cublasLtMatmul".to_owned()],
}),
}
}
pub fn absent() -> Self {
Self {
probe: Err(CudaLoadError {
message: "CUDA runtime absent".to_owned(),
}),
}
}
}
impl DynamicCudaLoader for FakeCudaLoader {
fn discover(&self) -> Result<CudaRuntimeProbe, CudaLoadError> {
self.probe.clone()
}
}
pub fn discover_cuda_runtime() -> Result<CudaRuntimeProbe, CudaLoadError> {
CudaRuntimeLoader::new().discover()
}
fn complete_fake_evidence() -> CudaAbiEvidence {
CudaAbiEvidence {
driver_library: "fake-libcuda".to_owned(),
cublas_library: "fake-libcublas".to_owned(),
cublaslt_library: "fake-libcublasLt".to_owned(),
driver_version: Some(12_000),
driver_symbols: DRIVER_SYMBOLS
.iter()
.map(|name| CudaSymbolEvidence {
name: (*name).to_owned(),
present: true,
})
.collect(),
cublas_symbols: CUBLAS_SYMBOLS
.iter()
.map(|name| CudaSymbolEvidence {
name: (*name).to_owned(),
present: true,
})
.collect(),
cublaslt_symbols: CUBLASLT_SYMBOLS
.iter()
.map(|name| CudaSymbolEvidence {
name: (*name).to_owned(),
present: true,
})
.collect(),
}
}
fn candidate_paths(search_dirs: &[PathBuf], names: &[&str]) -> Vec<PathBuf> {
let mut candidates = Vec::new();
for directory in search_dirs {
for name in names {
candidates.push(directory.join(name));
}
}
candidates.extend(names.iter().map(PathBuf::from));
candidates
}
fn symbol_evidence(library: &Library, names: &[&str]) -> Vec<CudaSymbolEvidence> {
names
.iter()
.map(|name| CudaSymbolEvidence {
name: (*name).to_owned(),
present: symbol_present(library, name),
})
.collect()
}
fn open_library(path: &Path) -> Result<Library, libloading::Error> {
unsafe { Library::new(path) }
}
fn symbol_present(library: &Library, name: &str) -> bool {
let mut bytes = name.as_bytes().to_vec();
bytes.push(0);
unsafe { library.get::<*mut c_void>(&bytes).is_ok() }
}
fn driver_version(library: &Library) -> Result<i32, CudaLoadError> {
type CuInit = unsafe extern "C" fn(u32) -> i32;
type CuDriverGetVersion = unsafe extern "C" fn(*mut i32) -> i32;
unsafe {
let cu_init = library
.get::<CuInit>(b"cuInit\0")
.map_err(|error| CudaLoadError {
message: error.to_string(),
})?;
let get_version = library
.get::<CuDriverGetVersion>(b"cuDriverGetVersion\0")
.map_err(|error| CudaLoadError {
message: error.to_string(),
})?;
let init_status = cu_init(0);
if init_status != 0 {
return Err(CudaLoadError {
message: format!("cuInit failed with status {init_status}"),
});
}
let mut version = 0;
let version_status = get_version(&mut version);
if version_status != 0 {
return Err(CudaLoadError {
message: format!("cuDriverGetVersion failed with status {version_status}"),
});
}
Ok(version)
}
}