use cubecl::prelude::*;
use crate::accelerate::Accelerator;
use crate::device::Device;
pub fn sniff_best_accelerator(enable: &[Accelerator], device: &Device) -> Option<Accelerator> {
for accelerator in enable {
let is_enabled = match accelerator {
#[cfg(feature = "cuda")]
Accelerator::Cuda => match device.to_cuda() {
Ok(dev) => probe_runtime::<cubecl::cuda::CudaRuntime>("CUDA", &dev),
Err(_) => false,
},
#[cfg(feature = "rocm")]
Accelerator::Rocm => match device.to_amd() {
Ok(dev) => probe_runtime::<cubecl::hip::HipRuntime>("ROCM", &dev),
Err(_) => false,
},
#[cfg(feature = "vulkan")]
Accelerator::Vulkan => match device.to_wgpu() {
Ok(dev) => probe_runtime::<cubecl::wgpu::WgpuRuntime>("VULKAN", &dev),
Err(_) => false,
},
#[cfg(feature = "metal")]
Accelerator::Metal => match device.to_wgpu() {
Ok(dev) => probe_runtime::<cubecl::wgpu::WgpuRuntime>("METAL", &dev),
Err(_) => false,
},
#[cfg(docsrs)]
#[allow(unreachable_patterns)]
_ => unreachable!(),
};
if is_enabled {
return Some(*accelerator);
}
}
None
}
fn probe_runtime<R: Runtime>(name: &'static str, device: &R::Device) -> bool {
let client = R::client(device);
match cubecl::future::block_on(client.sync()) {
Ok(()) => true,
Err(err) => {
tracing::debug!(err = ?err, "could not use {name} runtime");
false
},
}
}