use cubecl::device::DeviceId;
use cubecl::prelude::*;
use crate::accelerate::Accelerator;
use crate::device::Device;
use crate::probe::open_client;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct BackendDevices {
pub accelerator: Accelerator,
pub available: bool,
pub devices: Vec<Device>,
}
pub fn enumerate_devices(enable: &[Accelerator]) -> Vec<BackendDevices> {
enable
.iter()
.map(|accelerator| match accelerator {
#[cfg(feature = "cuda")]
Accelerator::Cuda => match Device::Default.to_cuda() {
Ok(dev) => query_runtime::<cubecl::cuda::CudaRuntime>(*accelerator, &dev),
Err(_) => unavailable(*accelerator),
},
#[cfg(feature = "rocm")]
Accelerator::Rocm => match Device::Default.to_amd() {
Ok(dev) => query_runtime::<cubecl::hip::HipRuntime>(*accelerator, &dev),
Err(_) => unavailable(*accelerator),
},
#[cfg(feature = "vulkan")]
Accelerator::Vulkan => match Device::Default.to_wgpu() {
Ok(dev) => query_runtime::<cubecl::wgpu::WgpuRuntime>(*accelerator, &dev),
Err(_) => unavailable(*accelerator),
},
#[cfg(feature = "metal")]
Accelerator::Metal => match Device::Default.to_wgpu() {
Ok(dev) => query_runtime::<cubecl::wgpu::WgpuRuntime>(*accelerator, &dev),
Err(_) => unavailable(*accelerator),
},
#[cfg(docsrs)]
#[expect(
unreachable_patterns,
reason = "the arm only keeps the match exhaustive on docs.rs"
)]
_ => unreachable!(),
})
.collect()
}
fn unavailable(accelerator: Accelerator) -> BackendDevices {
BackendDevices {
accelerator,
available: false,
devices: Vec::new(),
}
}
fn query_runtime<R: Runtime>(accelerator: Accelerator, device: &R::Device) -> BackendDevices {
let Some(client) = open_client::<R>(accelerator, device) else {
return unavailable(accelerator);
};
let mut devices: Vec<Device> = Vec::new();
for type_id in 0..=3 {
for id in client.enumerate_devices(type_id) {
if let Some(device) = to_device(id)
&& !devices.contains(&device)
{
devices.push(device);
}
}
}
BackendDevices {
accelerator,
available: true,
devices,
}
}
fn to_device(id: DeviceId) -> Option<Device> {
let index = id.index_id as usize;
match id.type_id {
0 => Some(Device::Discrete { index }),
1 => Some(Device::Integrated { index }),
2 => Some(Device::Virtual { index }),
3 => Some(Device::Cpu),
_ => None,
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn device_kinds_map_from_type_ids() {
assert_eq!(
to_device(DeviceId::new(0, 1)),
Some(Device::Discrete { index: 1 }),
);
assert_eq!(
to_device(DeviceId::new(1, 0)),
Some(Device::Integrated { index: 0 }),
);
assert_eq!(to_device(DeviceId::new(2, 2)), Some(Device::Virtual { index: 2 }),);
assert_eq!(to_device(DeviceId::new(3, 0)), Some(Device::Cpu));
}
#[test]
fn unknown_type_ids_are_skipped() {
assert_eq!(to_device(DeviceId::new(4, 0)), None);
}
#[test]
fn no_backends_lists_nothing() {
assert!(enumerate_devices(&[]).is_empty());
}
#[cfg(feature = "vulkan")]
#[test]
fn vulkan_reports_at_least_one_device() {
let reported = enumerate_devices(&[Accelerator::Vulkan]);
assert_eq!(reported.len(), 1);
let vulkan = &reported[0];
assert_eq!(vulkan.accelerator, Accelerator::Vulkan);
assert!(vulkan.available, "the vulkan backend did not start");
assert!(
!vulkan.devices.is_empty(),
"the vulkan backend started but listed no devices",
);
let mut unique = vulkan.devices.clone();
unique.dedup();
assert_eq!(
unique.len(),
vulkan.devices.len(),
"a device was listed more than once: {:?}",
vulkan.devices,
);
}
}