From f9d90f864b327bae8233d910f28d8c09bb33f9bf Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Thu, 16 Apr 2026 22:54:23 +0200 Subject: [PATCH] feat(isv): add 4 pinned scratch scalars for ISV signal sources MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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) --- .../ml/src/cuda_pipeline/gpu_dqn_trainer.rs | 118 ++++++++++++++++++ 1 file changed, 118 insertions(+) diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index b80c31519..9fc4d20f2 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -1296,6 +1296,18 @@ pub struct GpuDqnTrainer { q_mag_pre_risk: CudaSlice, // [B, b1_size] saved for backward risk_grad_buf: CudaSlice, // [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::(), + ); + 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::(), + ); + 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::(), + ); + 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::(), + ); + 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()?;