diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index f77d67bd6..8bf58eb0c 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -7945,9 +7945,13 @@ impl GpuDqnTrainer { let block_dim = 256_u32; let grid_dim = ((batch_size as u32 + block_dim - 1) / block_dim).max(1); let q_out_ptr = self.q_out_buf.raw_ptr(); - let null_atom_stats = 0u64; - let null_q_var = 0u64; + let atom_stats_ptr = self.atom_stats_buf.raw_ptr(); + let q_var_ptr = self.q_var_buf_trainer.raw_ptr(); let atom_positions_buf_ptr = self.atom_positions_buf.raw_ptr(); + // Zero atom_stats before accumulation (kernel uses atomicAdd) + unsafe { + cudarc::driver::sys::cuMemsetD32Async(atom_stats_ptr, 0, 2, self.stream.cu_stream()); + } unsafe { self.stream .launch_builder(&self.expected_q_kernel) @@ -7961,8 +7965,8 @@ impl GpuDqnTrainer { .arg(&b2) .arg(&b3) .arg(&self.per_sample_support_ptr) - .arg(&null_atom_stats) - .arg(&null_q_var) + .arg(&atom_stats_ptr) + .arg(&q_var_ptr) .arg(&atom_positions_buf_ptr) .launch(LaunchConfig { grid_dim: (grid_dim, 1, 1),