Skip to main content

aria_kernel/
compute.rs

1//! Local compute preference (`auto|cpu|cuda`) vs hybrid execution (cloud routing).
2
3use crate::cuda_rt;
4use crate::{EngineError, SimdMode};
5
6/// CLI / config preference. Orthogonal to `hybrid_execution`.
7#[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/// Resolved backend used by Session GEMM.
29#[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
43/// Resolve preference. `Cuda` never silently falls back to CPU.
44pub 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}