From d625ca28e8331d842f3d305262e6218e72b3a77f Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Mon, 20 Apr 2026 01:39:20 +0200 Subject: [PATCH] =?UTF-8?q?fix:=20q=5Freadback=20buffer=2012=E2=86=92total?= =?UTF-8?q?=5Factions=20(13=20with=20Hold=20action)?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Hardcoded 12 Q-values in pinned readback buffer and per-branch Q-gap slice caused panic with b0=4 (Hold action: 4+3+3+3=13 total actions). Now uses dynamic total_actions from config. Co-Authored-By: Claude Opus 4.6 (1M context) --- .../ml/src/cuda_pipeline/gpu_dqn_trainer.rs | 26 ++++++++++--------- 1 file changed, 14 insertions(+), 12 deletions(-) diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index 8c56d8714..8bcc420d9 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -1319,7 +1319,7 @@ pub struct GpuDqnTrainer { q_stats_buf: CudaSlice, /// GPU buffer for atom utilization accumulation [2 floats: sum_entropy, sum_utilized] atom_stats_buf: CudaSlice, - /// Pinned device-mapped readback for q_stats [7] + q_out sample 0 [12] = 19 floats. + /// Pinned device-mapped readback for q_stats [7] + q_out sample 0 [total_actions] floats. /// GPU writes via q_readback_dev_ptr, CPU reads via q_readback_pinned — zero sync. q_readback_dev_ptr: u64, /// DtoH uses pinned DMA-capable destination for faster async transfer. @@ -5395,10 +5395,12 @@ impl GpuDqnTrainer { .map_err(|e| MLError::ModelError(format!("alloc q_stats_f32: {e}")))?; let atom_stats_buf = stream.alloc_zeros::(2) .map_err(|e| MLError::ModelError(format!("alloc atom_stats: {e}")))?; - // q_readback — pinned host buffer for DMA-capable DtoH of q_stats[7] + q_out[12] + // q_readback — pinned host buffer for q_stats[7] + q_out sample 0 [total_actions] + let total_actions = config.branch_0_size + config.branch_1_size + config.branch_2_size + config.branch_3_size; + let q_readback_size = 7 + total_actions; // 7 stats + 13 Q-values = 20 let q_readback_pinned: *mut f32 = unsafe { let flags = cudarc::driver::sys::CU_MEMHOSTALLOC_DEVICEMAP; - cudarc::driver::result::malloc_host(19 * std::mem::size_of::(), flags) + cudarc::driver::result::malloc_host(q_readback_size * std::mem::size_of::(), flags) .map_err(|e| MLError::ModelError(format!("pinned q_readback alloc: {e}")))? as *mut f32 }; @@ -8834,12 +8836,12 @@ impl GpuDqnTrainer { }) .map_err(|e| MLError::ModelError(format!("q_stats_kernel: {e}")))?; } - // Copy q_out sample 0 (12 floats) to pinned buffer [7..19] for per-branch Q-gap. + // Copy q_out sample 0 (total_actions floats) to pinned buffer [7..] for per-branch Q-gap. let q_out_dst = self.q_readback_dev_ptr + 7 * std::mem::size_of::() as u64; unsafe { cudarc::driver::result::memcpy_dtod_async( q_out_dst, self.q_out_buf.raw_ptr(), - 12 * std::mem::size_of::(), self.stream.cu_stream(), + total_actions as usize * std::mem::size_of::(), self.stream.cu_stream(), ).map_err(|e| MLError::ModelError(format!("q_out sample0 DtoD: {e}")))?; } // No sync — one-step lag is fine for monitoring. populate_q_out + q_stats_reduce @@ -8854,14 +8856,14 @@ impl GpuDqnTrainer { std::ptr::copy_nonoverlapping(self.q_readback_pinned, h.as_mut_ptr(), 7); h }; - let q12: [f32; 12] = unsafe { - let mut q = [0.0_f32; 12]; - std::ptr::copy_nonoverlapping(self.q_readback_pinned.add(7), q.as_mut_ptr(), 12); - q + let total_actions = self.total_actions() as usize; + let mut q_vals = vec![0.0_f32; total_actions]; + unsafe { + std::ptr::copy_nonoverlapping(self.q_readback_pinned.add(7), q_vals.as_mut_ptr(), total_actions); }; - // Per-branch Q-gap: first sample's 12 Q-values. - // Branch layout: [dir(3), mag(3), ord(3), urg(3)]. + // Per-branch Q-gap: first sample's Q-values. + // Branch layout: [dir(4), mag(3), ord(3), urg(3)] = 13 total. let branch_sizes = [ self.config.branch_0_size, self.config.branch_1_size, self.config.branch_2_size, self.config.branch_3_size, @@ -8869,7 +8871,7 @@ impl GpuDqnTrainer { let mut offset = 0; for d in 0..4 { let bs = branch_sizes[d]; - let slice = &q12[offset..offset + bs]; + let slice = &q_vals[offset..offset + bs]; let max_q = slice.iter().cloned().fold(f32::NEG_INFINITY, f32::max); let mean_q = slice.iter().sum::() / bs as f32; self.last_per_branch_q_gaps[d] = (max_q - mean_q).max(0.0);