diff --git a/crates/ml/src/cuda_pipeline/batched_forward.rs b/crates/ml/src/cuda_pipeline/batched_forward.rs index 9e398c8a2..0d25ce399 100644 --- a/crates/ml/src/cuda_pipeline/batched_forward.rs +++ b/crates/ml/src/cuda_pipeline/batched_forward.rs @@ -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, param_sizes: &[usize; 20], stream: &Arc, diff --git a/crates/ml/src/cuda_pipeline/gpu_backtest_evaluator.rs b/crates/ml/src/cuda_pipeline/gpu_backtest_evaluator.rs index c35a8edca..31fc152c7 100644 --- a/crates/ml/src/cuda_pipeline/gpu_backtest_evaluator.rs +++ b/crates/ml/src/cuda_pipeline/gpu_backtest_evaluator.rs @@ -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, ¶m_sizes, &self.stream) + super::batched_forward::bf16_weight_ptrs(flat, ¶m_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, ¶m_sizes, &self.stream) + super::batched_forward::bf16_weight_ptrs(flat, ¶m_sizes, &self.stream) }; // Direct kernel launches (no CUDA Graph for the backtest step loop). diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index feb3a8e4f..e95ac4326 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -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, ¶m_sizes, &self.stream); + let w_ptrs = bf16_weight_ptrs(&self.params_buf, ¶m_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, ¶m_sizes, &self.stream); + let w_ptrs = bf16_weight_ptrs(&self.params_buf, ¶m_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, ¶m_sizes, &self.stream); + let w_ptrs = bf16_weight_ptrs(&self.params_buf, ¶m_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, ¶m_sizes, &self.stream); + let w_ptrs = bf16_weight_ptrs(&self.params_buf, ¶m_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, ¶m_sizes, &self.stream); + let w_ptrs = bf16_weight_ptrs(&self.params_buf, ¶m_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, ¶m_sizes, &self.stream); + let w_ptrs = bf16_weight_ptrs(&self.params_buf, ¶m_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, ¶m_sizes, &self.stream); + let w_ptrs = bf16_weight_ptrs(&self.params_buf, ¶m_sizes, &self.stream); // Construct d_adv_logits pointers per branch let f32_size = std::mem::size_of::(); @@ -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 = 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 = dbg.iter().map(|v| v.to_f32()).collect(); + } self.params_initialized = true; } @@ -2231,7 +2241,19 @@ impl GpuDqnTrainer { ) -> Result { // ── 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 = 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 = 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, ¶m_sizes, &self.stream); + let on_w_ptrs = bf16_weight_ptrs(&self.params_buf, ¶m_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 = 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 = 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 = 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 = 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, ¶m_sizes, &self.stream); - let tg_w_ptrs = f32_weight_ptrs(&self.target_params_buf, ¶m_sizes, &self.stream); + 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); // ── 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, ¶m_sizes, &self.stream); + let w_ptrs = bf16_weight_ptrs(&self.params_buf, ¶m_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); diff --git a/crates/ml/src/cuda_pipeline/gpu_experience_collector.rs b/crates/ml/src/cuda_pipeline/gpu_experience_collector.rs index 01f3c0eed..ad6545cd9 100644 --- a/crates/ml/src/cuda_pipeline/gpu_experience_collector.rs +++ b/crates/ml/src/cuda_pipeline/gpu_experience_collector.rs @@ -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;