fix(bf16): comprehensive EventTrackingGuard for all post-graph methods

Expert analysis (zen debug): cudarc device_ptr() records stale events
after CUDA graph replay, causing race conditions. Added guards to:
compute_q_values, apply_iqn_trunk_gradient, apply_ensemble_diversity,
apply_cql_gradient, apply_cql_clipped_saxpy, apply_spectral_norm,
target_ema_update.

Smoke test: model trains correctly (Sharpe improves -25→-1),
but Q-value/loss/grad readback still returns garbage due to
cudarc device_ptr corruption in forward_online's raw_bf16_ptr.

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-03-28 16:53:44 +01:00
parent 0873938478
commit fc5deaa965

View File

@@ -725,6 +725,7 @@ impl GpuDqnTrainer {
_online_dueling: &mut DuelingWeightSet,
) -> Result<(), MLError> {
let b = self.config.batch_size;
let _eg = EventTrackingGuard::new(self.stream.context());
let sd = self.config.state_dim;
let sh1 = self.config.shared_h1;
let sh2 = self.config.shared_h2;
@@ -917,6 +918,7 @@ impl GpuDqnTrainer {
scale: f32,
) -> Result<(), MLError> {
let b = self.config.batch_size;
let _eg = EventTrackingGuard::new(self.stream.context());
let sd = self.config.state_dim;
let sh1 = self.config.shared_h1;
let sh2 = self.config.shared_h2;
@@ -1176,6 +1178,7 @@ impl GpuDqnTrainer {
pub fn apply_cql_gradient(
&mut self,
) -> Result<bool, MLError> {
let _eg = EventTrackingGuard::new(self.stream.context());
let cql_kernel = match &self.cql_logit_grad_kernel {
Some(k) => k.clone(),
None => return Ok(false),
@@ -1314,6 +1317,7 @@ impl GpuDqnTrainer {
/// Computes norm of scratch, clips to `cql_budget`, then SAXPYs into grad_buf.
pub fn apply_cql_clipped_saxpy(&mut self, cql_budget: f32) -> Result<(), MLError> {
// Compute CQL gradient norm
let _eg = EventTrackingGuard::new(self.stream.context());
self.stream.memset_zeros(&mut self.grad_norm_buf)
.map_err(|e| MLError::ModelError(format!("zero cql_grad_norm: {e}")))?;
@@ -1372,6 +1376,7 @@ impl GpuDqnTrainer {
online_branching: &mut BranchingWeightSet,
) -> Result<(), MLError> {
let sh1 = self.config.shared_h1 as i32;
let _eg = EventTrackingGuard::new(self.stream.context());
let sh2 = self.config.shared_h2 as i32;
let sd = self.config.state_dim as i32;
let sigma_max = self.config.spectral_norm_sigma_max;
@@ -2204,15 +2209,9 @@ 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;
}
@@ -2241,18 +2240,8 @@ 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;
}
@@ -2590,6 +2579,7 @@ impl GpuDqnTrainer {
states: &CudaSlice<half::bf16>,
batch_size: usize,
) -> Result<&CudaSlice<half::bf16>, MLError> {
let _eg = EventTrackingGuard::new(self.stream.context());
if batch_size > self.config.batch_size {
return Err(MLError::ModelError(format!(
"compute_q_values: batch_size {batch_size} exceeds trainer batch_size {}",
@@ -2621,20 +2611,11 @@ impl GpuDqnTrainer {
{
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.
@@ -2855,25 +2836,16 @@ impl GpuDqnTrainer {
// 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(())
}
@@ -3945,6 +3917,7 @@ impl GpuDqnTrainer {
tau: f32,
) -> Result<(), MLError> {
// Sync stream and clear any stale errors from graph capture phase.
let _eg = EventTrackingGuard::new(self.stream.context());
// cudarc stores errors from cuStreamWaitEvent on disabled events during
// graph capture. check_err() consumes them so bind_to_thread() succeeds.
unsafe { cudarc::driver::sys::cuStreamSynchronize(self.stream.cu_stream()); }