diff --git a/crates/ml-core/src/device.rs b/crates/ml-core/src/device.rs index c5436e6ab..66ae20cb6 100644 --- a/crates/ml-core/src/device.rs +++ b/crates/ml-core/src/device.rs @@ -64,7 +64,10 @@ impl MlDevice { pub fn cuda_if_available(ordinal: usize) -> Self { match Self::cuda(ordinal) { Ok(dev) => dev, - Err(_) => MlDevice::Cpu, + Err(e) => { + tracing::warn!("CUDA device {ordinal} unavailable ({e}), falling back to CPU"); + MlDevice::Cpu + } } } diff --git a/crates/ml/src/trainers/dqn/trainer/constructor.rs b/crates/ml/src/trainers/dqn/trainer/constructor.rs index 107b76f3f..e8fc349c7 100644 --- a/crates/ml/src/trainers/dqn/trainer/constructor.rs +++ b/crates/ml/src/trainers/dqn/trainer/constructor.rs @@ -118,11 +118,19 @@ impl DQNTrainer { } // Use override device if provided (hyperopt shares one CUDA context), - // otherwise auto-detect GPU + // otherwise require CUDA GPU — no CPU fallback in CUDA builds. let device = if let Some(dev) = override_device { dev } else { - MlDevice::cuda_if_available(0) + match MlDevice::cuda(0) { + Ok(dev) => dev, + Err(e) => { + tracing::error!("CUDA device init failed: {e}"); + return Err(anyhow::anyhow!( + "DQNTrainer requires CUDA GPU. CUDA init failed: {e}" + )); + } + } }; // Fork a dedicated CudaStream once. All GPU components (experience collector,