diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index 6d229bfb9..56f81c27a 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -7237,89 +7237,14 @@ impl GpuDqnTrainer { pub fn reduce_current_q_stats(&mut self) -> Result { let _eg = EventTrackingGuard::new(self.stream.context()); - // Populate q_out_buf from the training forward pass logits. - let batch_size = self.config.batch_size; - let na = self.config.num_atoms as i32; - let b0 = self.config.branch_0_size as i32; - let b1 = self.config.branch_1_size as i32; - let b2 = self.config.branch_2_size as i32; - let b3 = self.config.branch_3_size as i32; - let n = batch_size as i32; + // q_out_buf is already populated by the graph replay (populate_q_out in aux_child). + // All temporal/ISV/post-processing ops already ran inside the graph. + // Just reduce the existing Q-values to stats — no kernel re-execution. + let n = self.config.batch_size as i32; let total_actions = self.total_actions() as i32; - let eq_blocks = ((batch_size as u32 + 255) / 256).max(1); - let on_v_ptr = self.ptrs.on_v_logits_buf; - let on_b_ptr = self.ptrs.on_b_logits_buf; - let q_out_ptr = self.q_out_buf.raw_ptr(); - self.stream.memset_zeros(&mut self.atom_stats_buf) - .map_err(|e| MLError::ModelError(format!("atom_stats_buf memset: {e}")))?; - let atom_stats_ptr = self.atom_stats_buf.raw_ptr(); - let support_ptr = self.per_sample_support_ptr; - let q_var_ptr = self.q_var_buf_trainer.raw_ptr(); + let num_atoms = self.config.num_atoms as i32; - // All temporal/ISV ops already ran inside the graph replay. - // Logit buffers (on_v_logits, on_b_logits) are current — just compute Q-stats. - - unsafe { - self.stream - .launch_builder(&self.expected_q_kernel) - .arg(&on_v_ptr) - .arg(&on_b_ptr) - .arg(&q_out_ptr) - .arg(&n) - .arg(&na) - .arg(&b0) - .arg(&b1) - .arg(&b2) - .arg(&b3) - .arg(&support_ptr) - .arg(&atom_stats_ptr) - .arg(&q_var_ptr) - .arg(&self.atom_positions_buf) - .launch(LaunchConfig { - grid_dim: (eq_blocks, 1, 1), - block_dim: (256, 1, 1), - shared_mem_bytes: 0, - }) - .map_err(|e| MLError::ModelError(format!("compute_expected_q (q_stats): {e}")))?; - } - - // Q-mean centering: removes bootstrapping drift before attention and stats. - self.launch_q_mean_center(batch_size)?; - - // NOTE: update_q_mean_ema moved after stream sync (line ~6371) to avoid - // reading pinned memory before GPU kernel finishes writing. - - // G11: Anchor Q-values to Flat/neutral actions per branch - self.apply_q_anchoring(batch_size)?; - - // Apply risk budget: scale magnitude Q-values, produce CVaR alpha + commitment lambda - self.apply_risk_budget(batch_size)?; - - // Branch confidence routing: ISV gate × Q-value separation confidence - self.apply_branch_confidence_routing(batch_size)?; - - // G5: Epistemic gate — high variance → conservative magnitude - self.apply_epistemic_gate(batch_size)?; - - // G6 + G10: informational penalties — computed after branches and expected_q - self.compute_temporal_consistency(batch_size)?; - self.compute_branch_independence(batch_size)?; - - // Cross-branch Q attention: q_out_buf → q_coord_buf [B, 12]. - // Runs after compute_expected_q so attention weights train on real Q-values. - self.launch_q_attention(batch_size)?; - - // Graph message passing: structural coordination across branches (in-place on q_coord_buf). - self.launch_graph_message_pass(batch_size)?; - - // Snapshot pre-denoise Q-values for backward pass (before denoiser modifies q_coord_buf). - self.snapshot_pre_denoise_q(batch_size)?; - - // Diffusion Q-refinement: 2-step denoiser conditioned on Var[Q]. - self.launch_q_denoise(batch_size)?; - - // Stats reduction on q_out_buf. - let num_atoms = na; + // Stats reduction on the already-computed q_out_buf. unsafe { self.stream .launch_builder(&self.q_stats_kernel) diff --git a/crates/ml/src/trainers/dqn/trainer/training_loop.rs b/crates/ml/src/trainers/dqn/trainer/training_loop.rs index 87a084d89..c5615848c 100644 --- a/crates/ml/src/trainers/dqn/trainer/training_loop.rs +++ b/crates/ml/src/trainers/dqn/trainer/training_loop.rs @@ -1565,10 +1565,9 @@ impl DQNTrainer { } // end guard frequency check guard_total_us += guard_start.elapsed().as_micros() as u64; - // Q-stats + v_range adaptation every 50 steps. - // Cost: ~4s per call (15+ kernels + cuStreamSynchronize for DtoH readback). - // Running every step was a 700s/epoch bottleneck. - if train_step_count % 50 == 0 { + // Q-stats + v_range adaptation every step. + // Cost: ~1 kernel (q_stats_reduce) + sync for 76B DtoH readback. + // q_out_buf is already populated by graph replay — no re-computation. if let Some(ref mut fused) = self.fused_ctx { if let Ok(stats) = fused.reduce_current_q_stats() { self.cached_avg_q = stats.q_mean as f64; @@ -1587,7 +1586,6 @@ impl DQNTrainer { fused.update_eval_v_range(stats.q_mean, q_std, q_gap, per_branch_q_gaps); } } - } // end Q-stats frequency check train_step_count += 1; self.gradient_logging_step += 1; }