fix: q_readback buffer 12→total_actions (13 with Hold action)
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) <noreply@anthropic.com>
This commit is contained in:
@@ -1319,7 +1319,7 @@ pub struct GpuDqnTrainer {
|
||||
q_stats_buf: CudaSlice<f32>,
|
||||
/// GPU buffer for atom utilization accumulation [2 floats: sum_entropy, sum_utilized]
|
||||
atom_stats_buf: CudaSlice<f32>,
|
||||
/// 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::<f32>(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::<f32>(), flags)
|
||||
cudarc::driver::result::malloc_host(q_readback_size * std::mem::size_of::<f32>(), 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::<f32>() as u64;
|
||||
unsafe {
|
||||
cudarc::driver::result::memcpy_dtod_async(
|
||||
q_out_dst, self.q_out_buf.raw_ptr(),
|
||||
12 * std::mem::size_of::<f32>(), self.stream.cu_stream(),
|
||||
total_actions as usize * std::mem::size_of::<f32>(), 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::<f32>() / bs as f32;
|
||||
self.last_per_branch_q_gaps[d] = (max_q - mean_q).max(0.0);
|
||||
|
||||
Reference in New Issue
Block a user