diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index b72566d07..94bad02de 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -3256,20 +3256,20 @@ impl GpuDqnTrainer { // Extract raw pointers for online branch logits (contiguous [B, (B0+B1+B2)*NA]) let on_b_base = self.on_b_logits_buf.raw_ptr(); let on_b0_ptr = on_b_base; - let on_b1_ptr = on_b_base + (b * b0 * na * 4) as u64; - let on_b2_ptr = on_b1_ptr + (b * b1 * na * 4) as u64; + let on_b1_ptr = on_b_base + (b * b0 * na * std::mem::size_of::()) as u64; + let on_b2_ptr = on_b1_ptr + (b * b1 * na * std::mem::size_of::()) as u64; // Target branch logits let tg_b_base = self.tg_b_logits_buf.raw_ptr(); let tg_b0_ptr = tg_b_base; - let tg_b1_ptr = tg_b_base + (b * b0 * na * 4) as u64; - let tg_b2_ptr = tg_b1_ptr + (b * b1 * na * 4) as u64; + let tg_b1_ptr = tg_b_base + (b * b0 * na * std::mem::size_of::()) as u64; + let tg_b2_ptr = tg_b1_ptr + (b * b1 * na * std::mem::size_of::()) as u64; // Online-next branch logits (Double DQN action selector) let on_next_b_base = self.on_next_b_logits_buf.raw_ptr(); let on_next_b0_ptr = on_next_b_base; - let on_next_b1_ptr = on_next_b_base + (b * b0 * na * 4) as u64; - let on_next_b2_ptr = on_next_b1_ptr + (b * b1 * na * 4) as u64; + let on_next_b1_ptr = on_next_b_base + (b * b0 * na * std::mem::size_of::()) as u64; + let on_next_b2_ptr = on_next_b1_ptr + (b * b1 * na * std::mem::size_of::()) as u64; // N-step returns: use gamma^n for the Bellman projection. // The experience collector pre-computes R_n = sum(gamma^i * r_i). @@ -3402,20 +3402,20 @@ impl GpuDqnTrainer { // Extract raw pointers for online branch logits (contiguous [B, (B0+B1+B2)*NA]) let on_b_base = self.on_b_logits_buf.raw_ptr(); let on_b0_ptr = on_b_base; - let on_b1_ptr = on_b_base + (b * b0 * na * 4) as u64; - let on_b2_ptr = on_b1_ptr + (b * b1 * na * 4) as u64; + let on_b1_ptr = on_b_base + (b * b0 * na * std::mem::size_of::()) as u64; + let on_b2_ptr = on_b1_ptr + (b * b1 * na * std::mem::size_of::()) as u64; // Target branch logits let tg_b_base = self.tg_b_logits_buf.raw_ptr(); let tg_b0_ptr = tg_b_base; - let tg_b1_ptr = tg_b_base + (b * b0 * na * 4) as u64; - let tg_b2_ptr = tg_b1_ptr + (b * b1 * na * 4) as u64; + let tg_b1_ptr = tg_b_base + (b * b0 * na * std::mem::size_of::()) as u64; + let tg_b2_ptr = tg_b1_ptr + (b * b1 * na * std::mem::size_of::()) as u64; // Online-next branch logits (Double DQN action selector) let on_next_b_base = self.on_next_b_logits_buf.raw_ptr(); let on_next_b0_ptr = on_next_b_base; - let on_next_b1_ptr = on_next_b_base + (b * b0 * na * 4) as u64; - let on_next_b2_ptr = on_next_b1_ptr + (b * b1 * na * 4) as u64; + let on_next_b1_ptr = on_next_b_base + (b * b0 * na * std::mem::size_of::()) as u64; + let on_next_b2_ptr = on_next_b1_ptr + (b * b1 * na * std::mem::size_of::()) as u64; // N-step returns: use gamma^n for the Bellman projection. let gamma = self.config.gamma.powi(self.config.n_steps as i32);