refactor: split forward ops into per-stream functions
submit_forward_ops_main: Pass 1 + 2 + loss + grad + backward (main stream) submit_forward_ops_ddqn: Pass 3 Double DQN (double_dqn_stream) No set_stream() or event ops inside either function. Graph capture will record each on its own stream. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -4370,6 +4370,128 @@ impl GpuDqnTrainer {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Submit main-stream forward ops for graph capture (Pass 1 + 2 + loss + grad + backward).
|
||||
///
|
||||
/// Contains everything that runs on `self.stream`:
|
||||
/// - Zero accumulators
|
||||
/// - Pass 1: online forward on states + stochastic depth
|
||||
/// - Pass 2: target forward on next_states
|
||||
/// - Curiosity inference
|
||||
/// - MSE + C51 loss/gradient blend
|
||||
/// - bf16 d_logits cast
|
||||
/// - cuBLAS backward
|
||||
///
|
||||
/// Does NOT contain Pass 3 (Double DQN) or any event/set_stream ops.
|
||||
/// Pass 3 is submitted separately via `submit_forward_ops_ddqn()`.
|
||||
pub(crate) fn submit_forward_ops_main(&mut self) -> Result<(), MLError> {
|
||||
// ── Zero accumulators (capturable: memset_zeros uses cuMemsetD32Async) ─
|
||||
self.stream
|
||||
.memset_zeros(&mut self.total_loss_buf)
|
||||
.map_err(|e| MLError::ModelError(format!("zero total_loss: {e}")))?;
|
||||
self.stream
|
||||
.memset_zeros(&mut self.mse_loss_buf)
|
||||
.map_err(|e| MLError::ModelError(format!("zero mse_loss: {e}")))?;
|
||||
self.stream
|
||||
.memset_zeros(&mut self.grad_buf)
|
||||
.map_err(|e| MLError::ModelError(format!("zero grad_buf: {e}")))?;
|
||||
// Zero C51 gradient output buffers (atomicAdd accumulates)
|
||||
self.stream
|
||||
.memset_zeros(&mut self.d_value_logits_buf)
|
||||
.map_err(|e| MLError::ModelError(format!("zero d_value_logits: {e}")))?;
|
||||
self.stream
|
||||
.memset_zeros(&mut self.d_adv_logits_buf)
|
||||
.map_err(|e| MLError::ModelError(format!("zero d_adv_logits: {e}")))?;
|
||||
|
||||
// ── 1. Forward (cuBLAS SGEMM — Pass 1 + 2, no Pass 3) ──────────
|
||||
self.launch_cublas_forward()?;
|
||||
|
||||
// ── 2+3. Loss + gradient (blended MSE + C51 via c51_alpha ramp) ─
|
||||
self.launch_curiosity_inference()?;
|
||||
|
||||
// MSE path → scratch buffers
|
||||
self.stream.memset_zeros(&mut self.d_value_logits_mse)
|
||||
.map_err(|e| MLError::ModelError(format!("zero d_value_mse: {e}")))?;
|
||||
self.stream.memset_zeros(&mut self.d_adv_logits_mse)
|
||||
.map_err(|e| MLError::ModelError(format!("zero d_adv_mse: {e}")))?;
|
||||
self.launch_mse_loss()?;
|
||||
self.launch_mse_grad_to_scratch()?;
|
||||
|
||||
// C51 path → main buffers (already zeroed above)
|
||||
self.stream.memset_zeros(&mut self.mixup_barrier_buf)
|
||||
.map_err(|e| MLError::ModelError(format!("zero mixup_barrier: {e}")))?;
|
||||
self.launch_c51_loss()?;
|
||||
self.launch_c51_grad()?;
|
||||
|
||||
// Blend: main = α * C51 + (1-α) * MSE
|
||||
{
|
||||
let alpha = self.c51_alpha;
|
||||
let na = self.config.num_atoms;
|
||||
let b = self.config.batch_size;
|
||||
let b0 = self.config.branch_0_size;
|
||||
let b1 = self.config.branch_1_size;
|
||||
let b2 = self.config.branch_2_size;
|
||||
let n_val = (b * na) as i32;
|
||||
let n_adv = (b * (b0 + b1 + b2) * na) as i32;
|
||||
let scale_mse = 1.0 - alpha;
|
||||
let val_blocks = ((b * na + 255) / 256) as u32;
|
||||
let adv_blocks = ((b * (b0 + b1 + b2) * na + 255) / 256) as u32;
|
||||
let cfg_val = LaunchConfig { grid_dim: (val_blocks, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
|
||||
let cfg_adv = LaunchConfig { grid_dim: (adv_blocks, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
|
||||
|
||||
let val_ptr = self.d_value_logits_buf.raw_ptr();
|
||||
let val_mse_ptr = self.d_value_logits_mse.raw_ptr();
|
||||
let adv_ptr = self.d_adv_logits_buf.raw_ptr();
|
||||
let adv_mse_ptr = self.d_adv_logits_mse.raw_ptr();
|
||||
|
||||
unsafe {
|
||||
self.stream.launch_builder(&self.scale_f32_kernel)
|
||||
.arg(&val_ptr)
|
||||
.arg(&alpha).arg(&n_val)
|
||||
.launch(cfg_val).map_err(|e| MLError::ModelError(format!("scale c51 val: {e}")))?;
|
||||
self.stream.launch_builder(&self.saxpy_f32_kernel)
|
||||
.arg(&val_ptr).arg(&val_mse_ptr)
|
||||
.arg(&scale_mse).arg(&n_val)
|
||||
.launch(cfg_val).map_err(|e| MLError::ModelError(format!("blend mse val: {e}")))?;
|
||||
self.stream.launch_builder(&self.scale_f32_kernel)
|
||||
.arg(&adv_ptr)
|
||||
.arg(&alpha).arg(&n_adv)
|
||||
.launch(cfg_adv).map_err(|e| MLError::ModelError(format!("scale c51 adv: {e}")))?;
|
||||
self.stream.launch_builder(&self.saxpy_f32_kernel)
|
||||
.arg(&adv_ptr).arg(&adv_mse_ptr)
|
||||
.arg(&scale_mse).arg(&n_adv)
|
||||
.launch(cfg_adv).map_err(|e| MLError::ModelError(format!("blend mse adv: {e}")))?;
|
||||
}
|
||||
}
|
||||
|
||||
// ── 3.5. Cast f32 d_logits → bf16 staging for cuBLAS backward ─
|
||||
self.cast_d_logits_to_bf16()?;
|
||||
|
||||
// ── 4. Backward (cuBLAS SGEMM, chain rule through layers) ─
|
||||
self.launch_cublas_backward()?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Submit Pass 3 (Double DQN online forward on next_states) for
|
||||
/// the double_dqn_stream graph. Uses its own CublasForward handle.
|
||||
///
|
||||
/// No set_stream() or event ops — this function is captured entirely
|
||||
/// on `self.double_dqn_stream` via `self.cublas_forward_ddqn`.
|
||||
fn submit_forward_ops_ddqn(&mut self) -> Result<(), MLError> {
|
||||
let param_sizes = compute_param_sizes(&self.config);
|
||||
let on_w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, ¶m_sizes);
|
||||
|
||||
self.cublas_forward_ddqn.forward_online_raw(
|
||||
&self.double_dqn_stream, self.ptrs.next_states_buf, &on_w_ptrs,
|
||||
self.ptrs.on_next_h_s1_scratch, self.ptrs.on_next_h_s2_scratch,
|
||||
self.ptrs.on_next_h_v_scratch,
|
||||
self.ptrs.on_next_h_b_scratch, self.ptrs.on_next_h_b_scratch,
|
||||
self.ptrs.on_next_h_b_scratch,
|
||||
self.ptrs.on_next_v_logits_buf, self.ptrs.on_next_b_logits_buf,
|
||||
)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Submit MSE-only forward ops — no C51 loss/grad/blend.
|
||||
///
|
||||
/// Used during MSE warmup (c51_alpha ≈ 0). Eliminates C51 loss, C51 grad,
|
||||
@@ -4696,18 +4818,9 @@ impl GpuDqnTrainer {
|
||||
}
|
||||
}
|
||||
|
||||
// ── Pass 1 complete: record event so double_dqn_stream can start ──
|
||||
// Pass 2 (target fwd) and Pass 3 (online fwd on next_states for Double
|
||||
// DQN) are independent — they read the same next_states_buf but write to
|
||||
// completely separate output buffers. Run them concurrently on 2 streams.
|
||||
self.pass1_event.record(&self.stream)
|
||||
.map_err(|e| MLError::ModelError(format!("pass1 event record: {e}")))?;
|
||||
|
||||
// double_dqn_stream waits for Pass 1 + stochastic depth to finish.
|
||||
self.double_dqn_stream.wait(&self.pass1_event)
|
||||
.map_err(|e| MLError::ModelError(format!("double_dqn wait pass1: {e}")))?;
|
||||
|
||||
// ── Pass 2: Target forward on NEXT_STATES — main stream (graph-safe)
|
||||
// No event sync needed — Pass 3 (Double DQN) is submitted separately
|
||||
// via submit_forward_ops_ddqn() and captured on its own stream/graph.
|
||||
cublas.forward_target_raw(
|
||||
&self.stream, self.ptrs.next_states_buf, &tg_w_ptrs,
|
||||
self.ptrs.tg_h_s1_scratch, self.ptrs.tg_h_s2_buf,
|
||||
@@ -4715,28 +4828,6 @@ impl GpuDqnTrainer {
|
||||
self.ptrs.tg_v_logits_buf, self.ptrs.tg_b_logits_buf,
|
||||
)?;
|
||||
|
||||
// ── Pass 3: Online forward on NEXT_STATES — double_dqn_stream (graph-safe)
|
||||
// Redirect cuBLAS handle to the double_dqn_stream so GEMMs execute there.
|
||||
// The bias/relu kernel launches already use the stream parameter passed to
|
||||
// forward_online_raw, so they will run on double_dqn_stream automatically.
|
||||
cublas.set_stream(&self.double_dqn_stream)?;
|
||||
cublas.forward_online_raw(
|
||||
&self.double_dqn_stream, self.ptrs.next_states_buf, &on_w_ptrs,
|
||||
self.ptrs.on_next_h_s1_scratch, self.ptrs.on_next_h_s2_scratch,
|
||||
self.ptrs.on_next_h_v_scratch,
|
||||
self.ptrs.on_next_h_b_scratch, self.ptrs.on_next_h_b_scratch,
|
||||
self.ptrs.on_next_h_b_scratch, // reuse for all 3 branches (scratch)
|
||||
self.ptrs.on_next_v_logits_buf, self.ptrs.on_next_b_logits_buf,
|
||||
)?;
|
||||
// Restore cuBLAS handle to the main stream.
|
||||
cublas.set_stream(&self.stream)?;
|
||||
|
||||
// ── Join: main stream waits for Pass 3 to complete before loss kernels ──
|
||||
self.pass3_event.record(&self.double_dqn_stream)
|
||||
.map_err(|e| MLError::ModelError(format!("pass3 event record: {e}")))?;
|
||||
self.stream.wait(&self.pass3_event)
|
||||
.map_err(|e| MLError::ModelError(format!("join pass3: {e}")))?;
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user