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:
jgrusewski
2026-04-16 22:54:23 +02:00
parent 6dd2f3e68b
commit f9d90f864b

View File

@@ -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()?;