From 954bca690d8691face8a4abc71a21104fc32be7f Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sat, 28 Mar 2026 16:27:52 +0100 Subject: [PATCH] fix(bf16): graph-safe weight pointers + rename alloc_bf16 MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Root cause identified: bf16_weight_ptrs called device_ptr() inside CUDA graph capture, which records cudarc events — not allowed during capture. Created bf16_weight_ptrs_from_base() that takes pre-resolved u64 base pointer from CachedPtrs (no device_ptr calls, graph-safe). Replaced ALL 12 bf16_weight_ptrs calls in gpu_dqn_trainer with bf16_weight_ptrs_from_base using self.ptrs.params_buf/target_params_buf. Debug forward (no graph) confirms cuBLAS GemmEx produces valid logits. Q-value divergence persists in graph mode — Adam or readback issue TBD. Co-Authored-By: Claude Opus 4.6 (1M context) --- .../ml/src/cuda_pipeline/batched_forward.rs | 8 +++ .../ml/src/cuda_pipeline/gpu_dqn_trainer.rs | 64 ++++++++++++++----- 2 files changed, 57 insertions(+), 15 deletions(-) diff --git a/crates/ml/src/cuda_pipeline/batched_forward.rs b/crates/ml/src/cuda_pipeline/batched_forward.rs index 0d25ce399..139c74738 100644 --- a/crates/ml/src/cuda_pipeline/batched_forward.rs +++ b/crates/ml/src/cuda_pipeline/batched_forward.rs @@ -629,7 +629,15 @@ pub fn bf16_weight_ptrs( let _no_drop = ManuallyDrop::new(guard); ptr }; + bf16_weight_ptrs_from_base(base, param_sizes) +} +/// Compute 20 raw BF16 device pointers from a pre-resolved base pointer. +/// Graph-safe: no device_ptr calls, no event recording. +pub fn bf16_weight_ptrs_from_base( + base: u64, + param_sizes: &[usize; 20], +) -> [u64; 20] { let mut ptrs = [0_u64; 20]; let mut byte_offset: u64 = 0; for i in 0..20 { diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index b86b36d95..f9755d664 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -55,7 +55,7 @@ use tracing::info; use crate::MLError; use super::gpu_attention::GpuAttention; use super::gpu_weights::{DuelingWeightSet, BranchingWeightSet}; -use super::batched_forward::{CublasForward, bf16_weight_ptrs}; +use super::batched_forward::{CublasForward, bf16_weight_ptrs_from_base}; use super::batched_backward::{CublasBackward, alloc_backward_scratch, raw_f32_ptr as bw_raw_ptr}; // ── Precompiled cubins (build.rs → include_bytes! → ZERO runtime nvcc) ────── @@ -789,7 +789,7 @@ impl GpuDqnTrainer { let dy = bw_raw_ptr(&self.bw_d_h_s2, &self.stream); let x = bw_raw_ptr(&self.save_h_s1, &self.stream); let param_sizes = compute_param_sizes(&self.config); - let w_ptrs = bf16_weight_ptrs(&self.params_buf, ¶m_sizes, &self.stream); + let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, ¶m_sizes); let w = w_ptrs[2]; // w_s2 let scratch_base = self.ptrs.iqn_trunk_m; @@ -831,7 +831,7 @@ impl GpuDqnTrainer { let dy = bw_raw_ptr(&self.bw_d_h_s1, &self.stream); let x = bw_raw_ptr(&self.states_buf, &self.stream); let param_sizes = compute_param_sizes(&self.config); - let w_ptrs = bf16_weight_ptrs(&self.params_buf, ¶m_sizes, &self.stream); + let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, ¶m_sizes); let w = w_ptrs[0]; // w_s1 let scratch_base = self.ptrs.iqn_trunk_m; @@ -947,7 +947,7 @@ impl GpuDqnTrainer { // Only upstream gradient (dX) is needed -- skip dW/db for value head. { let param_sizes = compute_param_sizes(&self.config); - let w_ptrs = bf16_weight_ptrs(&self.params_buf, ¶m_sizes, &self.stream); + let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, ¶m_sizes); let w_v2 = w_ptrs[6]; // W_v2 [NA, VH] let dx = self.ptrs.bw_d_h_v; @@ -985,7 +985,7 @@ impl GpuDqnTrainer { // Only upstream gradient (dX) is needed -- skip dW/db for value head. { let param_sizes = compute_param_sizes(&self.config); - let w_ptrs = bf16_weight_ptrs(&self.params_buf, ¶m_sizes, &self.stream); + let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, ¶m_sizes); let w_v1 = w_ptrs[4]; // W_v1 [VH, SH2] let dy = self.ptrs.bw_d_h_v; @@ -1024,7 +1024,7 @@ impl GpuDqnTrainer { let dy = bw_raw_ptr(&self.bw_d_h_s2, &self.stream); let x = bw_raw_ptr(&self.save_h_s1, &self.stream); let param_sizes = compute_param_sizes(&self.config); - let w_ptrs = bf16_weight_ptrs(&self.params_buf, ¶m_sizes, &self.stream); + let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, ¶m_sizes); let w = w_ptrs[2]; // w_s2 let scratch_base = self.ptrs.iqn_trunk_m; @@ -1065,7 +1065,7 @@ impl GpuDqnTrainer { let dy = bw_raw_ptr(&self.bw_d_h_s1, &self.stream); let x = bw_raw_ptr(&self.states_buf, &self.stream); let param_sizes = compute_param_sizes(&self.config); - let w_ptrs = bf16_weight_ptrs(&self.params_buf, ¶m_sizes, &self.stream); + let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, ¶m_sizes); let w = w_ptrs[0]; // w_s1 let scratch_base = self.ptrs.iqn_trunk_m; @@ -1253,7 +1253,7 @@ impl GpuDqnTrainer { // Extract weight pointers for cuBLAS backward let param_sizes = compute_param_sizes(&self.config); - let w_ptrs = bf16_weight_ptrs(&self.params_buf, ¶m_sizes, &self.stream); + let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, ¶m_sizes); // Construct d_adv_logits pointers per branch let f32_size = std::mem::size_of::(); @@ -2271,8 +2271,6 @@ impl GpuDqnTrainer { self.capture_training_graphs(online_dueling, online_branching)?; } self.replay_forward()?; - - // Return placeholder — real scalars come from replay_adam_and_readback() Ok(FusedTrainScalars { total_loss: 0.0, grad_norm: 0.0 }) } @@ -2602,7 +2600,7 @@ impl GpuDqnTrainer { // Writes logits into on_v_logits_buf [batch_size, NA] and on_b_logits_buf [batch_size, (B0+B1+B2)*NA]. // Activation buffers (save_h_*) are used as scratch here (no backward will use them). let param_sizes = compute_param_sizes(&self.config); - let on_w_ptrs = bf16_weight_ptrs(&self.params_buf, ¶m_sizes, &self.stream); + let on_w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, ¶m_sizes); self.cublas_forward.forward_online( &self.stream, @@ -2840,6 +2838,42 @@ impl GpuDqnTrainer { /// /// Steps: zero accumulators → cuBLAS forward → C51 loss → C51 grad → cuBLAS backward. /// After this graph, `grad_buf` contains C51's gradients. + /// Debug: run forward WITHOUT graph capture, with intermediate readbacks. + #[allow(dead_code)] + fn debug_forward_no_graph(&mut self) -> Result<(), MLError> { + let _eg = EventTrackingGuard::new(self.stream.context()); + + // Zero accumulators + self.stream.memset_zeros(&mut self.total_loss_buf).ok(); + self.stream.memset_zeros(&mut self.grad_buf).ok(); + self.stream.memset_zeros(&mut self.d_value_logits_buf).ok(); + self.stream.memset_zeros(&mut self.d_adv_logits_buf).ok(); + + // Check params before forward + unsafe { cudarc::driver::sys::cuStreamSynchronize(self.stream.cu_stream()); } + let mut dbg = vec![half::bf16::ZERO; 8]; + self.stream.memcpy_dtoh(&self.params_buf.slice(..8), &mut dbg).ok(); + eprintln!("[DBG] params[0..8] = {:?}", dbg.iter().map(|v| v.to_f32()).collect::>()); + + // Check states + self.stream.memcpy_dtoh(&self.states_buf.slice(..8), &mut dbg).ok(); + eprintln!("[DBG] states[0..8] = {:?}", dbg.iter().map(|v| v.to_f32()).collect::>()); + + // Forward + self.launch_cublas_forward()?; + unsafe { cudarc::driver::sys::cuStreamSynchronize(self.stream.cu_stream()); } + + // Check h_s1 (first hidden layer output) + self.stream.memcpy_dtoh(&self.save_h_s1.slice(..8), &mut dbg).ok(); + eprintln!("[DBG] h_s1[0..8] = {:?}", dbg.iter().map(|v| v.to_f32()).collect::>()); + + // Check logits + self.stream.memcpy_dtoh(&self.on_v_logits_buf.slice(..8), &mut dbg).ok(); + eprintln!("[DBG] v_logits[0..8] = {:?}", dbg.iter().map(|v| v.to_f32()).collect::>()); + + Ok(()) + } + fn submit_forward_ops(&mut self) -> Result<(), MLError> { // ── Zero accumulators (capturable: memset_zeros uses cuMemsetD32Async) ─ self.stream @@ -3178,10 +3212,10 @@ impl GpuDqnTrainer { fn launch_cublas_forward(&self) -> Result<(), MLError> { let cublas = &self.cublas_forward; - // Compute 20 raw F32 device pointers (online + target) into flat param buffers. + // Compute 20 raw device pointers from CachedPtrs (no device_ptr calls — graph-safe). let param_sizes = compute_param_sizes(&self.config); - let on_w_ptrs = bf16_weight_ptrs(&self.params_buf, ¶m_sizes, &self.stream); - let tg_w_ptrs = bf16_weight_ptrs(&self.target_params_buf, ¶m_sizes, &self.stream); + let on_w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, ¶m_sizes); + let tg_w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.target_params_buf, ¶m_sizes); // ── Pass 1: Online network forward on STATES ────────────────── // Writes into save_h_s1, save_h_s2, save_h_v, save_h_b0/b1/b2, @@ -3558,7 +3592,7 @@ impl GpuDqnTrainer { let bw = &self.cublas_backward; let param_sizes = compute_param_sizes(&self.config); - let w_ptrs = bf16_weight_ptrs(&self.params_buf, ¶m_sizes, &self.stream); + let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, ¶m_sizes); let grad_base = bw_raw_ptr(&self.grad_buf, &self.stream); let states_ptr = bw_raw_ptr(&self.states_buf, &self.stream);