//! Device abstraction replacing `candle_core::Device`. //! //! Provides `MlDevice` enum with CPU and CUDA variants. //! The CUDA variant holds `Arc` and `Arc` for //! direct cudarc interop. use std::sync::Arc; use crate::MLError; #[cfg(feature = "cuda")] use cudarc::driver::{CudaContext, CudaStream}; /// Device abstraction for the ML pipeline. /// /// Replaces `candle_core::Device`. Only two variants: CPU (for checkpoint /// serialization and preprocessing) and CUDA (for training and inference). #[derive(Clone)] pub enum MlDevice { /// Host CPU — used for checkpoint I/O and lightweight preprocessing. Cpu, /// CUDA GPU with context + stream handles. #[cfg(feature = "cuda")] Cuda { context: Arc, stream: Arc, }, } impl std::fmt::Debug for MlDevice { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { MlDevice::Cpu => write!(f, "MlDevice::Cpu"), #[cfg(feature = "cuda")] MlDevice::Cuda { .. } => write!(f, "MlDevice::Cuda"), } } } impl MlDevice { /// Create a CUDA device with the given ordinal (0-indexed). /// /// Initializes a `CudaDevice` and forks a non-default stream for all /// subsequent operations. #[cfg(feature = "cuda")] pub fn cuda(ordinal: usize) -> Result { let context = CudaContext::new(ordinal).map_err(|e| { MLError::DeviceError(format!("Failed to open CUDA device {ordinal}: {e}")) })?; // Check for stale CUDA errors from previous context users in the same process. // CudaContext::new reuses the primary context (cuDevicePrimaryCtxRetain), // so deferred errors from a previous test's Drop persist. if let Err(e) = context.check_err() { tracing::warn!("Stale CUDA error on device {ordinal} (may indicate prior kernel failure): {e}"); } let stream = context.new_stream().map_err(|e| { MLError::DeviceError(format!("Failed to create CUDA stream on device {ordinal}: {e}")) })?; Ok(MlDevice::Cuda { context, stream, }) } /// Try to create a CUDA device, falling back to CPU if unavailable. /// /// This is the recommended way to get a device when you want GPU /// acceleration but can tolerate CPU fallback. #[cfg(feature = "cuda")] pub fn cuda_if_available(ordinal: usize) -> Self { match Self::cuda(ordinal) { Ok(dev) => dev, Err(e) => { tracing::warn!("CUDA device {ordinal} unavailable ({e}), falling back to CPU"); MlDevice::Cpu } } } /// CPU-only fallback for `cuda_if_available` when CUDA feature is disabled. #[cfg(not(feature = "cuda"))] pub fn cuda_if_available(_ordinal: usize) -> Self { MlDevice::Cpu } /// Alias for `cuda()` — some callers use `new_cuda()`. #[cfg(feature = "cuda")] pub fn new_cuda(ordinal: usize) -> Result { Self::cuda(ordinal) } /// Identity accessor for compatibility — returns a reference to self. /// /// Some callers migrated from Candle use `obj.device()` to get the device /// handle. For `MlDevice` this is a no-op identity. pub fn device(&self) -> &Self { self } /// Returns `true` if this is a CUDA device. pub fn is_cuda(&self) -> bool { match self { MlDevice::Cpu => false, #[cfg(feature = "cuda")] MlDevice::Cuda { .. } => true, } } /// Returns `true` if this is the CPU device. pub fn is_cpu(&self) -> bool { matches!(self, MlDevice::Cpu) } /// Get the CUDA stream, or error if this is a CPU device. #[cfg(feature = "cuda")] pub fn cuda_stream(&self) -> Result<&Arc, MLError> { match self { MlDevice::Cuda { stream, .. } => Ok(stream), MlDevice::Cpu => Err(MLError::DeviceError( "cuda_stream() called on CPU device".to_owned(), )), } } /// Get the CudaContext handle, or error if this is a CPU device. #[cfg(feature = "cuda")] pub fn cuda_context(&self) -> Result<&Arc, MLError> { match self { MlDevice::Cuda { context, .. } => Ok(context), MlDevice::Cpu => Err(MLError::DeviceError( "cuda_context() called on CPU device".to_owned(), )), } } } impl std::fmt::Display for MlDevice { fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { match self { MlDevice::Cpu => write!(f, "cpu"), #[cfg(feature = "cuda")] MlDevice::Cuda { .. } => write!(f, "cuda"), } } }