refactor(bf16): Spec C dead code cleanup — delete 911 LOC of F32 paths
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<u16> 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::<f32> → size_of::<half::bf16>): - 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) <noreply@anthropic.com>
This commit is contained in:
@@ -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<u8>,
|
||||
|
||||
// ── 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<half::bf16>
|
||||
// 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<u16>, // [B, SD]
|
||||
h_s1_bf16: CudaSlice<u16>, // [B, SH1]
|
||||
h_s2_bf16: CudaSlice<u16>, // [B, SH2]
|
||||
h_v_bf16: CudaSlice<u16>, // [B, VH]
|
||||
h_b0_bf16: CudaSlice<u16>, // [B, AH]
|
||||
h_b1_bf16: CudaSlice<u16>, // [B, AH]
|
||||
h_b2_bf16: CudaSlice<u16>, // [B, AH]
|
||||
/// Online logit BF16 buffers
|
||||
v_logits_bf16: CudaSlice<u16>, // [B, NA]
|
||||
b_logits_bf16: CudaSlice<u16>, // [B, (B0+B1+B2)*NA]
|
||||
|
||||
/// Target network BF16 activation buffers (inference scratch — reused)
|
||||
tgt_h_s1_bf16: CudaSlice<u16>, // [B, SH1]
|
||||
tgt_h_s2_bf16: CudaSlice<u16>, // [B, SH2]
|
||||
tgt_h_v_bf16: CudaSlice<u16>, // [B, VH]
|
||||
tgt_h_b_bf16: CudaSlice<u16>, // [B, AH] (shared across branches)
|
||||
/// Target logit BF16 buffers
|
||||
tgt_v_logits_bf16: CudaSlice<u16>, // [B, NA]
|
||||
tgt_b_logits_bf16: CudaSlice<u16>, // [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::<u16>(batch_size * state_dim)
|
||||
.map_err(|e| MLError::ModelError(format!("BF16 states_bf16 alloc: {e}")))?;
|
||||
let h_s1_bf16 = stream.alloc_zeros::<u16>(batch_size * shared_h1)
|
||||
.map_err(|e| MLError::ModelError(format!("BF16 h_s1_bf16 alloc: {e}")))?;
|
||||
let h_s2_bf16 = stream.alloc_zeros::<u16>(batch_size * shared_h2)
|
||||
.map_err(|e| MLError::ModelError(format!("BF16 h_s2_bf16 alloc: {e}")))?;
|
||||
let h_v_bf16 = stream.alloc_zeros::<u16>(batch_size * value_h)
|
||||
.map_err(|e| MLError::ModelError(format!("BF16 h_v_bf16 alloc: {e}")))?;
|
||||
let h_b0_bf16 = stream.alloc_zeros::<u16>(batch_size * adv_h)
|
||||
.map_err(|e| MLError::ModelError(format!("BF16 h_b0_bf16 alloc: {e}")))?;
|
||||
let h_b1_bf16 = stream.alloc_zeros::<u16>(batch_size * adv_h)
|
||||
.map_err(|e| MLError::ModelError(format!("BF16 h_b1_bf16 alloc: {e}")))?;
|
||||
let h_b2_bf16 = stream.alloc_zeros::<u16>(batch_size * adv_h)
|
||||
.map_err(|e| MLError::ModelError(format!("BF16 h_b2_bf16 alloc: {e}")))?;
|
||||
let v_logits_bf16 = stream.alloc_zeros::<u16>(batch_size * num_atoms)
|
||||
.map_err(|e| MLError::ModelError(format!("BF16 v_logits_bf16 alloc: {e}")))?;
|
||||
let b_logits_bf16 = stream.alloc_zeros::<u16>(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::<u16>(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::<u16>(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::<u16>(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::<u16>(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::<u16>(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::<u16>(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<CudaStream>,
|
||||
// ── Inputs ──────────────────────────────────────────────────────
|
||||
states_f32: &CudaSlice<half::bf16>, // [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::<u16>()) 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<CudaStream>,
|
||||
next_states_f32: &CudaSlice<half::bf16>, // [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::<u16>()) 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<CudaStream>) -> 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<CudaStream>) -> 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<CudaStream>) -> 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<CudaStream>) -> 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<u16> {
|
||||
&self.v_logits_bf16
|
||||
}
|
||||
|
||||
/// Reference to online branch logits BF16 buffer.
|
||||
pub fn b_logits_bf16_buf(&self) -> &CudaSlice<u16> {
|
||||
&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<half::bf16> [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<CudaStream>,
|
||||
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<CudaStream>,
|
||||
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<CudaStream>,
|
||||
output: &CudaSlice<half::bf16>,
|
||||
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<CudaStream>,
|
||||
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<CudaStream>,
|
||||
output: &CudaSlice<half::bf16>,
|
||||
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<CudaStream>,
|
||||
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<CudaStream>) -> 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<CudaStream>) -> 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<CudaStream>) -> u64 {
|
||||
raw_u16_ptr(&self.h_s2_bf16, stream)
|
||||
}
|
||||
}
|
||||
|
||||
// ── Raw device pointer helpers ────────────────────────────────────────────────
|
||||
@@ -1199,54 +582,35 @@ fn raw_bf16_ptr(slice: &CudaSlice<half::bf16>, stream: &Arc<CudaStream>) -> u64
|
||||
ptr
|
||||
}
|
||||
|
||||
/// Extract raw u16 device pointer (BF16 buffers stored as u16).
|
||||
fn raw_u16_ptr(slice: &CudaSlice<u16>, stream: &Arc<CudaStream>) -> 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<CudaStream>,
|
||||
) -> 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<u16>` 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<u16>,
|
||||
bf16_byte_offsets: &[u64; 20],
|
||||
stream: &Arc<CudaStream>,
|
||||
) -> [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
|
||||
}
|
||||
|
||||
@@ -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<u16>, // [TOTAL_PARAMS] flat online BF16
|
||||
bf16_target_params_buf: CudaSlice<u16>, // [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<half::bf16>, // [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<half::bf16>,
|
||||
|
||||
/// 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::<u16>() as u64);
|
||||
}
|
||||
bf16_padded_total = (offset / std::mem::size_of::<u16>() 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::<half::bf16>()) 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<CudaStream> {
|
||||
&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<half::bf16>,
|
||||
bf16_buf: &CudaSlice<u16>,
|
||||
) -> 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::<half::bf16>() 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<u16>) -> [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<u32>, stream: &CudaStream) -> u64 {
|
||||
ptr
|
||||
}
|
||||
|
||||
/// Extract raw CUDA device pointer from a `CudaSlice<u16>` (BF16 weight buffers).
|
||||
fn raw_device_ptr_u16(slice: &CudaSlice<u16>, 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,
|
||||
|
||||
@@ -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::<f32>();
|
||||
let n_bytes = bs * std::mem::size_of::<half::bf16>();
|
||||
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::<f32>();
|
||||
let bf16_size = std::mem::size_of::<half::bf16>();
|
||||
|
||||
// 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<half::bf16> 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<cudarc::driver::CudaStream>,
|
||||
head_idx: usize,
|
||||
) -> Result<DuelingWeightSet> {
|
||||
let clone_f32 = |slice: &cudarc::driver::CudaSlice<half::bf16>| -> Result<cudarc::driver::CudaSlice<half::bf16>> {
|
||||
let clone_bf16 = |slice: &cudarc::driver::CudaSlice<half::bf16>| -> Result<cudarc::driver::CudaSlice<half::bf16>> {
|
||||
let n = slice.len();
|
||||
let dst = stream.alloc_zeros::<half::bf16>(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::<f32>();
|
||||
let n_bytes = n * std::mem::size_of::<half::bf16>();
|
||||
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<cudarc::driver::CudaStream>,
|
||||
head_idx: usize,
|
||||
) -> Result<BranchingWeightSet> {
|
||||
let clone_f32 = |slice: &cudarc::driver::CudaSlice<half::bf16>| -> Result<cudarc::driver::CudaSlice<half::bf16>> {
|
||||
let clone_bf16 = |slice: &cudarc::driver::CudaSlice<half::bf16>| -> Result<cudarc::driver::CudaSlice<half::bf16>> {
|
||||
let n = slice.len();
|
||||
let dst = stream.alloc_zeros::<half::bf16>(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::<f32>();
|
||||
let n_bytes = n * std::mem::size_of::<half::bf16>();
|
||||
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;
|
||||
|
||||
Reference in New Issue
Block a user