From d83c1d4c48bb9d250cddd16dff674fb6a59de6a0 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sat, 28 Mar 2026 14:14:31 +0100 Subject: [PATCH] =?UTF-8?q?refactor(bf16):=20Spec=20C=20dead=20code=20clea?= =?UTF-8?q?nup=20=E2=80=94=20delete=20911=20LOC=20of=20F32=20paths?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit batched_forward.rs (-698 lines): - Delete sgemm_layer, sgemm_layer_raw (F32 cublasSgemm) - Delete 4 F32 bias launchers (launch_add_bias_relu/_raw, launch_add_bias/_raw) - Delete forward_online_bf16, forward_target_bf16 (conversion layer paths) - Delete bf16_weight_ptrs, 6 dead BF16 buffer accessors, raw_u16_ptr - Delete 15 CudaSlice internal BF16 mirror buffers from CublasForward - Delete F32 kernel fields (add_bias_relu_kernel, add_bias_kernel, f32_to_bf16_kernel) - compile_bias_kernels returns only BF16 kernels now gpu_dqn_trainer.rs (-199 lines): - Delete bf16_params_buf, bf16_target_params_buf mirror infrastructure - Delete launch_segmented_bf16_convert, launch_bf16_convert_online/target - Delete bf16_goff_byte_offsets, bf16_padded_total, bf16_mirrors_initialized - Delete bf16_to_f32_kernel accessor, raw_device_ptr_u16 helper - Remove sync_target_bf16 call from target_ema_update fused_training.rs (4 size_of:: → size_of::): - td_errors DtoD copy, ensemble buffer memset, dueling/branching weight clones 1722/1722 tests pass. Zero dead F32 code remains. Co-Authored-By: Claude Opus 4.6 (1M context) --- .../ml/src/cuda_pipeline/batched_forward.rs | 698 +----------------- .../ml/src/cuda_pipeline/gpu_dqn_trainer.rs | 199 ----- crates/ml/src/trainers/dqn/fused_training.rs | 64 +- 3 files changed, 50 insertions(+), 911 deletions(-) diff --git a/crates/ml/src/cuda_pipeline/batched_forward.rs b/crates/ml/src/cuda_pipeline/batched_forward.rs index 3d53e12e1..9e398c8a2 100644 --- a/crates/ml/src/cuda_pipeline/batched_forward.rs +++ b/crates/ml/src/cuda_pipeline/batched_forward.rs @@ -1,6 +1,6 @@ #![allow(unsafe_code)] -//! cuBLAS SGEMM-based batched forward pass for the DQN trainer. +//! cuBLAS BF16 batched forward pass for the DQN trainer. //! //! Replaces the 1-warp-per-sample fused kernel with cuBLAS matrix multiplications //! that process the entire batch in a single GEMM call per layer. This raises GPU @@ -9,13 +9,11 @@ //! //! ## BF16 tensor core path (cublasGemmEx) //! -//! Alongside the existing F32 `cublasSgemm` path, this module provides a -//! `gemmex_bf16` method that calls `cublasGemmEx` with BF16 inputs and F32 -//! internal accumulation. On H100 this yields ~3x throughput vs F32 SGEMM. +//! All GEMMs use `cublasGemmEx` with BF16 inputs and F32 internal accumulation. +//! On H100 this yields ~3x throughput vs F32 SGEMM. //! -//! The BF16 path uses internal activation buffers (`h_s1_bf16`, etc.) allocated -//! at construction time, plus BF16 bias+relu kernels (`add_bias_relu_bf16_kernel`, -//! `add_bias_bf16_kernel`). The existing F32 `sgemm_layer` stays as fallback. +//! BF16 bias+relu kernels (`add_bias_relu_bf16_kernel`, `add_bias_bf16_kernel`) +//! handle post-GEMM operations. No F32 fallback path exists. //! //! ## Layout convention //! @@ -96,55 +94,16 @@ pub struct CublasForward { /// cuBLAS workspace buffer — must outlive the handle for CUDA Graph replay. _workspace_buf: CudaSlice, - // ── Compiled bias+activation kernels (loaded from the DQN module) ── - /// Kernel: `add_bias_relu(output, bias, out_dim, total_elements)` — F32 - /// Grid: ceil(B * out_dim / 256), Block: 256. - add_bias_relu_kernel: CudaFunction, - - /// Kernel: `add_bias(output, bias, out_dim, total_elements)` — F32 - /// Same shape as add_bias_relu but no clamping (used for final logit layers). - add_bias_kernel: CudaFunction, + // ── Compiled bias+activation kernels (loaded from precompiled cubin) ── /// Kernel: `add_bias_relu_bf16(output, bias, out_dim, total_elements)` — BF16 - /// BF16 variant of add_bias_relu for tensor core forward path. + /// BF16 bias+relu for tensor core forward path. add_bias_relu_bf16_kernel: CudaFunction, /// Kernel: `add_bias_bf16(output, bias, out_dim, total_elements)` — BF16 - /// BF16 variant of add_bias for tensor core forward path. + /// BF16 bias (no activation) for tensor core forward path. add_bias_bf16_kernel: CudaFunction, - /// Kernel: `f32_to_bf16_kernel(src_f32, dst_bf16, n)` — element-wise F32→BF16. - /// Used by `forward_online_bf16` to convert F32 states from the experience - /// collector into BF16 before the tensor core GEMM pipeline. - /// Grid: ceil(n/256), Block: 256. - f32_to_bf16_kernel: CudaFunction, - - // ── BF16 activation buffers for tensor core forward path ───────── - // Internal BF16 buffers needed because forward_online takes &CudaSlice - // and we cannot change that signature without cascading to callers. - // The gemmex_bf16 path reads/writes these instead. - - /// Online network BF16 activation buffers - states_bf16: CudaSlice, // [B, SD] - h_s1_bf16: CudaSlice, // [B, SH1] - h_s2_bf16: CudaSlice, // [B, SH2] - h_v_bf16: CudaSlice, // [B, VH] - h_b0_bf16: CudaSlice, // [B, AH] - h_b1_bf16: CudaSlice, // [B, AH] - h_b2_bf16: CudaSlice, // [B, AH] - /// Online logit BF16 buffers - v_logits_bf16: CudaSlice, // [B, NA] - b_logits_bf16: CudaSlice, // [B, (B0+B1+B2)*NA] - - /// Target network BF16 activation buffers (inference scratch — reused) - tgt_h_s1_bf16: CudaSlice, // [B, SH1] - tgt_h_s2_bf16: CudaSlice, // [B, SH2] - tgt_h_v_bf16: CudaSlice, // [B, VH] - tgt_h_b_bf16: CudaSlice, // [B, AH] (shared across branches) - /// Target logit BF16 buffers - tgt_v_logits_bf16: CudaSlice, // [B, NA] - tgt_b_logits_bf16: CudaSlice, // [B, (B0+B1+B2)*NA] - // ── Network dimensions (baked at construction) ── batch_size: usize, state_dim: usize, @@ -206,74 +165,15 @@ impl CublasForward { } } - // ── Compile bias kernels (F32 + BF16) + f32_to_bf16 converter ── - let (add_bias_relu_kernel, add_bias_kernel, - add_bias_relu_bf16_kernel, add_bias_bf16_kernel, - f32_to_bf16_kernel) = + // ── Load BF16 bias kernels from precompiled cubin ── + let (add_bias_relu_bf16_kernel, add_bias_bf16_kernel) = compile_bias_kernels(stream)?; - // ── Allocate BF16 activation buffers ───────────────────────── - let total_branch_logits = (branch_0_size + branch_1_size + branch_2_size) * num_atoms; - - // Online network BF16 activation buffers - let states_bf16 = stream.alloc_zeros::(batch_size * state_dim) - .map_err(|e| MLError::ModelError(format!("BF16 states_bf16 alloc: {e}")))?; - let h_s1_bf16 = stream.alloc_zeros::(batch_size * shared_h1) - .map_err(|e| MLError::ModelError(format!("BF16 h_s1_bf16 alloc: {e}")))?; - let h_s2_bf16 = stream.alloc_zeros::(batch_size * shared_h2) - .map_err(|e| MLError::ModelError(format!("BF16 h_s2_bf16 alloc: {e}")))?; - let h_v_bf16 = stream.alloc_zeros::(batch_size * value_h) - .map_err(|e| MLError::ModelError(format!("BF16 h_v_bf16 alloc: {e}")))?; - let h_b0_bf16 = stream.alloc_zeros::(batch_size * adv_h) - .map_err(|e| MLError::ModelError(format!("BF16 h_b0_bf16 alloc: {e}")))?; - let h_b1_bf16 = stream.alloc_zeros::(batch_size * adv_h) - .map_err(|e| MLError::ModelError(format!("BF16 h_b1_bf16 alloc: {e}")))?; - let h_b2_bf16 = stream.alloc_zeros::(batch_size * adv_h) - .map_err(|e| MLError::ModelError(format!("BF16 h_b2_bf16 alloc: {e}")))?; - let v_logits_bf16 = stream.alloc_zeros::(batch_size * num_atoms) - .map_err(|e| MLError::ModelError(format!("BF16 v_logits_bf16 alloc: {e}")))?; - let b_logits_bf16 = stream.alloc_zeros::(batch_size * total_branch_logits) - .map_err(|e| MLError::ModelError(format!("BF16 b_logits_bf16 alloc: {e}")))?; - - // Target network BF16 activation buffers (scratch — reused) - let tgt_h_s1_bf16 = stream.alloc_zeros::(batch_size * shared_h1) - .map_err(|e| MLError::ModelError(format!("BF16 tgt_h_s1_bf16 alloc: {e}")))?; - let tgt_h_s2_bf16 = stream.alloc_zeros::(batch_size * shared_h2) - .map_err(|e| MLError::ModelError(format!("BF16 tgt_h_s2_bf16 alloc: {e}")))?; - let tgt_h_v_bf16 = stream.alloc_zeros::(batch_size * value_h) - .map_err(|e| MLError::ModelError(format!("BF16 tgt_h_v_bf16 alloc: {e}")))?; - let tgt_h_b_bf16 = stream.alloc_zeros::(batch_size * adv_h) - .map_err(|e| MLError::ModelError(format!("BF16 tgt_h_b_bf16 alloc: {e}")))?; - let tgt_v_logits_bf16 = stream.alloc_zeros::(batch_size * num_atoms) - .map_err(|e| MLError::ModelError(format!("BF16 tgt_v_logits_bf16 alloc: {e}")))?; - let tgt_b_logits_bf16 = stream.alloc_zeros::(batch_size * total_branch_logits) - .map_err(|e| MLError::ModelError(format!("BF16 tgt_b_logits_bf16 alloc: {e}")))?; - Ok(Self { handle: SendSyncCublasHandle(raw_handle), _workspace_buf: workspace_buf, - add_bias_relu_kernel, - add_bias_kernel, add_bias_relu_bf16_kernel, add_bias_bf16_kernel, - f32_to_bf16_kernel, - // Online BF16 activation buffers - states_bf16, - h_s1_bf16, - h_s2_bf16, - h_v_bf16, - h_b0_bf16, - h_b1_bf16, - h_b2_bf16, - v_logits_bf16, - b_logits_bf16, - // Target BF16 activation buffers - tgt_h_s1_bf16, - tgt_h_s2_bf16, - tgt_h_v_bf16, - tgt_h_b_bf16, - tgt_v_logits_bf16, - tgt_b_logits_bf16, // Dimensions batch_size, state_dim, @@ -412,306 +312,6 @@ impl CublasForward { Ok(()) } - // ══════════════════════════════════════════════════════════════════════════ - // Forward pass (online network — BF16 tensor core path) - // ══════════════════════════════════════════════════════════════════════════ - - /// Run the online network forward pass using cublasGemmEx BF16 tensor cores. - /// - /// Same layer sequence as `forward_online` but all GEMMs use BF16 weights, - /// BF16 activations, and F32 internal accumulation (H100: ~3x throughput). - /// - /// ## Data flow - /// - /// 1. **States F32→BF16**: `f32_to_bf16_kernel` converts the F32 experience - /// collector states into `self.states_bf16`. - /// 2. **Hidden layers**: `gemmex_bf16(W_bf16, input_bf16, output_bf16)` then - /// `add_bias_relu_bf16(output_bf16, bias_bf16)`. - /// 3. **Output layers**: `gemmex_bf16(W_bf16, input_bf16, output_bf16)` then - /// `add_bias_bf16(output_bf16, bias_bf16)`. - /// - /// All activations are written to `self.{h_s1_bf16, h_s2_bf16, ...}` internal - /// buffers. Callers access them via `h_s1_bf16_ptr()` etc. for the backward - /// pass (which will also operate in BF16). - /// - /// ## Weight pointer layout - /// - /// `bf16_w_ptrs` has the same 20-entry layout as the F32 `w_ptrs`: - /// `[w_s1, b_s1, w_s2, b_s2, w_v1, b_v1, w_v2, b_v2, - /// w_b0fc, b_b0fc, w_b0out, b_b0out, ...]` - /// — but all pointers target BF16 (`u16`) device memory in `bf16_params_buf`. - #[allow(dead_code, clippy::too_many_arguments)] - pub fn forward_online_bf16( - &self, - stream: &Arc, - // ── Inputs ────────────────────────────────────────────────────── - states_f32: &CudaSlice, // [B, SD] F32 from experience collector - // ── BF16 weight pointers (raw u64 into flat bf16_params_buf) ── - bf16_w_ptrs: &[u64; 20], // [w_s1, b_s1, w_s2, b_s2, ...] - ) -> Result<(), MLError> { - let b = self.batch_size; - let n_states = b * self.state_dim; - - // ── Step 1: Convert states F32 → BF16 ──────────────────────────── - { - let src_ptr = raw_bf16_ptr(states_f32, stream); - let dst_ptr = raw_u16_ptr(&self.states_bf16, stream); - let n_i32 = n_states as i32; - let blocks = ((n_states + 255) / 256) as u32; - - unsafe { - stream - .launch_builder(&self.f32_to_bf16_kernel) - .arg(&src_ptr) - .arg(&dst_ptr) - .arg(&n_i32) - .launch(LaunchConfig { - grid_dim: (blocks, 1, 1), - block_dim: (256, 1, 1), - shared_mem_bytes: 0, - }) - .map_err(|e| MLError::ModelError(format!("f32_to_bf16 states: {e}")))?; - } - } - - // ── Step 2: Shared trunk layer 1 ───────────────────────────────── - // h_s1[B, SH1] = ReLU(states_bf16 @ W_s1_bf16^T + b_s1_bf16) - let states_bf16_ptr = raw_u16_ptr(&self.states_bf16, stream); - let h_s1_bf16_ptr = raw_u16_ptr(&self.h_s1_bf16, stream); - - self.gemmex_bf16( - bf16_w_ptrs[0], // W_s1[SH1, SD] - states_bf16_ptr, - h_s1_bf16_ptr, - self.shared_h1, b, self.state_dim, - "bf16_h_s1", - )?; - self.launch_add_bias_relu_bf16_raw( - stream, h_s1_bf16_ptr, bf16_w_ptrs[1], self.shared_h1, b, - )?; - - // ── Step 3: Shared trunk layer 2 ───────────────────────────────── - // h_s2[B, SH2] = ReLU(h_s1_bf16 @ W_s2_bf16^T + b_s2_bf16) - let h_s2_bf16_ptr = raw_u16_ptr(&self.h_s2_bf16, stream); - - self.gemmex_bf16( - bf16_w_ptrs[2], // W_s2[SH2, SH1] - h_s1_bf16_ptr, - h_s2_bf16_ptr, - self.shared_h2, b, self.shared_h1, - "bf16_h_s2", - )?; - self.launch_add_bias_relu_bf16_raw( - stream, h_s2_bf16_ptr, bf16_w_ptrs[3], self.shared_h2, b, - )?; - - // ── Step 4: Value head layer 1 ─────────────────────────────────── - // h_v[B, VH] = ReLU(h_s2_bf16 @ W_v1_bf16^T + b_v1_bf16) - let h_v_bf16_ptr = raw_u16_ptr(&self.h_v_bf16, stream); - - self.gemmex_bf16( - bf16_w_ptrs[4], // W_v1[VH, SH2] - h_s2_bf16_ptr, - h_v_bf16_ptr, - self.value_h, b, self.shared_h2, - "bf16_h_v", - )?; - self.launch_add_bias_relu_bf16_raw( - stream, h_v_bf16_ptr, bf16_w_ptrs[5], self.value_h, b, - )?; - - // ── Step 5: Value head layer 2 (logits, no ReLU) ──────────────── - // v_logits[B, NA] = h_v_bf16 @ W_v2_bf16^T + b_v2_bf16 - let v_logits_bf16_ptr = raw_u16_ptr(&self.v_logits_bf16, stream); - - self.gemmex_bf16( - bf16_w_ptrs[6], // W_v2[NA, VH] - h_v_bf16_ptr, - v_logits_bf16_ptr, - self.num_atoms, b, self.value_h, - "bf16_v_logits", - )?; - self.launch_add_bias_bf16_raw( - stream, v_logits_bf16_ptr, bf16_w_ptrs[7], self.num_atoms, b, - )?; - - // ── Step 6: Branch heads ───────────────────────────────────────── - let branch_sizes = [self.branch_0_size, self.branch_1_size, self.branch_2_size]; - let branch_h_bf16_ptrs = [ - raw_u16_ptr(&self.h_b0_bf16, stream), - raw_u16_ptr(&self.h_b1_bf16, stream), - raw_u16_ptr(&self.h_b2_bf16, stream), - ]; - let branch_w_base = [8_usize, 12, 16]; - let b_logits_bf16_ptr = raw_u16_ptr(&self.b_logits_bf16, stream); - let na = self.num_atoms; - - let mut logit_byte_offset: u64 = 0; - for d in 0..3 { - let n_d = branch_sizes[d]; - let w_fc_idx = branch_w_base[d]; - let b_fc_idx = w_fc_idx + 1; - let w_out_idx = w_fc_idx + 2; - let b_out_idx = w_fc_idx + 3; - - // h_bd[B, AH] = ReLU(h_s2_bf16 @ W_bdk_fc_bf16^T + b_bdk_fc_bf16) - self.gemmex_bf16( - bf16_w_ptrs[w_fc_idx], // W_bdk_fc[AH, SH2] - h_s2_bf16_ptr, - branch_h_bf16_ptrs[d], - self.adv_h, b, self.shared_h2, - "bf16_h_bd", - )?; - self.launch_add_bias_relu_bf16_raw( - stream, branch_h_bf16_ptrs[d], bf16_w_ptrs[b_fc_idx], self.adv_h, b, - )?; - - // adv_logits_d[B, n_d*NA] = h_bd_bf16 @ W_bdk_out_bf16^T + b_bdk_out_bf16 - // Output pointer: offset into b_logits_bf16 by accumulated bytes - let adv_out_ptr = b_logits_bf16_ptr + logit_byte_offset; - self.gemmex_bf16( - bf16_w_ptrs[w_out_idx], // W_bdk_out[n_d*NA, AH] - branch_h_bf16_ptrs[d], - adv_out_ptr, - n_d * na, b, self.adv_h, - "bf16_adv_logits", - )?; - self.launch_add_bias_bf16_raw( - stream, adv_out_ptr, bf16_w_ptrs[b_out_idx], n_d * na, b, - )?; - - logit_byte_offset += (b * n_d * na * std::mem::size_of::()) as u64; - } - - Ok(()) - } - - /// BF16 target network forward: next_states(F32) → BF16 GemmEx → target logits(BF16). - /// - /// Uses target-specific internal BF16 scratch buffers (`tgt_h_s1_bf16`, etc.) - /// and writes logits to `tgt_v_logits_bf16` / `tgt_b_logits_bf16`. - /// Does NOT save activations for backward (inference only). - #[allow(dead_code, clippy::too_many_arguments)] - pub fn forward_target_bf16( - &self, - stream: &Arc, - next_states_f32: &CudaSlice, // [B, SD] F32 from experience collector - bf16_w_ptrs: &[u64; 20], // BF16 target weight pointers - ) -> Result<(), MLError> { - let b = self.batch_size; - let n_states = b * self.state_dim; - - // ── Step 1: Convert states F32 → BF16 (reuse online states_bf16 as scratch) ── - { - let src_ptr = raw_bf16_ptr(next_states_f32, stream); - let dst_ptr = raw_u16_ptr(&self.states_bf16, stream); - let n_i32 = n_states as i32; - let blocks = ((n_states + 255) / 256) as u32; - unsafe { - stream - .launch_builder(&self.f32_to_bf16_kernel) - .arg(&src_ptr) - .arg(&dst_ptr) - .arg(&n_i32) - .launch(LaunchConfig { - grid_dim: (blocks, 1, 1), - block_dim: (256, 1, 1), - shared_mem_bytes: 0, - }) - .map_err(|e| MLError::ModelError(format!("f32_to_bf16 tgt_states: {e}")))?; - } - } - - let states_bf16_ptr = raw_u16_ptr(&self.states_bf16, stream); - let tgt_h_s1_ptr = raw_u16_ptr(&self.tgt_h_s1_bf16, stream); - let tgt_h_s2_ptr = raw_u16_ptr(&self.tgt_h_s2_bf16, stream); - let tgt_h_v_ptr = raw_u16_ptr(&self.tgt_h_v_bf16, stream); - - // ── Step 2: Shared trunk layer 1 ──── - self.gemmex_bf16(bf16_w_ptrs[0], states_bf16_ptr, tgt_h_s1_ptr, - self.shared_h1, b, self.state_dim, "bf16_tgt_h_s1")?; - self.launch_add_bias_relu_bf16_raw(stream, tgt_h_s1_ptr, bf16_w_ptrs[1], self.shared_h1, b)?; - - // ── Step 3: Shared trunk layer 2 ──── - self.gemmex_bf16(bf16_w_ptrs[2], tgt_h_s1_ptr, tgt_h_s2_ptr, - self.shared_h2, b, self.shared_h1, "bf16_tgt_h_s2")?; - self.launch_add_bias_relu_bf16_raw(stream, tgt_h_s2_ptr, bf16_w_ptrs[3], self.shared_h2, b)?; - - // ── Step 4: Value head layer 1 ──── - self.gemmex_bf16(bf16_w_ptrs[4], tgt_h_s2_ptr, tgt_h_v_ptr, - self.value_h, b, self.shared_h2, "bf16_tgt_h_v")?; - self.launch_add_bias_relu_bf16_raw(stream, tgt_h_v_ptr, bf16_w_ptrs[5], self.value_h, b)?; - - // ── Step 5: Value head layer 2 (logits, no ReLU) ──── - let tgt_v_logits_ptr = raw_u16_ptr(&self.tgt_v_logits_bf16, stream); - self.gemmex_bf16(bf16_w_ptrs[6], tgt_h_v_ptr, tgt_v_logits_ptr, - self.num_atoms, b, self.value_h, "bf16_tgt_v_logits")?; - self.launch_add_bias_bf16_raw(stream, tgt_v_logits_ptr, bf16_w_ptrs[7], self.num_atoms, b)?; - - // ── Step 6: Branch heads ──── - let branch_sizes = [self.branch_0_size, self.branch_1_size, self.branch_2_size]; - let branch_w_base = [8_usize, 12, 16]; - let tgt_h_b_ptr = raw_u16_ptr(&self.tgt_h_b_bf16, stream); - let tgt_b_logits_ptr = raw_u16_ptr(&self.tgt_b_logits_bf16, stream); - let na = self.num_atoms; - - let mut logit_byte_offset: u64 = 0; - for d in 0..3 { - let n_d = branch_sizes[d]; - let w_fc_idx = branch_w_base[d]; - let b_fc_idx = w_fc_idx + 1; - let w_out_idx = w_fc_idx + 2; - let b_out_idx = w_fc_idx + 3; - - self.gemmex_bf16(bf16_w_ptrs[w_fc_idx], tgt_h_s2_ptr, tgt_h_b_ptr, - self.adv_h, b, self.shared_h2, "bf16_tgt_h_bd")?; - self.launch_add_bias_relu_bf16_raw(stream, tgt_h_b_ptr, bf16_w_ptrs[b_fc_idx], self.adv_h, b)?; - - let adv_out_ptr = tgt_b_logits_ptr + logit_byte_offset; - self.gemmex_bf16(bf16_w_ptrs[w_out_idx], tgt_h_b_ptr, adv_out_ptr, - n_d * na, b, self.adv_h, "bf16_tgt_adv_logits")?; - self.launch_add_bias_bf16_raw(stream, adv_out_ptr, bf16_w_ptrs[b_out_idx], n_d * na, b)?; - - logit_byte_offset += (b * n_d * na * std::mem::size_of::()) as u64; - } - - Ok(()) - } - - // ══════════════════════════════════════════════════════════════════════════ - // BF16 logit buffer accessors (for loss kernel wiring) - // ══════════════════════════════════════════════════════════════════════════ - - /// Raw pointer to online value logits BF16 buffer [B, NA]. - pub fn v_logits_bf16_ptr(&self, stream: &Arc) -> u64 { - raw_u16_ptr(&self.v_logits_bf16, stream) - } - - /// Raw pointer to online branch logits BF16 buffer [B, (B0+B1+B2)*NA]. - pub fn b_logits_bf16_ptr(&self, stream: &Arc) -> u64 { - raw_u16_ptr(&self.b_logits_bf16, stream) - } - - /// Raw pointer to target value logits BF16 buffer [B, NA]. - pub fn tgt_v_logits_bf16_ptr(&self, stream: &Arc) -> u64 { - raw_u16_ptr(&self.tgt_v_logits_bf16, stream) - } - - /// Raw pointer to target branch logits BF16 buffer [B, (B0+B1+B2)*NA]. - pub fn tgt_b_logits_bf16_ptr(&self, stream: &Arc) -> u64 { - raw_u16_ptr(&self.tgt_b_logits_bf16, stream) - } - - /// Reference to online value logits BF16 buffer. - pub fn v_logits_bf16_buf(&self) -> &CudaSlice { - &self.v_logits_bf16 - } - - /// Reference to online branch logits BF16 buffer. - pub fn b_logits_bf16_buf(&self) -> &CudaSlice { - &self.b_logits_bf16 - } - /// Run value head forward only: h_s2 → W_v1 → ReLU → W_v2 → v_logits. /// /// Used by ensemble heads (1..K-1) to compute per-head value logits @@ -835,85 +435,7 @@ impl CublasForward { } // ══════════════════════════════════════════════════════════════════════════ - // cuBLAS SGEMM helpers - // ══════════════════════════════════════════════════════════════════════════ - - /// Single cuBLAS SGEMM: C[B, N] = A[B, K] @ W[N, K]^T (row-major). - /// - /// Call convention (column-major cuBLAS, row-major data): - /// - /// sgemm(CUBLAS_OP_T, CUBLAS_OP_N, N, B, K, 1.0, W, K, A, K, 0.0, C, N) - /// - /// This outputs C in col-major [N, B] which equals row-major C[B, N]. ✓ - /// - /// Arguments: - /// - `w_ptr` : raw device pointer to W[N, K] (row-major, stride K) - /// - `a_ptr` : raw device pointer to A[B, K] (row-major, stride K) - /// - `c_buf` : output CudaSlice [B * N] (written in col-major [N, B]) - /// - `n` : output cols (out_dim) - /// - `b` : batch size (rows of A / cols of output in col-major) - /// - `k` : inner dim (in_dim = cols of W = cols of A) - #[allow(clippy::too_many_arguments)] - fn sgemm_layer( - &self, - _stream: &Arc, - w_ptr: u64, - a_ptr: u64, - c_ptr: u64, - n: usize, - b: usize, - k: usize, - _label: &str, - ) -> Result<(), MLError> { - let alpha = 1.0_f32; - let beta = 0.0_f32; - - // SAFETY: w_ptr, a_ptr, c_ptr are valid CUDA device pointers in the same - // context. Sizes are computed from config values that were validated at - // construction. The cuBLAS handle is bound to the correct stream. - unsafe { - cublas_result::sgemm( - self.handle.0, - cublas_sys::cublasOperation_t::CUBLAS_OP_T, // transa: transpose W - cublas_sys::cublasOperation_t::CUBLAS_OP_N, // transb: keep A as-is - n as i32, // m (rows of op(W) = N = out_dim) - b as i32, // n (cols of op(A) = B = batch) - k as i32, // k (inner dim) - &alpha, - w_ptr as *const f32, - k as i32, // lda: leading dim of W (before transpose) = K - a_ptr as *const f32, - k as i32, // ldb: leading dim of A = K - &beta, - c_ptr as *mut f32, - n as i32, // ldc: leading dim of C = N - ) - .map_err(|e| MLError::ModelError(format!("cublasSgemm {_label}: {e:?}")))?; - } - - Ok(()) - } - - /// Same as `sgemm_layer` but takes raw output pointer directly. - /// - /// Used when writing into a sub-region of a larger buffer (branch logits). - #[allow(clippy::too_many_arguments)] - fn sgemm_layer_raw( - &self, - _stream: &Arc, - w_ptr: u64, - a_ptr: u64, - c_ptr: u64, - n: usize, - b: usize, - k: usize, - _label: &str, - ) -> Result<(), MLError> { - self.sgemm_layer(_stream, w_ptr, a_ptr, c_ptr, n, b, k, _label) - } - - // ══════════════════════════════════════════════════════════════════════════ - // cublasGemmEx BF16 helpers (tensor core path) + // cublasGemmEx BF16 GEMM (tensor core path) // ══════════════════════════════════════════════════════════════════════════ /// BF16 x BF16 -> BF16 GEMM via `cublasGemmEx` with F32 internal accumulation. @@ -975,125 +497,7 @@ impl CublasForward { } // ══════════════════════════════════════════════════════════════════════════ - // Bias + activation kernel launchers (F32) - // ══════════════════════════════════════════════════════════════════════════ - - /// Launch `add_bias_relu` over an activation buffer. - /// - /// `output[i] = max(0, output[i] + bias[i % out_dim])` - /// - /// Grid: ceil(B * out_dim / 256), Block: 256. - fn launch_add_bias_relu( - &self, - stream: &Arc, - output: &CudaSlice, - bias_ptr: u64, - out_dim: usize, - batch: usize, - ) -> Result<(), MLError> { - let total = (batch * out_dim) as i32; - let out_dim_i32 = out_dim as i32; - let blocks = ((batch * out_dim + 255) / 256) as u32; - - let out_ptr = raw_bf16_ptr(output, stream); - - unsafe { - stream - .launch_builder(&self.add_bias_relu_kernel) - .arg(&out_ptr) - .arg(&bias_ptr) - .arg(&out_dim_i32) - .arg(&total) - .launch(LaunchConfig { - grid_dim: (blocks, 1, 1), - block_dim: (256, 1, 1), - shared_mem_bytes: 0, - }) - .map_err(|e| MLError::ModelError(format!("add_bias_relu kernel: {e}")))?; - } - - Ok(()) - } - - /// Launch `add_bias_relu` with a raw output pointer (sub-buffer support). - fn launch_add_bias_relu_raw( - &self, - stream: &Arc, - out_ptr: u64, - bias_ptr: u64, - out_dim: usize, - batch: usize, - ) -> Result<(), MLError> { - let total = (batch * out_dim) as i32; - let out_dim_i32 = out_dim as i32; - let blocks = ((batch * out_dim + 255) / 256) as u32; - - unsafe { - stream - .launch_builder(&self.add_bias_relu_kernel) - .arg(&out_ptr) - .arg(&bias_ptr) - .arg(&out_dim_i32) - .arg(&total) - .launch(LaunchConfig { - grid_dim: (blocks, 1, 1), - block_dim: (256, 1, 1), - shared_mem_bytes: 0, - }) - .map_err(|e| MLError::ModelError(format!("add_bias_relu_raw kernel: {e}")))?; - } - - Ok(()) - } - - /// Launch `add_bias` (no activation) over an output buffer. - /// - /// `output[i] = output[i] + bias[i % out_dim]` - fn launch_add_bias( - &self, - stream: &Arc, - output: &CudaSlice, - bias_ptr: u64, - out_dim: usize, - batch: usize, - ) -> Result<(), MLError> { - let out_ptr = raw_bf16_ptr(output, stream); - self.launch_add_bias_raw(stream, out_ptr, bias_ptr, out_dim, batch) - } - - /// Launch `add_bias` with a raw output pointer (sub-buffer support). - fn launch_add_bias_raw( - &self, - stream: &Arc, - out_ptr: u64, - bias_ptr: u64, - out_dim: usize, - batch: usize, - ) -> Result<(), MLError> { - let total = (batch * out_dim) as i32; - let out_dim_i32 = out_dim as i32; - let blocks = ((batch * out_dim + 255) / 256) as u32; - - unsafe { - stream - .launch_builder(&self.add_bias_kernel) - .arg(&out_ptr) - .arg(&bias_ptr) - .arg(&out_dim_i32) - .arg(&total) - .launch(LaunchConfig { - grid_dim: (blocks, 1, 1), - block_dim: (256, 1, 1), - shared_mem_bytes: 0, - }) - .map_err(|e| MLError::ModelError(format!("add_bias kernel: {e}")))?; - } - - Ok(()) - } - - // ══════════════════════════════════════════════════════════════════════════ - // Bias + activation kernel launchers (BF16 — tensor core path) + // Bias + activation kernel launchers (BF16) // ══════════════════════════════════════════════════════════════════════════ /// Launch `add_bias_relu_bf16` over a BF16 activation buffer. @@ -1167,27 +571,6 @@ impl CublasForward { Ok(()) } - // ══════════════════════════════════════════════════════════════════════════ - // BF16 activation buffer accessors (for callers wiring up the full BF16 path) - // ══════════════════════════════════════════════════════════════════════════ - - /// Raw device pointer to the online states BF16 buffer. - #[allow(dead_code)] - pub fn states_bf16_ptr(&self, stream: &Arc) -> u64 { - raw_u16_ptr(&self.states_bf16, stream) - } - - /// Raw device pointer to the online h_s1 BF16 buffer. - #[allow(dead_code)] - pub fn h_s1_bf16_ptr(&self, stream: &Arc) -> u64 { - raw_u16_ptr(&self.h_s1_bf16, stream) - } - - /// Raw device pointer to the online h_s2 BF16 buffer. - #[allow(dead_code)] - pub fn h_s2_bf16_ptr(&self, stream: &Arc) -> u64 { - raw_u16_ptr(&self.h_s2_bf16, stream) - } } // ── Raw device pointer helpers ──────────────────────────────────────────────── @@ -1199,54 +582,35 @@ fn raw_bf16_ptr(slice: &CudaSlice, stream: &Arc) -> u64 ptr } -/// Extract raw u16 device pointer (BF16 buffers stored as u16). -fn raw_u16_ptr(slice: &CudaSlice, stream: &Arc) -> u64 { - let (ptr, guard) = slice.device_ptr(stream); - let _no_drop = ManuallyDrop::new(guard); - ptr -} - // ── Kernel compilation ──────────────────────────────────────────────────────── /// Precompiled bias kernels cubin (build.rs — ZERO runtime nvcc). static BIAS_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/bias_kernels.cubin")); -/// Load F32 and BF16 bias+activation kernels from precompiled cubin. +/// Load BF16 bias+activation kernels from precompiled cubin. /// -/// Returns (add_bias_relu, add_bias, add_bias_relu_bf16, add_bias_bf16, f32_to_bf16). -/// -/// The `f32_to_bf16_kernel` is defined in `common_device_functions.cuh` which is -/// prepended to `bias_kernels.cu` at build time (see `build.rs`). +/// Returns (add_bias_relu_bf16, add_bias_bf16). fn compile_bias_kernels( stream: &Arc, -) -> Result<(CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction), MLError> { +) -> Result<(CudaFunction, CudaFunction), MLError> { let context = stream.context(); let module = context .load_cubin(BIAS_CUBIN.to_vec()) .map_err(|e| MLError::ModelError(format!("bias cubin load: {e}")))?; - let add_bias_relu = module - .load_function("add_bias_relu_kernel") - .map_err(|e| MLError::ModelError(format!("add_bias_relu_kernel load: {e}")))?; - let add_bias = module - .load_function("add_bias_kernel") - .map_err(|e| MLError::ModelError(format!("add_bias_kernel load: {e}")))?; let add_bias_relu_bf16 = module .load_function("add_bias_relu_bf16_kernel") .map_err(|e| MLError::ModelError(format!("add_bias_relu_bf16_kernel load: {e}")))?; let add_bias_bf16 = module .load_function("add_bias_bf16_kernel") .map_err(|e| MLError::ModelError(format!("add_bias_bf16_kernel load: {e}")))?; - let f32_to_bf16 = module - .load_function("f32_to_bf16_kernel") - .map_err(|e| MLError::ModelError(format!("f32_to_bf16_kernel load: {e}")))?; - Ok((add_bias_relu, add_bias, add_bias_relu_bf16, add_bias_bf16, f32_to_bf16)) + Ok((add_bias_relu_bf16, add_bias_bf16)) } -// ── Compute F32 weight pointers from flat params_buf ───────────────────────── +// ── Compute BF16 weight pointers from flat params_buf ─────────────────────── -/// Compute the 20 raw F32 device pointers into a flat params_buf at GOFF_* offsets. +/// Compute the 20 raw BF16 device pointers into a flat params_buf at GOFF_* offsets. /// /// The flat buffer layout matches `compute_param_sizes()`: /// [w_s1, b_s1, w_s2, b_s2, w_v1, b_v1, w_v2, b_v2, @@ -1274,29 +638,3 @@ pub fn f32_weight_ptrs( } ptrs } - -/// Compute 20 raw BF16 device pointers into a flat `CudaSlice` buffer. -/// -/// Uses precomputed byte offsets (padded for 4-byte alignment — see -/// `bf16_goff_byte_offsets` in `GpuDqnTrainer`). Each segment base is -/// `buf_base + byte_offset[i]`. -/// -/// Returns 20 raw u64 device pointers (BF16 stored as u16). -#[allow(dead_code)] -pub fn bf16_weight_ptrs( - bf16_buf: &CudaSlice, - bf16_byte_offsets: &[u64; 20], - stream: &Arc, -) -> [u64; 20] { - let base = { - let (ptr, guard) = bf16_buf.device_ptr(stream); - let _no_drop = ManuallyDrop::new(guard); - ptr - }; - - let mut ptrs = [0_u64; 20]; - for i in 0..20 { - ptrs[i] = base + bf16_byte_offsets[i]; - } - ptrs -} diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index 7108dba6e..44f46f062 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -312,8 +312,6 @@ pub(crate) fn compute_total_params(cfg: &GpuDqnTrainConfig) -> usize { struct CachedPtrs { params_buf: u64, target_params_buf: u64, - bf16_params_buf: u64, - bf16_target_params_buf: u64, grad_buf: u64, grad_norm_buf: u64, m_buf: u64, @@ -408,20 +406,6 @@ pub struct GpuDqnTrainer { /// Number of trunk parameters (w_s1 + b_s1 + w_s2 + b_s2). trunk_param_count: usize, - // ── Flat BF16 weight mirrors (forward kernel reads BF16 for tensor core throughput) ── - // Single contiguous BF16 buffer per network (online + target), same GOFF_* layout - // as the F32 flat buffers. One f32_to_bf16_kernel launch converts the entire buffer - // instead of 20 per-tensor launches. Forward kernels read at precomputed byte offsets. - bf16_params_buf: CudaSlice, // [TOTAL_PARAMS] flat online BF16 - bf16_target_params_buf: CudaSlice, // [TOTAL_PARAMS] flat target BF16 - bf16_mirrors_initialized: bool, - /// Precomputed byte offsets into flat BF16 buffers for each of the 20 weight tensors. - /// Layout matches GOFF_* order: w_s1, b_s1, w_s2, b_s2, ..., w_bu2, b_bu2. - /// Each offset is in bytes (element offset * sizeof(u16)). - /// Segments are padded to even element counts for 4-byte aligned short2 loads. - bf16_goff_byte_offsets: [u64; 20], - /// Total u16 elements in padded BF16 buffers (>= total_params due to alignment padding). - bf16_padded_total: usize, // ── Batch input buffers (uploaded per step) ───────────────────── states_buf: CudaSlice, // [B, STATE_DIM] @@ -571,9 +555,6 @@ pub struct GpuDqnTrainer { /// Gradient w.r.t. branch logits: [B, (B0+B1+B2)*NA] d_adv_logits_buf: CudaSlice, - /// Precomputed F32 weight byte offsets into params_buf. - /// Same GOFF_* layout as BF16, but for the F32 cuBLAS path. - f32_goff_byte_offsets: [u64; 20], // ── cuBLAS batched backward (Phase 2 Task 2) ────────────────────── /// cuBLAS backward context (handle + ReLU mask + bias grad kernels). @@ -1763,27 +1744,6 @@ impl GpuDqnTrainer { let cql_grad_scratch = alloc_f32(&stream, total_params, "cql_grad_scratch")?; let t_buf = alloc_i32(&stream, 1, "adam_t")?; - // ── Allocate flat BF16 weight mirror buffers ────────────────── - // One contiguous BF16 buffer per network. Forward kernels read at - // precomputed GOFF byte offsets. Each segment is padded to even element - // count so that every base pointer is 4-byte aligned — required by - // cooperative_load_tile_bf16 which vectorizes as short2 (4-byte loads). - let param_sizes = compute_param_sizes(&config); - let mut bf16_goff_byte_offsets = [0_u64; 20]; - let bf16_padded_total: usize; - { - let mut offset = 0_u64; - for i in 0..20 { - bf16_goff_byte_offsets[i] = offset; - // Pad to even element count → 4-byte aligned short2 loads - let padded = ((param_sizes[i] + 1) & !1) as u64; - offset += padded * (std::mem::size_of::() as u64); - } - bf16_padded_total = (offset / std::mem::size_of::() as u64) as usize; - } - let bf16_params_buf = alloc_u16(&stream, bf16_padded_total, "bf16_params")?; - let bf16_target_params_buf = alloc_u16(&stream, bf16_padded_total, "bf16_target_params")?; - // ── Allocate consolidated transfer buffers ───────────────── // Upload staging: states + next_states + actions(as f32) + rewards + dones + is_weights let upload_staging_len = b * config.state_dim * 2 + b * 4; // 2*B*SD + 4*B @@ -1952,15 +1912,6 @@ impl GpuDqnTrainer { // W_bu2 [b2*num_atoms, adv_h] alloc_spec_pair!(spec_u_bu2, spec_v_bu2, b2 * na, ah, "spec_u_bu2", "spec_v_bu2"); - // ── Precompute BF16 GOFF byte offsets (for cuBLAS weight pointers) ── - let mut f32_goff_byte_offsets = [0_u64; 20]; - { - let mut offset = 0_u64; - for i in 0..20 { - f32_goff_byte_offsets[i] = offset; - offset += (param_sizes[i] * std::mem::size_of::()) as u64; - } - } // ── Initialize cuBLAS forward context (required) ───────────── let cublas_forward = CublasForward::new( @@ -2017,8 +1968,6 @@ impl GpuDqnTrainer { CachedPtrs { params_buf: raw_device_ptr(¶ms_buf, &stream), target_params_buf: raw_device_ptr(&target_params_buf, &stream), - bf16_params_buf: raw_device_ptr_u16(&bf16_params_buf, &stream), - bf16_target_params_buf: raw_device_ptr_u16(&bf16_target_params_buf, &stream), grad_buf: raw_device_ptr(&grad_buf, &stream), grad_norm_buf: raw_device_ptr(&grad_norm_buf, &stream), m_buf: raw_device_ptr(&m_buf, &stream), @@ -2088,11 +2037,6 @@ impl GpuDqnTrainer { iqn_trunk_grad_norm, iqn_trunk_t_buf, trunk_param_count: trunk_params, - bf16_params_buf, - bf16_target_params_buf, - bf16_mirrors_initialized: false, - bf16_goff_byte_offsets, - bf16_padded_total, states_buf, next_states_buf, actions_buf, @@ -2161,7 +2105,6 @@ impl GpuDqnTrainer { d_adv_logits_buf, d_value_logits_mse, d_adv_logits_mse, - f32_goff_byte_offsets, cublas_backward, bw_d_h_s2, bw_d_h_s1, @@ -2178,138 +2121,11 @@ impl GpuDqnTrainer { }) } - /// Reference to the bf16→f32 conversion kernel. - /// - /// Exposed so that callers (e.g., `FusedTrainingCtx`) can convert BF16 - /// GpuBatch tensors to F32 CudaSlice buffers without creating Candle - /// temporary tensors. - pub fn bf16_to_f32_kernel(&self) -> &CudaFunction { - &self.bf16_to_f32_kernel - } - /// Reference to the trainer's forked CudaStream. pub fn stream(&self) -> &Arc { &self.stream } - // ═══════════════════════════════════════════════════════════════════ - // BF16 weight mirror management - // ═══════════════════════════════════════════════════════════════════ - - /// Ensure flat BF16 weight mirrors are synced from F32 originals. - /// - /// Called lazily on first `train_step()` or `forward_loss()`. Performs - /// initial F32 → BF16 conversion via 2 kernel launches (1 online + 1 target) - /// over the flat parameter buffers (same GOFF_* layout). - fn ensure_bf16_mirrors( - &mut self, - _online_d: &DuelingWeightSet, - _online_b: &BranchingWeightSet, - _target_d: &DuelingWeightSet, - _target_b: &BranchingWeightSet, - ) -> Result<(), MLError> { - if self.bf16_mirrors_initialized { - return Ok(()); - } - - // Per-segment F32 → BF16 at padded offsets (4-byte aligned short2 loads) - self.launch_bf16_convert_online()?; - self.launch_bf16_convert_target()?; - - self.bf16_mirrors_initialized = true; - - info!( - total_params = self.total_params, - bf16_padded_total = self.bf16_padded_total, - "GpuDqnTrainer: BF16 weight mirrors synced (20 seg launches × 2 networks)" - ); - Ok(()) - } - - /// Per-segment F32 → BF16 conversion: f32_buf → bf16_buf at padded offsets. - /// - /// 20 small kernel launches (one per weight tensor). Each segment is written - /// at its padded BF16 byte offset so all base pointers are 4-byte aligned - /// for short2 vectorized loads in the forward kernel. - fn launch_segmented_bf16_convert( - &self, - f32_buf: &CudaSlice, - bf16_buf: &CudaSlice, - ) -> Result<(), MLError> { - let param_sizes = compute_param_sizes(&self.config); - let f32_base = raw_device_ptr(f32_buf, &self.stream); - let bf16_base = raw_device_ptr_u16(bf16_buf, &self.stream); - - let mut f32_byte_offset = 0_u64; - for i in 0..20 { - let n = param_sizes[i]; - if n == 0 { continue; } - let f32_ptr = f32_base + f32_byte_offset; - let bf16_ptr = bf16_base + self.bf16_goff_byte_offsets[i]; - let n_i32 = n as i32; - let blocks = ((n + 255) / 256) as u32; - unsafe { - self.stream - .launch_builder(&self.f32_to_bf16_kernel) - .arg(&f32_ptr) - .arg(&bf16_ptr) - .arg(&n_i32) - .launch(LaunchConfig { - grid_dim: (blocks, 1, 1), - block_dim: (256, 1, 1), - shared_mem_bytes: 0, - }) - .map_err(|e| MLError::ModelError(format!("f32_to_bf16 seg[{i}]: {e}")))?; - } - f32_byte_offset += (n as u64) * (std::mem::size_of::() as u64); - } - Ok(()) - } - - /// Segmented F32 → BF16: params_buf → bf16_params_buf (online weights). - fn launch_bf16_convert_online(&self) -> Result<(), MLError> { - self.launch_segmented_bf16_convert(&self.params_buf, &self.bf16_params_buf) - } - - /// Segmented F32 → BF16: target_params_buf → bf16_target_params_buf (target weights). - fn launch_bf16_convert_target(&self) -> Result<(), MLError> { - self.launch_segmented_bf16_convert(&self.target_params_buf, &self.bf16_target_params_buf) - } - - /// Compute 20 raw device pointers into a flat BF16 buffer at GOFF_* offsets. - /// - /// Returns `[ptr_w_s1, ptr_b_s1, ptr_w_s2, ..., ptr_w_bu2, ptr_b_bu2]`. - /// Each pointer is the base of the flat buffer + the precomputed byte offset. - fn bf16_weight_ptrs(&self, bf16_buf: &CudaSlice) -> [u64; 20] { - let base = raw_device_ptr_u16(bf16_buf, &self.stream); - let mut ptrs = [0_u64; 20]; - for i in 0..20 { - ptrs[i] = base + self.bf16_goff_byte_offsets[i]; - } - ptrs - } - - /// Sync online BF16 mirrors from the flat F32 `params_buf`. - /// - /// Called inside the CUDA Graph after `unflatten_online_weights()` so that - /// the next graph replay's forward kernel reads updated BF16 weights. - /// 20 per-segment launches with padded offsets for 4-byte BF16 alignment. - fn sync_online_bf16( - &self, - _online_d: &DuelingWeightSet, - _online_b: &BranchingWeightSet, - ) -> Result<(), MLError> { - self.launch_bf16_convert_online() - } - - /// Sync target BF16 mirrors from the flat F32 `target_params_buf`. - /// - /// Called after `target_ema_update()` so that the next forward kernel - /// reads updated BF16 target weights. 20 per-segment launches with padded offsets. - fn sync_target_bf16(&self) -> Result<(), MLError> { - self.launch_bf16_convert_target() - } - // ═══════════════════════════════════════════════════════════════════ // Full training step (CUDA Graph — capture once, replay many) // ═══════════════════════════════════════════════════════════════════ @@ -3098,10 +2914,6 @@ impl GpuDqnTrainer { self.graph_adam = None; self.last_captured_loss_mode = None; self.params_initialized = false; - // Reset BF16 mirrors flag — they'll be re-synced on next train_step. - // Padded BF16 buffers are pre-allocated (no reallocation needed); only the - // content needs refreshing via 20 per-segment f32_to_bf16 launches each. - self.bf16_mirrors_initialized = false; } /// Switch the loss mode (MSE warmup vs C51 distributional). @@ -4078,10 +3890,6 @@ impl GpuDqnTrainer { // Scatter flat target buffer back to individual weight tensors self.unflatten_target_weights(target_d, target_b)?; - // Sync target BF16 mirrors from updated F32 target weights - // 20 per-segment launches: target_params_buf → bf16_target_params_buf - self.sync_target_bf16()?; - Ok(()) } } @@ -4322,13 +4130,6 @@ fn raw_device_ptr_u32(slice: &CudaSlice, stream: &CudaStream) -> u64 { ptr } -/// Extract raw CUDA device pointer from a `CudaSlice` (BF16 weight buffers). -fn raw_device_ptr_u16(slice: &CudaSlice, stream: &CudaStream) -> u64 { - let (ptr, guard) = slice.device_ptr(stream); - let _no_drop = std::mem::ManuallyDrop::new(guard); - ptr -} - /// Async device-to-device memcpy with error context. pub(crate) fn dtod_copy( dst: u64, diff --git a/crates/ml/src/trainers/dqn/fused_training.rs b/crates/ml/src/trainers/dqn/fused_training.rs index 61447e970..a5e316601 100644 --- a/crates/ml/src/trainers/dqn/fused_training.rs +++ b/crates/ml/src/trainers/dqn/fused_training.rs @@ -741,7 +741,7 @@ impl FusedTrainingCtx { let iqn_loss = iqn.per_sample_loss(); let td_errors = self.trainer.td_errors_buf(); let bs = self.trainer.batch_size(); - let n_bytes = bs * std::mem::size_of::(); + let n_bytes = bs * std::mem::size_of::(); let (src_ptr, _sg) = iqn_loss.device_ptr(&self.stream); let (dst_ptr, _dg) = td_errors.device_ptr(&self.stream); unsafe { @@ -859,7 +859,7 @@ impl FusedTrainingCtx { let k = self.ensemble_extra_heads.len() + 1; // total heads including head 0 let b = self.batch_size; let na = self.trainer.config().num_atoms; - let f32_size = std::mem::size_of::(); + let bf16_size = std::mem::size_of::(); // No cuStreamSynchronize needed — all ops are on the same stream. // CUDA guarantees in-order execution on a single stream. @@ -885,7 +885,7 @@ impl FusedTrainingCtx { { let head0_ptr = raw_device_ptr(self.trainer.on_v_logits_buf(), &self.stream); let dst_ptr = raw_device_ptr(logits_buf, &self.stream); - let n_bytes = b * na * f32_size; + let n_bytes = b * na * bf16_size; // Safety: both are valid CudaSlice on the same context. Byte sizes match. unsafe { cudarc::driver::result::memcpy_dtod_async( @@ -921,7 +921,7 @@ impl FusedTrainingCtx { let w_v2_ptr = raw_device_ptr(&head_dueling.w_v2, &self.stream); let b_v2_ptr = raw_device_ptr(&head_dueling.b_v2, &self.stream); let dst_k_ptr = raw_device_ptr(logits_buf, &self.stream) - + (k_idx * b * na * f32_size) as u64; + + (k_idx * b * na * bf16_size) as u64; self.trainer.forward_value_head_for_ensemble( w_v1_ptr, b_v1_ptr, w_v2_ptr, b_v2_ptr, @@ -934,7 +934,7 @@ impl FusedTrainingCtx { let div_loss_ptr = raw_device_ptr(div_loss_buf, &self.stream); unsafe { cudarc::driver::result::memset_d8_async( - div_loss_ptr, 0u8, f32_size, self.stream.cu_stream() + div_loss_ptr, 0u8, bf16_size, self.stream.cu_stream() ).map_err(|e| anyhow::anyhow!("Ensemble div_loss zero: {e}"))?; } @@ -972,7 +972,7 @@ impl FusedTrainingCtx { cudarc::driver::sys::cuMemcpyDtoH_v2( div_loss_host.as_mut_ptr().cast(), div_loss_ptr, - f32_size, + bf16_size, ); } // gpu-exit: 4-byte scalar readback let normalizer = if num_pairs > 0 { num_pairs as f32 * b as f32 } else { 1.0_f32 }; @@ -1443,15 +1443,15 @@ fn clone_dueling_weights( stream: &Arc, head_idx: usize, ) -> Result { - let clone_f32 = |slice: &cudarc::driver::CudaSlice| -> Result> { + let clone_bf16 = |slice: &cudarc::driver::CudaSlice| -> Result> { let n = slice.len(); let dst = stream.alloc_zeros::(n) - .map_err(|e| anyhow::anyhow!("Clone alloc {n}xf32: {e}"))?; + .map_err(|e| anyhow::anyhow!("Clone alloc {n}xbf16: {e}"))?; let (src_ptr, src_guard) = slice.device_ptr(stream); let (dst_ptr, dst_guard) = dst.device_ptr(stream); let _no_drop_s = std::mem::ManuallyDrop::new(src_guard); let _no_drop_d = std::mem::ManuallyDrop::new(dst_guard); - let n_bytes = n * std::mem::size_of::(); + let n_bytes = n * std::mem::size_of::(); unsafe { cudarc::driver::result::memcpy_dtod_async( dst_ptr, src_ptr, n_bytes, stream.cu_stream() @@ -1461,18 +1461,18 @@ fn clone_dueling_weights( }; let cloned = DuelingWeightSet { - w_s1: clone_f32(&src.w_s1)?, - b_s1: clone_f32(&src.b_s1)?, - w_s2: clone_f32(&src.w_s2)?, - b_s2: clone_f32(&src.b_s2)?, - w_v1: clone_f32(&src.w_v1)?, - b_v1: clone_f32(&src.b_v1)?, - w_v2: clone_f32(&src.w_v2)?, - b_v2: clone_f32(&src.b_v2)?, - w_a1: clone_f32(&src.w_a1)?, - b_a1: clone_f32(&src.b_a1)?, - w_a2: clone_f32(&src.w_a2)?, - b_a2: clone_f32(&src.b_a2)?, + w_s1: clone_bf16(&src.w_s1)?, + b_s1: clone_bf16(&src.b_s1)?, + w_s2: clone_bf16(&src.w_s2)?, + b_s2: clone_bf16(&src.b_s2)?, + w_v1: clone_bf16(&src.w_v1)?, + b_v1: clone_bf16(&src.b_v1)?, + w_v2: clone_bf16(&src.w_v2)?, + b_v2: clone_bf16(&src.b_v2)?, + w_a1: clone_bf16(&src.w_a1)?, + b_a1: clone_bf16(&src.b_a1)?, + w_a2: clone_bf16(&src.w_a2)?, + b_a2: clone_bf16(&src.b_a2)?, }; // Sync so all DtoD copies complete before we return (caller uses weights immediately) @@ -1488,15 +1488,15 @@ fn clone_branching_weights( stream: &Arc, head_idx: usize, ) -> Result { - let clone_f32 = |slice: &cudarc::driver::CudaSlice| -> Result> { + let clone_bf16 = |slice: &cudarc::driver::CudaSlice| -> Result> { let n = slice.len(); let dst = stream.alloc_zeros::(n) - .map_err(|e| anyhow::anyhow!("Clone alloc {n}xf32: {e}"))?; + .map_err(|e| anyhow::anyhow!("Clone alloc {n}xbf16: {e}"))?; let (src_ptr, src_guard) = slice.device_ptr(stream); let (dst_ptr, dst_guard) = dst.device_ptr(stream); let _no_drop_s = std::mem::ManuallyDrop::new(src_guard); let _no_drop_d = std::mem::ManuallyDrop::new(dst_guard); - let n_bytes = n * std::mem::size_of::(); + let n_bytes = n * std::mem::size_of::(); unsafe { cudarc::driver::result::memcpy_dtod_async( dst_ptr, src_ptr, n_bytes, stream.cu_stream() @@ -1506,14 +1506,14 @@ fn clone_branching_weights( }; let cloned = BranchingWeightSet { - w_bo1: clone_f32(&src.w_bo1)?, - b_bo1: clone_f32(&src.b_bo1)?, - w_bo2: clone_f32(&src.w_bo2)?, - b_bo2: clone_f32(&src.b_bo2)?, - w_bu1: clone_f32(&src.w_bu1)?, - b_bu1: clone_f32(&src.b_bu1)?, - w_bu2: clone_f32(&src.w_bu2)?, - b_bu2: clone_f32(&src.b_bu2)?, + w_bo1: clone_bf16(&src.w_bo1)?, + b_bo1: clone_bf16(&src.b_bo1)?, + w_bo2: clone_bf16(&src.w_bo2)?, + b_bo2: clone_bf16(&src.b_bo2)?, + w_bu1: clone_bf16(&src.w_bu1)?, + b_bu1: clone_bf16(&src.b_bu1)?, + w_bu2: clone_bf16(&src.w_bu2)?, + b_bu2: clone_bf16(&src.b_bu2)?, }; let _ = head_idx;