1use crate::cuda_rt;
4use crate::{EngineError, SimdMode};
5
6#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
8pub enum ComputePref {
9 #[default]
10 Auto,
11 Cpu,
12 Cuda,
13}
14
15impl ComputePref {
16 pub fn parse(raw: &str) -> Result<Self, EngineError> {
17 match raw.trim().to_ascii_lowercase().as_str() {
18 "" | "auto" => Ok(Self::Auto),
19 "cpu" => Ok(Self::Cpu),
20 "cuda" | "gpu" => Ok(Self::Cuda),
21 other => Err(EngineError::InvalidParam(format!(
22 "compute must be auto|cpu|cuda, got {other:?}"
23 ))),
24 }
25 }
26}
27
28#[derive(Debug, Clone, Copy, PartialEq, Eq)]
30pub enum ComputeBackend {
31 Cpu,
32 Cuda,
33}
34
35pub fn cpu_simd_label() -> String {
36 match SimdMode::auto() {
37 SimdMode::Neon => "neon".into(),
38 SimdMode::Avx2 => "avx2".into(),
39 SimdMode::Scalar => "scalar".into(),
40 }
41}
42
43pub fn resolve_compute(pref: ComputePref) -> Result<(ComputeBackend, String), EngineError> {
45 match pref {
46 ComputePref::Cpu => Ok((
47 ComputeBackend::Cpu,
48 format!("cpu simd={}", cpu_simd_label()),
49 )),
50 ComputePref::Cuda => {
51 let info = cuda_rt::device_info().map_err(|e| {
52 EngineError::Unsupported(format!(
53 "compute=cuda requested but CUDA/cuBLAS is unavailable: {e}"
54 ))
55 })?;
56 Ok((ComputeBackend::Cuda, format!("cuda {info}")))
57 }
58 ComputePref::Auto => match cuda_rt::device_info() {
59 Ok(info) => Ok((ComputeBackend::Cuda, format!("cuda {info}"))),
60 Err(_) => Ok((
61 ComputeBackend::Cpu,
62 format!("cpu simd={}", cpu_simd_label()),
63 )),
64 },
65 }
66}
67
68#[cfg(test)]
69mod tests {
70 use super::*;
71
72 #[test]
73 fn parse_compute_pref() {
74 assert_eq!(ComputePref::parse("auto").unwrap(), ComputePref::Auto);
75 assert_eq!(ComputePref::parse("CPU").unwrap(), ComputePref::Cpu);
76 assert_eq!(ComputePref::parse("cuda").unwrap(), ComputePref::Cuda);
77 assert_eq!(ComputePref::parse("gpu").unwrap(), ComputePref::Cuda);
78 assert!(matches!(
79 ComputePref::parse("tpu"),
80 Err(EngineError::InvalidParam(_))
81 ));
82 }
83
84 #[test]
85 fn auto_never_errors() {
86 let (backend, label) = resolve_compute(ComputePref::Auto).unwrap();
87 assert!(!label.is_empty());
88 match backend {
89 ComputeBackend::Cpu => assert!(label.contains("cpu")),
90 ComputeBackend::Cuda => assert!(label.contains("cuda")),
91 }
92 }
93
94 #[test]
95 fn explicit_cuda_fails_without_device() {
96 if cuda_rt::device_info().is_err() {
97 let err = resolve_compute(ComputePref::Cuda).unwrap_err();
98 assert!(matches!(err, EngineError::Unsupported(_)), "{err}");
99 }
100 }
101}