refactor: rename f32_weight_ptrs → bf16_weight_ptrs for consistency

All weight pointer computation uses size_of::<half::bf16>() — the
function name should reflect the actual type. Renamed across
batched_forward.rs, gpu_dqn_trainer.rs, gpu_experience_collector.rs.

Smoke test debug: params_buf valid after flatten, NaN after first
graph_adam replay. Root cause: either Adam produces NaN from the
first backward's gradients, or the C51/MSE loss kernel produces
NaN from valid logits. Need to check first-iteration loss output.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-03-28 16:09:28 +01:00
parent 5d78108c77
commit c858f84ada
4 changed files with 59 additions and 17 deletions

View File

@@ -619,7 +619,7 @@ fn compile_bias_kernels(
/// w_b2fc, b_b2fc, w_b2out, b_b2out]
///
/// Returns 20 raw u64 device pointers (base + byte_offset for each tensor).
pub fn f32_weight_ptrs(
pub fn bf16_weight_ptrs(
params_buf: &CudaSlice<half::bf16>,
param_sizes: &[usize; 20],
stream: &Arc<CudaStream>,

View File

@@ -820,7 +820,7 @@ impl GpuBacktestEvaluator {
let flat = self.cublas_params_flat.as_ref().ok_or_else(|| {
MLError::ModelError("cublas_params_flat unexpectedly None".to_owned())
})?;
super::batched_forward::f32_weight_ptrs(flat, &param_sizes, &self.stream)
super::batched_forward::bf16_weight_ptrs(flat, &param_sizes, &self.stream)
};
self.submit_dqn_step_loop_cublas(&w_ptrs, dqn_cfg)?;
@@ -884,7 +884,7 @@ impl GpuBacktestEvaluator {
let flat = self.cublas_params_flat.as_ref().ok_or_else(|| {
MLError::ModelError("cublas_params_flat unexpectedly None".to_owned())
})?;
super::batched_forward::f32_weight_ptrs(flat, &param_sizes, &self.stream)
super::batched_forward::bf16_weight_ptrs(flat, &param_sizes, &self.stream)
};
// Direct kernel launches (no CUDA Graph for the backtest step loop).

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, f32_weight_ptrs};
use super::batched_forward::{CublasForward, bf16_weight_ptrs};
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 = f32_weight_ptrs(&self.params_buf, &param_sizes, &self.stream);
let w_ptrs = bf16_weight_ptrs(&self.params_buf, &param_sizes, &self.stream);
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 = f32_weight_ptrs(&self.params_buf, &param_sizes, &self.stream);
let w_ptrs = bf16_weight_ptrs(&self.params_buf, &param_sizes, &self.stream);
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 = f32_weight_ptrs(&self.params_buf, &param_sizes, &self.stream);
let w_ptrs = bf16_weight_ptrs(&self.params_buf, &param_sizes, &self.stream);
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 = f32_weight_ptrs(&self.params_buf, &param_sizes, &self.stream);
let w_ptrs = bf16_weight_ptrs(&self.params_buf, &param_sizes, &self.stream);
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 = f32_weight_ptrs(&self.params_buf, &param_sizes, &self.stream);
let w_ptrs = bf16_weight_ptrs(&self.params_buf, &param_sizes, &self.stream);
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 = f32_weight_ptrs(&self.params_buf, &param_sizes, &self.stream);
let w_ptrs = bf16_weight_ptrs(&self.params_buf, &param_sizes, &self.stream);
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 = f32_weight_ptrs(&self.params_buf, &param_sizes, &self.stream);
let w_ptrs = bf16_weight_ptrs(&self.params_buf, &param_sizes, &self.stream);
// Construct d_adv_logits pointers per branch
let f32_size = std::mem::size_of::<half::bf16>();
@@ -2203,7 +2203,17 @@ impl GpuDqnTrainer {
// ── First call: flatten weights ──────────────────────────────
if !self.params_initialized {
{
let mut dbg = vec![half::bf16::ZERO; 8];
self.stream.memcpy_dtoh(&online_dueling.w_s1.slice(..8), &mut dbg).ok();
let w: Vec<f32> = dbg.iter().map(|v| v.to_f32()).collect();
}
self.flatten_online_weights(online_dueling, online_branching)?;
{
let mut dbg = vec![half::bf16::ZERO; 8];
self.stream.memcpy_dtoh(&self.params_buf.slice(..8), &mut dbg).ok();
let w: Vec<f32> = dbg.iter().map(|v| v.to_f32()).collect();
}
self.params_initialized = true;
}
@@ -2231,7 +2241,19 @@ impl GpuDqnTrainer {
) -> Result<FusedTrainScalars, MLError> {
// ── First call: flatten weights ──────────────────────────────
if !self.params_initialized {
{
let _eg = EventTrackingGuard::new(self.stream.context());
let mut dbg = vec![half::bf16::ZERO; 8];
self.stream.memcpy_dtoh(&online_dueling.w_s1.slice(..8), &mut dbg).ok();
let w: Vec<f32> = dbg.iter().map(|v| v.to_f32()).collect();
}
self.flatten_online_weights(online_dueling, online_branching)?;
{
let _eg = EventTrackingGuard::new(self.stream.context());
let mut dbg = vec![half::bf16::ZERO; 8];
self.stream.memcpy_dtoh(&self.params_buf.slice(..8), &mut dbg).ok();
let w: Vec<f32> = dbg.iter().map(|v| v.to_f32()).collect();
}
self.params_initialized = true;
}
@@ -2580,7 +2602,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 = f32_weight_ptrs(&self.params_buf, &param_sizes, &self.stream);
let on_w_ptrs = bf16_weight_ptrs(&self.params_buf, &param_sizes, &self.stream);
self.cublas_forward.forward_online(
&self.stream,
@@ -2596,6 +2618,26 @@ impl GpuDqnTrainer {
&self.on_b_logits_buf,
)?;
// Debug: check params before and after graph capture
{
let _eg = EventTrackingGuard::new(self.stream.context());
unsafe { cudarc::driver::sys::cuStreamSynchronize(self.stream.cu_stream()); }
let mut dbg = vec![half::bf16::ZERO; 8];
// Check states
self.stream.memcpy_dtoh(&self.states_buf.slice(..8), &mut dbg).ok();
let s: Vec<f32> = dbg.iter().map(|v| v.to_f32()).collect();
// Check params (first 8 weights)
self.stream.memcpy_dtoh(&self.params_buf.slice(..8), &mut dbg).ok();
let w: Vec<f32> = dbg.iter().map(|v| v.to_f32()).collect();
// Check h_s1 (first hidden activation — output of first GEMM + bias + relu)
self.stream.memcpy_dtoh(&self.save_h_s1.slice(..8), &mut dbg).ok();
let h: Vec<f32> = dbg.iter().map(|v| v.to_f32()).collect();
// Check total_params
// Check logits
self.stream.memcpy_dtoh(&self.on_v_logits_buf.slice(..8), &mut dbg).ok();
let l: Vec<f32> = dbg.iter().map(|v| v.to_f32()).collect();
}
// Step 2: compute_expected_q kernel — logits → expected Q-values.
// Writes into q_out_buf [batch_size, total_actions].
let n = batch_size as i32;
@@ -3138,8 +3180,8 @@ impl GpuDqnTrainer {
// Compute 20 raw F32 device pointers (online + target) into flat param buffers.
let param_sizes = compute_param_sizes(&self.config);
let on_w_ptrs = f32_weight_ptrs(&self.params_buf, &param_sizes, &self.stream);
let tg_w_ptrs = f32_weight_ptrs(&self.target_params_buf, &param_sizes, &self.stream);
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);
// ── Pass 1: Online network forward on STATES ──────────────────
// Writes into save_h_s1, save_h_s2, save_h_v, save_h_b0/b1/b2,
@@ -3516,7 +3558,7 @@ impl GpuDqnTrainer {
let bw = &self.cublas_backward;
let param_sizes = compute_param_sizes(&self.config);
let w_ptrs = f32_weight_ptrs(&self.params_buf, &param_sizes, &self.stream);
let w_ptrs = bf16_weight_ptrs(&self.params_buf, &param_sizes, &self.stream);
let grad_base = bw_raw_ptr(&self.grad_buf, &self.stream);
let states_ptr = bw_raw_ptr(&self.states_buf, &self.stream);

View File

@@ -25,7 +25,7 @@ use ml_core::cuda_autograd::GpuVarStore;
use tracing::{debug, info};
use crate::MLError;
use super::batched_forward::{CublasForward, f32_weight_ptrs};
use super::batched_forward::{CublasForward, bf16_weight_ptrs};
use super::gpu_curiosity_trainer::GpuCuriosityTrainer;
use super::gpu_dqn_trainer::{
GpuDqnTrainConfig, compute_param_sizes, compute_total_params, dtod_copy, raw_device_ptr,
@@ -1164,7 +1164,7 @@ impl GpuExperienceCollector {
.map_err(|e| MLError::ModelError(format!("memset current_timesteps: {e}")))?;
// ── Compute weight pointers for cuBLAS ──────────────────────────
let w_ptrs = f32_weight_ptrs(&self.online_params_flat, &self.param_sizes, &self.stream);
let w_ptrs = bf16_weight_ptrs(&self.online_params_flat, &self.param_sizes, &self.stream);
let n = n_episodes;
let blocks_256 = ((n + 255) / 256) as u32;