fix(bf16): graph-safe weight pointers + rename alloc_bf16

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) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-03-28 16:27:52 +01:00
parent b54c825123
commit 954bca690d
2 changed files with 57 additions and 15 deletions

View File

@@ -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 {

View File

@@ -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, &param_sizes, &self.stream);
let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, &param_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, &param_sizes, &self.stream);
let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, &param_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, &param_sizes, &self.stream);
let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, &param_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, &param_sizes, &self.stream);
let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, &param_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, &param_sizes, &self.stream);
let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, &param_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, &param_sizes, &self.stream);
let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, &param_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, &param_sizes, &self.stream);
let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, &param_sizes);
// Construct d_adv_logits pointers per branch
let f32_size = std::mem::size_of::<half::bf16>();
@@ -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, &param_sizes, &self.stream);
let on_w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, &param_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::<Vec<_>>());
// 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::<Vec<_>>());
// 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::<Vec<_>>());
// 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::<Vec<_>>());
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, &param_sizes, &self.stream);
let tg_w_ptrs = bf16_weight_ptrs(&self.target_params_buf, &param_sizes, &self.stream);
let on_w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, &param_sizes);
let tg_w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.target_params_buf, &param_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, &param_sizes, &self.stream);
let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, &param_sizes);
let grad_base = bw_raw_ptr(&self.grad_buf, &self.stream);
let states_ptr = bw_raw_ptr(&self.states_buf, &self.stream);