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:
@@ -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()); }
|
||||
|
||||
Reference in New Issue
Block a user