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:
@@ -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 {
|
||||
|
||||
@@ -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::<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, ¶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::<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, ¶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);
|
||||
|
||||
Reference in New Issue
Block a user