diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index 3b52fdb41..0eda19606 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -2992,7 +2992,7 @@ impl GpuDqnTrainer { pub fn replay_adam_and_readback(&mut self) -> Result { static ADAM_CTR: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0); let adam_step = ADAM_CTR.fetch_add(1, std::sync::atomic::Ordering::Relaxed); - if adam_step < 3 { tracing::info!(adam_step, "adam_readback: enter"); } + if adam_step < 3 { eprintln!("adam_readback: step {adam_step} — enter"); } // 1. Collect previous step's scalars (if pending) let prev_scalars = if self.readback_pending { @@ -3016,11 +3016,11 @@ impl GpuDqnTrainer { }; // 2. Launch current step: grad_norm + adam - if adam_step < 3 { tracing::info!(adam_step, "adam_readback: compute_grad_norm"); } + if adam_step < 3 { eprintln!("adam_readback: step {adam_step} — compute_grad_norm"); } self.compute_grad_norm_outside_graph()?; - if adam_step < 3 { tracing::info!(adam_step, "adam_readback: replay_adam"); } + if adam_step < 3 { eprintln!("adam_readback: step {adam_step} — replay_adam"); } self.replay_adam()?; - if adam_step < 3 { tracing::info!(adam_step, "adam_readback: async readback"); } + if adam_step < 3 { eprintln!("adam_readback: step {adam_step} — async readback start"); } // 3. Async DtoH for current step's scalars (non-blocking) unsafe { @@ -3047,6 +3047,8 @@ impl GpuDqnTrainer { ); } + if adam_step < 3 { eprintln!("adam_readback: step {adam_step} — async readback done, recording event"); } + // 4. Record event — marks the point after the async copies let event = self.stream.record_event(None).map_err(|e| MLError::ModelError(format!("readback event record: {e}")) @@ -3054,6 +3056,8 @@ impl GpuDqnTrainer { self.readback_event = Some(event); self.readback_pending = true; + if adam_step < 3 { eprintln!("adam_readback: step {adam_step} — returning"); } + // Return previous step's values self.scalars_readback_host = [prev_scalars.total_loss, prev_scalars.grad_norm]; Ok(prev_scalars)