feat(isv): add 4 pinned scratch scalars for ISV signal sources
td_error_scratch, ensemble_var_scratch, reward_scratch, atom_util_scratch — all pinned device-mapped. atom_util and reward_scratch written in reduce_current_q_stats. td_error and ensemble_var wired in later tasks. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -1296,6 +1296,18 @@ pub struct GpuDqnTrainer {
|
||||
q_mag_pre_risk: CudaSlice<f32>, // [B, b1_size] saved for backward
|
||||
risk_grad_buf: CudaSlice<f32>, // [risk_param_count]
|
||||
risk_adam_step: i32,
|
||||
|
||||
// ── ISV scratch scalars (pinned device-mapped, written by GPU kernels) ──
|
||||
// TODO(isv): Wire C51 loss kernel to write mean |TD-error| here
|
||||
td_error_scratch_pinned: *mut f32,
|
||||
td_error_scratch_dev_ptr: u64,
|
||||
// TODO(isv): Wire ensemble aggregate kernel to write mean variance here
|
||||
ensemble_var_scratch_pinned: *mut f32,
|
||||
ensemble_var_scratch_dev_ptr: u64,
|
||||
reward_scratch_pinned: *mut f32,
|
||||
reward_scratch_dev_ptr: u64,
|
||||
atom_util_scratch_pinned: *mut f32,
|
||||
atom_util_scratch_dev_ptr: u64,
|
||||
}
|
||||
|
||||
impl GpuDqnTrainer {
|
||||
@@ -1589,6 +1601,18 @@ impl Drop for GpuDqnTrainer {
|
||||
if !self.iqn_readiness_pinned.is_null() {
|
||||
let _ = unsafe { cudarc::driver::result::free_host(self.iqn_readiness_pinned.cast()) };
|
||||
}
|
||||
if !self.td_error_scratch_pinned.is_null() {
|
||||
let _ = unsafe { cudarc::driver::result::free_host(self.td_error_scratch_pinned.cast()) };
|
||||
}
|
||||
if !self.ensemble_var_scratch_pinned.is_null() {
|
||||
let _ = unsafe { cudarc::driver::result::free_host(self.ensemble_var_scratch_pinned.cast()) };
|
||||
}
|
||||
if !self.reward_scratch_pinned.is_null() {
|
||||
let _ = unsafe { cudarc::driver::result::free_host(self.reward_scratch_pinned.cast()) };
|
||||
}
|
||||
if !self.atom_util_scratch_pinned.is_null() {
|
||||
let _ = unsafe { cudarc::driver::result::free_host(self.atom_util_scratch_pinned.cast()) };
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -4680,6 +4704,87 @@ impl GpuDqnTrainer {
|
||||
(host_ptr as *mut f32, dev_ptr_out)
|
||||
};
|
||||
|
||||
// ── ISV scratch scalars (pinned device-mapped) ────────────────────
|
||||
let (td_error_scratch_pinned, td_error_scratch_dev_ptr) = {
|
||||
let mut host_ptr: *mut std::ffi::c_void = std::ptr::null_mut();
|
||||
let mut dev_ptr_out: u64 = 0;
|
||||
unsafe {
|
||||
let rc = cudarc::driver::sys::cuMemAllocHost_v2(
|
||||
&mut host_ptr,
|
||||
std::mem::size_of::<f32>(),
|
||||
);
|
||||
assert_eq!(rc, cudarc::driver::sys::cudaError_enum::CUDA_SUCCESS, "cuMemAllocHost for td_error_scratch");
|
||||
let rc2 = cudarc::driver::sys::cuMemHostGetDevicePointer_v2(
|
||||
&mut dev_ptr_out,
|
||||
host_ptr,
|
||||
0,
|
||||
);
|
||||
assert_eq!(rc2, cudarc::driver::sys::cudaError_enum::CUDA_SUCCESS, "cuMemHostGetDevicePointer for td_error_scratch");
|
||||
*(host_ptr as *mut f32) = 0.0;
|
||||
}
|
||||
(host_ptr as *mut f32, dev_ptr_out)
|
||||
};
|
||||
|
||||
let (ensemble_var_scratch_pinned, ensemble_var_scratch_dev_ptr) = {
|
||||
let mut host_ptr: *mut std::ffi::c_void = std::ptr::null_mut();
|
||||
let mut dev_ptr_out: u64 = 0;
|
||||
unsafe {
|
||||
let rc = cudarc::driver::sys::cuMemAllocHost_v2(
|
||||
&mut host_ptr,
|
||||
std::mem::size_of::<f32>(),
|
||||
);
|
||||
assert_eq!(rc, cudarc::driver::sys::cudaError_enum::CUDA_SUCCESS, "cuMemAllocHost for ensemble_var_scratch");
|
||||
let rc2 = cudarc::driver::sys::cuMemHostGetDevicePointer_v2(
|
||||
&mut dev_ptr_out,
|
||||
host_ptr,
|
||||
0,
|
||||
);
|
||||
assert_eq!(rc2, cudarc::driver::sys::cudaError_enum::CUDA_SUCCESS, "cuMemHostGetDevicePointer for ensemble_var_scratch");
|
||||
*(host_ptr as *mut f32) = 0.0;
|
||||
}
|
||||
(host_ptr as *mut f32, dev_ptr_out)
|
||||
};
|
||||
|
||||
let (reward_scratch_pinned, reward_scratch_dev_ptr) = {
|
||||
let mut host_ptr: *mut std::ffi::c_void = std::ptr::null_mut();
|
||||
let mut dev_ptr_out: u64 = 0;
|
||||
unsafe {
|
||||
let rc = cudarc::driver::sys::cuMemAllocHost_v2(
|
||||
&mut host_ptr,
|
||||
std::mem::size_of::<f32>(),
|
||||
);
|
||||
assert_eq!(rc, cudarc::driver::sys::cudaError_enum::CUDA_SUCCESS, "cuMemAllocHost for reward_scratch");
|
||||
let rc2 = cudarc::driver::sys::cuMemHostGetDevicePointer_v2(
|
||||
&mut dev_ptr_out,
|
||||
host_ptr,
|
||||
0,
|
||||
);
|
||||
assert_eq!(rc2, cudarc::driver::sys::cudaError_enum::CUDA_SUCCESS, "cuMemHostGetDevicePointer for reward_scratch");
|
||||
*(host_ptr as *mut f32) = 0.0;
|
||||
}
|
||||
(host_ptr as *mut f32, dev_ptr_out)
|
||||
};
|
||||
|
||||
let (atom_util_scratch_pinned, atom_util_scratch_dev_ptr) = {
|
||||
let mut host_ptr: *mut std::ffi::c_void = std::ptr::null_mut();
|
||||
let mut dev_ptr_out: u64 = 0;
|
||||
unsafe {
|
||||
let rc = cudarc::driver::sys::cuMemAllocHost_v2(
|
||||
&mut host_ptr,
|
||||
std::mem::size_of::<f32>(),
|
||||
);
|
||||
assert_eq!(rc, cudarc::driver::sys::cudaError_enum::CUDA_SUCCESS, "cuMemAllocHost for atom_util_scratch");
|
||||
let rc2 = cudarc::driver::sys::cuMemHostGetDevicePointer_v2(
|
||||
&mut dev_ptr_out,
|
||||
host_ptr,
|
||||
0,
|
||||
);
|
||||
assert_eq!(rc2, cudarc::driver::sys::cudaError_enum::CUDA_SUCCESS, "cuMemHostGetDevicePointer for atom_util_scratch");
|
||||
*(host_ptr as *mut f32) = 0.0;
|
||||
}
|
||||
(host_ptr as *mut f32, dev_ptr_out)
|
||||
};
|
||||
|
||||
// ── Cross-branch graph message passing ──────────────────────────
|
||||
let cpbi_module_graph = stream.context().load_cubin(EXPECTED_Q_CUBIN.to_vec())
|
||||
.map_err(|e| MLError::ModelError(format!("cpbi cubin (graph_msg): {e}")))?;
|
||||
@@ -5346,6 +5451,14 @@ impl GpuDqnTrainer {
|
||||
q_mag_pre_risk,
|
||||
risk_grad_buf,
|
||||
risk_adam_step: 0,
|
||||
td_error_scratch_pinned,
|
||||
td_error_scratch_dev_ptr,
|
||||
ensemble_var_scratch_pinned,
|
||||
ensemble_var_scratch_dev_ptr,
|
||||
reward_scratch_pinned,
|
||||
reward_scratch_dev_ptr,
|
||||
atom_util_scratch_pinned,
|
||||
atom_util_scratch_dev_ptr,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -6726,6 +6839,11 @@ impl GpuDqnTrainer {
|
||||
);
|
||||
}
|
||||
|
||||
// ISV: write atom utilization to pinned scratch for GPU-side ISV signal update
|
||||
unsafe { *self.atom_util_scratch_pinned = host[6].clamp(0.0, 1.0); }
|
||||
// ISV: write Q-mean as reward proxy to pinned scratch
|
||||
unsafe { *self.reward_scratch_pinned = host[3]; }
|
||||
|
||||
// xLSTM temporal context: update matrix memory and write context[8] on GPU.
|
||||
self.launch_qlstm_step()?;
|
||||
|
||||
|
||||
Reference in New Issue
Block a user