diag: all adam_readback logs via eprintln

This commit is contained in:
jgrusewski
2026-04-06 01:00:51 +02:00
parent 1ba5ce93e7
commit f3c110dfbd

View File

@@ -2992,7 +2992,7 @@ impl GpuDqnTrainer {
pub fn replay_adam_and_readback(&mut self) -> Result<FusedTrainScalars, MLError> {
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)