use cubecl_common::device::{Device, DeviceId};
#[derive(Clone, Copy, Debug, Hash, PartialEq, Eq, Default)]
pub enum WgpuBackend {
#[default]
Auto,
Vulkan,
Metal,
Dx12,
Gl,
WebGpu,
}
impl WgpuBackend {
const SHIFT: u32 = 13;
const MASK: u16 = 0x7;
fn from_bits(bits: u16) -> Self {
match bits {
1 => Self::Vulkan,
2 => Self::Metal,
3 => Self::Dx12,
4 => Self::Gl,
5 => Self::WebGpu,
_ => Self::Auto,
}
}
}
#[derive(Clone, Debug, Hash, PartialEq, Eq, Default)]
pub enum WgpuDeviceKind {
DiscreteGpu(usize),
IntegratedGpu(usize),
VirtualGpu(usize),
Cpu,
#[default]
DefaultDevice,
Existing(u32),
Other(usize),
}
impl WgpuDeviceKind {
pub const MAX_INDEX: usize = (1 << WgpuBackend::SHIFT) - 1;
fn type_id(&self) -> u16 {
match *self {
Self::DiscreteGpu(_) => 0,
Self::IntegratedGpu(_) => 1,
Self::VirtualGpu(_) => 2,
Self::Cpu => 3,
Self::DefaultDevice => 4,
Self::Existing(_) => 5,
Self::Other(_) => 6,
}
}
fn index(&self) -> u16 {
match *self {
Self::DiscreteGpu(index)
| Self::IntegratedGpu(index)
| Self::VirtualGpu(index)
| Self::Other(index) => index.min(Self::MAX_INDEX) as u16,
Self::Cpu | Self::DefaultDevice => 0,
Self::Existing(id) => id as u16,
}
}
}
#[derive(Clone, Debug, Hash, PartialEq, Eq, Default)]
pub struct WgpuDevice {
pub kind: WgpuDeviceKind,
pub backend: WgpuBackend,
}
impl WgpuDevice {
pub fn new(kind: WgpuDeviceKind) -> Self {
Self {
kind,
backend: WgpuBackend::Auto,
}
}
pub fn on(self, backend: WgpuBackend) -> Self {
match self.kind {
WgpuDeviceKind::Existing(_) => self,
kind => Self { kind, backend },
}
}
}
impl From<WgpuDeviceKind> for WgpuDevice {
fn from(kind: WgpuDeviceKind) -> Self {
Self::new(kind)
}
}
impl Device for WgpuDevice {
fn from_id(device_id: DeviceId) -> Self {
let index = (device_id.index_id & WgpuDeviceKind::MAX_INDEX as u16) as usize;
let kind = match device_id.type_id {
0 => WgpuDeviceKind::DiscreteGpu(index),
1 => WgpuDeviceKind::IntegratedGpu(index),
2 => WgpuDeviceKind::VirtualGpu(index),
3 => WgpuDeviceKind::Cpu,
5 => return Self::new(WgpuDeviceKind::Existing(device_id.index_id as u32)),
6 => WgpuDeviceKind::Other(index),
_ => WgpuDeviceKind::DefaultDevice,
};
Self {
kind,
backend: WgpuBackend::from_bits(
(device_id.index_id >> WgpuBackend::SHIFT) & WgpuBackend::MASK,
),
}
}
fn to_id(&self) -> DeviceId {
let index = self.kind.index();
let index_id = match self.kind {
WgpuDeviceKind::Existing(_) => index,
_ => ((self.backend as u16) << WgpuBackend::SHIFT) | index,
};
DeviceId::new(self.kind.type_id(), index_id)
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn a_device_round_trips_its_kind_and_its_backend() {
for kind in [
WgpuDeviceKind::DiscreteGpu(2),
WgpuDeviceKind::IntegratedGpu(1),
WgpuDeviceKind::VirtualGpu(0),
WgpuDeviceKind::Cpu,
WgpuDeviceKind::DefaultDevice,
WgpuDeviceKind::Other(3),
] {
for backend in [
WgpuBackend::Auto,
WgpuBackend::Vulkan,
WgpuBackend::Metal,
WgpuBackend::Dx12,
WgpuBackend::Gl,
WgpuBackend::WebGpu,
] {
let device = WgpuDevice::new(kind.clone()).on(backend);
assert_eq!(WgpuDevice::from_id(device.to_id()), device);
}
}
}
#[test]
fn pinning_a_backend_changes_the_id() {
let auto = WgpuDevice::new(WgpuDeviceKind::DiscreteGpu(0));
let vulkan = auto.clone().on(WgpuBackend::Vulkan);
assert_ne!(auto.to_id(), vulkan.to_id());
}
#[test]
fn an_index_too_wide_for_the_id_is_kept_at_the_largest() {
let too_wide = WgpuDevice::new(WgpuDeviceKind::DiscreteGpu(WgpuDeviceKind::MAX_INDEX + 1));
assert_eq!(
WgpuDevice::from_id(too_wide.to_id()),
WgpuDevice::new(WgpuDeviceKind::DiscreteGpu(WgpuDeviceKind::MAX_INDEX))
);
}
#[test]
fn an_existing_device_keeps_its_whole_index() {
let existing = WgpuDevice::new(WgpuDeviceKind::Existing(60_000));
let pinned = existing.clone().on(WgpuBackend::Vulkan);
assert_eq!(pinned, existing);
assert_eq!(WgpuDevice::from_id(existing.to_id()), existing);
}
}