|
|
|
|
@@ -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::<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, ¶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<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, ¶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);
|
|
|
|
|
|