use std::fmt;
use std::str::FromStr;
use serde::{Deserialize, Serialize};
use crate::EngineError;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub enum BackendKind {
CpuSimd,
Cuda,
Metal,
}
impl BackendKind {
pub const ALL: [BackendKind; 3] = [BackendKind::Cuda, BackendKind::Metal, BackendKind::CpuSimd];
pub const fn as_str(self) -> &'static str {
match self {
BackendKind::CpuSimd => "cpu",
BackendKind::Cuda => "cuda",
BackendKind::Metal => "metal",
}
}
pub const fn is_gpu(self) -> bool {
!matches!(self, BackendKind::CpuSimd)
}
}
impl fmt::Display for BackendKind {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.write_str(self.as_str())
}
}
impl FromStr for BackendKind {
type Err = EngineError;
fn from_str(s: &str) -> Result<Self, Self::Err> {
match s.trim().to_ascii_lowercase().as_str() {
"cpu" | "cpusimd" | "cpu-simd" | "cpu_simd" => Ok(BackendKind::CpuSimd),
"cuda" => Ok(BackendKind::Cuda),
"metal" => Ok(BackendKind::Metal),
other => Err(EngineError::plan(format!(
"unknown backend '{other}'; expected one of cpu, cuda, metal"
))),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub struct DeviceId {
pub backend: BackendKind,
pub ordinal: u32,
}
impl DeviceId {
pub const CPU: DeviceId = DeviceId {
backend: BackendKind::CpuSimd,
ordinal: 0,
};
pub const fn new(backend: BackendKind, ordinal: u32) -> Self {
Self { backend, ordinal }
}
}
impl fmt::Display for DeviceId {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "{}:{}", self.backend, self.ordinal)
}
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used)]
mod tests {
use super::*;
#[test]
fn parses_all_names_case_insensitively() {
assert_eq!("cpu".parse::<BackendKind>().unwrap(), BackendKind::CpuSimd);
assert_eq!(" CUDA ".parse::<BackendKind>().unwrap(), BackendKind::Cuda);
assert_eq!("Metal".parse::<BackendKind>().unwrap(), BackendKind::Metal);
assert!("tpu".parse::<BackendKind>().is_err());
}
#[test]
fn display_round_trips_through_parse() {
for kind in BackendKind::ALL {
assert_eq!(kind.as_str().parse::<BackendKind>().unwrap(), kind);
}
}
#[test]
fn device_id_display() {
assert_eq!(DeviceId::new(BackendKind::Cuda, 1).to_string(), "cuda:1");
assert_eq!(DeviceId::CPU.to_string(), "cpu:0");
assert!(!BackendKind::CpuSimd.is_gpu());
assert!(BackendKind::Metal.is_gpu());
}
}