//! Candle backend configuration and detection use candle_core::Device as CandleCoreDevice; use serde::{Deserialize, Serialize}; use tracing::info; /// Available Candle backends #[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] pub enum CandleBackend { /// CPU backend (always available) Cpu, /// CUDA backend (NVIDIA GPUs) Cuda, /// Metal backend (Apple Silicon) Metal, } impl Default for CandleBackend { fn default() -> Self { detect_best_backend() } } impl CandleBackend { /// Get the name of the backend pub fn name(&self) -> &'static str { match self { CandleBackend::Cpu => "cpu", CandleBackend::Cuda => "cuda", CandleBackend::Metal => "metal", } } /// Check if this backend supports GPU acceleration pub fn is_gpu(&self) -> bool { matches!(self, CandleBackend::Cuda | CandleBackend::Metal) } /// Check if this backend is available pub fn is_available(&self) -> bool { match self { CandleBackend::Cpu => true, CandleBackend::Cuda => cfg!(feature = "cuda"), CandleBackend::Metal => cfg!(feature = "metal") && cfg!(target_os = "macos"), } } /// Convert to Candle device pub fn to_candle_device(&self, ordinal: usize) -> CandleDevice { CandleDevice { backend: *self, ordinal, } } } /// Device wrapper for Candle #[derive(Debug, Clone, Serialize, Deserialize, Default)] pub struct CandleDevice { /// Backend type pub backend: CandleBackend, /// Device ordinal (for multi-GPU) pub ordinal: usize, } impl CandleDevice { /// Create a CPU device pub fn cpu() -> Self { Self { backend: CandleBackend::Cpu, ordinal: 0, } } /// Create a CUDA device pub fn cuda(ordinal: usize) -> Self { Self { backend: CandleBackend::Cuda, ordinal, } } /// Create a Metal device pub fn metal() -> Self { Self { backend: CandleBackend::Metal, ordinal: 0, } } /// Convert to Candle's native Device type pub fn to_candle(&self) -> candle_core::Result { match self.backend { CandleBackend::Cpu => Ok(CandleCoreDevice::Cpu), #[cfg(feature = "cuda")] CandleBackend::Cuda => CandleCoreDevice::new_cuda(self.ordinal), #[cfg(not(feature = "cuda"))] CandleBackend::Cuda => Ok(CandleCoreDevice::Cpu), #[cfg(feature = "metal")] CandleBackend::Metal => CandleCoreDevice::new_metal(self.ordinal), #[cfg(not(feature = "metal"))] CandleBackend::Metal => Ok(CandleCoreDevice::Cpu), } } /// Check if using GPU pub fn is_gpu(&self) -> bool { self.backend.is_gpu() } } /// Detect the best available backend based on features and hardware pub fn detect_best_backend() -> CandleBackend { // Prefer CUDA if available #[cfg(feature = "cuda")] { if is_cuda_runtime_available() { info!("Candle: Using CUDA backend"); return CandleBackend::Cuda; } } // Then Metal on macOS #[cfg(all(feature = "metal", target_os = "macos"))] { info!("Candle: Using Metal backend"); return CandleBackend::Metal; } // Fall back to CPU info!("Candle: Using CPU backend"); CandleBackend::Cpu } /// Check if CUDA runtime is available #[cfg(feature = "cuda")] fn is_cuda_runtime_available() -> bool { std::env::var("CUDA_VISIBLE_DEVICES").is_ok() || std::path::Path::new("/usr/local/cuda").exists() } #[cfg(not(feature = "cuda"))] fn is_cuda_runtime_available() -> bool { false } #[cfg(test)] mod tests { use super::*; #[test] fn test_backend_name() { assert_eq!(CandleBackend::Cpu.name(), "cpu"); assert_eq!(CandleBackend::Cuda.name(), "cuda"); assert_eq!(CandleBackend::Metal.name(), "metal"); } #[test] fn test_backend_is_gpu() { assert!(!CandleBackend::Cpu.is_gpu()); assert!(CandleBackend::Cuda.is_gpu()); assert!(CandleBackend::Metal.is_gpu()); } #[test] fn test_cpu_always_available() { assert!(CandleBackend::Cpu.is_available()); } #[test] fn test_device_creation() { let cpu = CandleDevice::cpu(); assert_eq!(cpu.backend, CandleBackend::Cpu); assert_eq!(cpu.ordinal, 0); let cuda = CandleDevice::cuda(1); assert_eq!(cuda.backend, CandleBackend::Cuda); assert_eq!(cuda.ordinal, 1); } #[test] fn test_detect_backend() { let backend = detect_best_backend(); // Should always succeed assert!(!backend.name().is_empty()); } }