#[derive(Debug, Clone, PartialEq, Eq)]
pub(crate) enum DeviceRequest {
Cpu,
Cuda(usize),
Metal,
Unknown(String),
}
pub(crate) fn parse_device_request(raw: &str) -> DeviceRequest {
let requested = raw.trim().to_ascii_lowercase();
if requested.is_empty() || requested == "cpu" {
return DeviceRequest::Cpu;
}
if requested == "cuda" {
return DeviceRequest::Cuda(0);
}
if let Some(idx) = requested.strip_prefix("cuda:") {
return DeviceRequest::Cuda(idx.parse::<usize>().unwrap_or(0));
}
if requested == "metal" {
return DeviceRequest::Metal;
}
DeviceRequest::Unknown(requested)
}
#[cfg(test)]
mod device_request_tests {
use super::{parse_device_request, DeviceRequest};
#[test]
fn unset_or_empty_is_cpu() {
assert_eq!(parse_device_request(""), DeviceRequest::Cpu);
assert_eq!(parse_device_request(" "), DeviceRequest::Cpu);
}
#[test]
fn explicit_cpu_is_cpu_case_and_space_insensitive() {
assert_eq!(parse_device_request("cpu"), DeviceRequest::Cpu);
assert_eq!(parse_device_request("CPU"), DeviceRequest::Cpu);
assert_eq!(parse_device_request(" Cpu "), DeviceRequest::Cpu);
}
#[test]
fn bare_cuda_is_device_zero() {
assert_eq!(parse_device_request("cuda"), DeviceRequest::Cuda(0));
assert_eq!(parse_device_request("CUDA"), DeviceRequest::Cuda(0));
}
#[test]
fn cuda_n_selects_the_index() {
assert_eq!(parse_device_request("cuda:0"), DeviceRequest::Cuda(0));
assert_eq!(parse_device_request("cuda:1"), DeviceRequest::Cuda(1));
assert_eq!(parse_device_request(" cuda:2 "), DeviceRequest::Cuda(2));
}
#[test]
fn cuda_with_garbage_index_clamps_to_zero() {
assert_eq!(parse_device_request("cuda:x"), DeviceRequest::Cuda(0));
assert_eq!(parse_device_request("cuda:"), DeviceRequest::Cuda(0));
}
#[test]
fn metal_is_metal() {
assert_eq!(parse_device_request("metal"), DeviceRequest::Metal);
assert_eq!(parse_device_request("Metal"), DeviceRequest::Metal);
}
#[test]
fn unrecognized_is_a_named_unknown() {
assert_eq!(parse_device_request("rocm"), DeviceRequest::Unknown("rocm".to_string()));
assert_eq!(parse_device_request("gpu"), DeviceRequest::Unknown("gpu".to_string()));
}
}