diff --git a/ml/src/liquid/candle_cfc.rs b/ml/src/liquid/candle_cfc.rs index 6253a78cc..d83d57673 100644 --- a/ml/src/liquid/candle_cfc.rs +++ b/ml/src/liquid/candle_cfc.rs @@ -10,7 +10,7 @@ use serde::{Deserialize, Serialize}; use crate::MLError; /// Device configuration for CfC training -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub enum DeviceConfig { Cpu, Cuda(usize), @@ -21,8 +21,11 @@ impl DeviceConfig { pub fn resolve(&self) -> Result { match self { DeviceConfig::Cpu => Ok(Device::Cpu), - DeviceConfig::Cuda(id) => Device::cuda_if_available(*id).map_err(|e| { - MLError::ConfigurationError(format!("CUDA device {} error: {}", id, e)) + DeviceConfig::Cuda(id) => Device::new_cuda(*id).map_err(|e| { + MLError::ConfigurationError(format!( + "CUDA device {} required but unavailable: {}", + id, e + )) }), DeviceConfig::Auto => Device::cuda_if_available(0) .map_err(|e| MLError::ConfigurationError(format!("Auto device error: {}", e))), @@ -35,7 +38,7 @@ impl DeviceConfig { /// This is distinct from `cells::CfCConfig` which configures the FixedPoint inference path. /// `CfCTrainConfig` configures the Candle-based differentiable training path with /// gradient descent, dropout, and market regime adaptation. -#[derive(Debug, Clone, Serialize, Deserialize)] +#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)] pub struct CfCTrainConfig { pub input_size: usize, pub hidden_size: usize,