From 875f263ec995997e82cc607e084176b4e28fa7c9 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sun, 29 Mar 2026 12:17:04 +0200 Subject: [PATCH] =?UTF-8?q?feat(bf16):=20f32=20backward=20dX=20scratch=20+?= =?UTF-8?q?=20bf16=20staging=20=E2=80=94=20eliminates=20dX=20truncation?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Backward pass dX computation now uses f32 scratch buffers with bf16 staging for the backward chain. Pattern: f32 GemmEx → relu_mask (f32) → f32→bf16 cast → next layer reads bf16 dY. Changes: - 6 bw_d_h_* scratch buffers: CudaSlice → CudaSlice - New bw_dy_bf16_staging: shared bf16 buffer for layer transitions - backward_fc_layer dX: gemmex_bf16 → gemmex_bf16_acc_f32 (f32 output) - launch_dx_only: gemmex_bf16 → gemmex_bf16_acc_f32 (f32 with beta) - relu_mask_kernel: reads/writes f32 (no bf16 clamp needed) - f32_to_bf16_cast_kernel: ±500 clamp at type boundary (in backward_kernels.cu) - cast_dx_to_staging: f32 scratch → bf16 staging per layer - IQN/ensemble backward: bf16→f32 cast for dX input, f32→bf16 for dY output - bw_d_h_s2_as_bf16(): attention backward receives bf16 via staging Hyperparameters updated for f32 Adam: - learning_rate: 1e-5 → 1e-4 (updates must exceed bf16 shadow step ~1e-3) - adam_epsilon: 1e-3 → 1e-8 (standard Adam, bf16 workaround no longer needed) - grad_norm NaN skip kept as defense-in-depth (source still under investigation) 895/895 unit + 359/359 ml-dqn tests pass. 7-11/11 smoke tests (intermittent NaN from unknown source — NOT backward dX). Co-Authored-By: Claude Opus 4.6 (1M context) --- .../ml/src/cuda_pipeline/backward_kernels.cu | 49 +++-- .../ml/src/cuda_pipeline/batched_backward.rs | 175 ++++++++++++----- .../ml/src/cuda_pipeline/gpu_dqn_trainer.rs | 178 +++++++++++------- .../ml/src/cuda_pipeline/relu_mask_kernel.cu | 12 +- crates/ml/src/trainers/dqn/fused_training.rs | 6 +- 5 files changed, 290 insertions(+), 130 deletions(-) diff --git a/crates/ml/src/cuda_pipeline/backward_kernels.cu b/crates/ml/src/cuda_pipeline/backward_kernels.cu index 99157dd2e..9681a91f0 100644 --- a/crates/ml/src/cuda_pipeline/backward_kernels.cu +++ b/crates/ml/src/cuda_pipeline/backward_kernels.cu @@ -9,19 +9,14 @@ * block=(256, 1, 1). */ -/* ReLU mask + bf16 overflow clamp for backward dX. +/* ReLU mask for f32 backward dX. * - * The dX GemmEx (bf16 A × bf16 B → bf16 C) writes bf16 output. When the - * f32 accumulated sum exceeds bf16 max (~65504), the bf16 write produces - * Inf. This clamp prevents Inf from cascading through the backward chain. - * Same pattern as the forward bias kernel ±500 clamp. - * - * This is NOT a NaN guard — it's bf16 overflow prevention at the type - * boundary, identical to the forward-pass bias kernel clamping. */ -#define BW_SAFE_MAX 500.0f + * The dX GemmEx now writes f32 output (bf16 A × bf16 B → f32 C). + * No clamp needed — f32 has sufficient dynamic range. + * The kernel just zeros elements where the saved activation <= 0. */ extern "C" __global__ void relu_mask_kernel( - __nv_bfloat16* __restrict__ dy, + float* __restrict__ dy, const __nv_bfloat16* __restrict__ activation, int n) { @@ -29,14 +24,38 @@ extern "C" __global__ void relu_mask_kernel( if (i >= n) return; float act_f = (float)activation[i]; if (act_f <= 0.0f) { - dy[i] = bf16_zero(); - } else { - float dy_f = (float)dy[i]; - dy_f = fminf(fmaxf(dy_f, -BW_SAFE_MAX), BW_SAFE_MAX); - dy[i] = bf16(dy_f); + dy[i] = 0.0f; } } +/* Cast f32 dX scratch → bf16 staging buffer with bf16-safe clamping. + * + * Called once per layer transition: the f32 dX from the current layer + * is cast to bf16 so it can be passed as dY to the next layer's + * backward_fc_layer (which reads bf16 inputs for the dW GemmEx). + * + * The clamp prevents bf16 overflow (max ~65504) when f32 accumulated + * sums happen to be very large. Without this, __float2bfloat16 produces + * Inf for values exceeding bf16 max, which poisons downstream weight + * gradients. This is the ONLY place where bf16 clamping is needed in + * the backward chain — the entire dX computation stays in f32. + * + * The limit matches the old per-layer BW_SAFE_MAX but is now applied + * once at the type boundary instead of after every ReLU mask. */ +#define STAGING_SAFE_MAX 500.0f + +extern "C" __global__ void f32_to_bf16_cast_kernel( + __nv_bfloat16* __restrict__ dst, + const float* __restrict__ src, + int n) +{ + int i = blockIdx.x * blockDim.x + threadIdx.x; + if (i >= n) return; + float v = src[i]; + v = fminf(fmaxf(v, -STAGING_SAFE_MAX), STAGING_SAFE_MAX); + dst[i] = bf16(v); +} + extern "C" __global__ void bias_grad_reduce_kernel( const __nv_bfloat16* __restrict__ dy, float* __restrict__ db, /* f32 gradient accumulator (grad_buf) */ diff --git a/crates/ml/src/cuda_pipeline/batched_backward.rs b/crates/ml/src/cuda_pipeline/batched_backward.rs index cf5424211..1836494f0 100644 --- a/crates/ml/src/cuda_pipeline/batched_backward.rs +++ b/crates/ml/src/cuda_pipeline/batched_backward.rs @@ -120,11 +120,15 @@ pub struct CublasBackward { handle: SendSyncCublasHandle, /// `relu_mask_kernel(dx, activation, n)` — element-wise ReLU derivative gate. + /// dx is f32 (backward scratch), activation is bf16 (saved forward output). relu_mask_kernel: CudaFunction, /// `bias_grad_reduce_kernel(dy, db, out_dim, batch_size)` — reduce dY over batch. bias_grad_kernel: CudaFunction, + /// `f32_to_bf16_cast_kernel(dst, src, n)` — cast f32 dX scratch to bf16 staging. + f32_to_bf16_cast_kernel: CudaFunction, + // ── Network dimensions (baked at construction) ── batch_size: usize, state_dim: usize, @@ -162,12 +166,13 @@ impl CublasBackward { } // ── Compile helper kernels ────────────────────────────────── - let (relu_mask_kernel, bias_grad_kernel) = compile_backward_kernels(stream)?; + let (relu_mask_kernel, bias_grad_kernel, f32_to_bf16_cast_kernel) = compile_backward_kernels(stream)?; Ok(Self { handle: SendSyncCublasHandle(raw_handle), relu_mask_kernel, bias_grad_kernel, + f32_to_bf16_cast_kernel, batch_size: config.batch_size, state_dim: config.state_dim, shared_h1: config.shared_h1, @@ -341,9 +346,9 @@ impl CublasBackward { self.launch_bias_grad(stream, dy, db, out_dim, batch)?; // ── Upstream gradient: dX[B, in] = dY[B, out] @ W[out, in] ── + // dX is f32 scratch — use gemmex_bf16_acc_f32 (bf16 A,B → f32 C). if dx != 0 { - // GemmEx BF16: N, N, m=in_dim, n=B, k=out_dim, beta=0.0 (overwrite) - self.gemmex_bf16( + self.gemmex_bf16_acc_f32( cublas_sys::cublasOperation_t::CUBLAS_OP_N, cublas_sys::cublasOperation_t::CUBLAS_OP_N, in_dim as i32, batch as i32, out_dim as i32, @@ -447,11 +452,13 @@ impl CublasBackward { w_ptrs: &[u64; 20], // Flat gradient accumulator (must be zeroed by caller) grad_buf_base: u64, - // Scratch buffers for inter-layer gradients - scratch_d_h_s2: u64, // [B, SH2] — accumulated from value + all 3 branches - scratch_d_h_s1: u64, // [B, SH1] - scratch_d_h_v: u64, // [B, VH] - scratch_d_h_b: &[u64; 3], // [B, AH] each + // Scratch buffers for inter-layer gradients (f32) + scratch_d_h_s2: u64, // [B, SH2] — f32, accumulated from value + all 3 branches + scratch_d_h_s1: u64, // [B, SH1] — f32 + scratch_d_h_v: u64, // [B, VH] — f32 + scratch_d_h_b: &[u64; 3], // [B, AH] each — f32 + // Shared bf16 staging buffer (overwritten per layer) + staging_bf16: u64, // max(B*SH2, B*SH1, B*VH, B*AH) — bf16 ) -> Result<(), MLError> { let b = self.batch_size; let na = self.num_atoms; @@ -587,19 +594,22 @@ impl CublasBackward { // setting dx=0 for the fc layer and doing the upstream gradient // via a separate GEMM call with beta=1.0 for branches d>0. - // Apply ReLU mask to branch FC upstream: d_h_b[d] *= (save_h_b[d] > 0) + // Apply ReLU mask to branch output dX (f32→f32): d_h_b[d] *= (save_h_b[d] > 0) self.relu_mask( stream, - scratch_d_h_b[d], // dx to gate - save_h_b[d], // saved post-ReLU activation + scratch_d_h_b[d], // f32 dx to gate + save_h_b[d], // bf16 saved post-ReLU activation b * self.adv_h, )?; - // dW for branch FC: dW[AH, SH2] += d_h_b[d]^T @ h_s2 + // Cast f32 d_h_b[d] → bf16 staging for use as dY in branch FC backward. + self.cast_dx_to_staging(stream, scratch_d_h_b[d], staging_bf16, b * self.adv_h)?; + + // dW for branch FC: dW[AH, SH2] += staging_bf16^T @ h_s2 // (dX is computed separately below to allow accumulation) self.launch_dw_only( stream, - scratch_d_h_b[d], // dY [B, AH] + staging_bf16, // dY [B, AH] — bf16 staging save_h_s2, // X [B, SH2] grad_buf_base + goff_w_bfc[d], // dW [AH, SH2] grad_buf_base + goff_b_bfc[d], // db [AH] @@ -608,14 +618,15 @@ impl CublasBackward { b, )?; - // dX for branch FC: d_h_s2 += d_h_b[d] @ W_bdk_fc + // dX for branch FC: d_h_s2 += staging_bf16 @ W_bdk_fc // Use beta = if d==0 { 0.0 } else { 1.0 } to accumulate branches. + // Output dX is f32 (gemmex_bf16_acc_f32). let beta_s2 = if d == 0 { 0.0_f32 } else { 1.0_f32 }; self.launch_dx_only( stream, - scratch_d_h_b[d], // dY [B, AH] + staging_bf16, // dY [B, AH] — bf16 staging w_fc, // W [AH, SH2] - scratch_d_h_s2, // dX [B, SH2] + scratch_d_h_s2, // dX [B, SH2] — f32 self.adv_h, // out_dim self.shared_h2, // in_dim b, @@ -643,12 +654,16 @@ impl CublasBackward { )?; // ── Value FC layer (ReLU) ───────────────────────────────────── + // ReLU mask on f32 dX_v self.relu_mask(stream, scratch_d_h_v, save_h_v, b * self.value_h)?; - // dW for value FC: dW[VH, SH2] += d_h_v^T @ h_s2 + // Cast f32 d_h_v → bf16 staging for value FC backward + self.cast_dx_to_staging(stream, scratch_d_h_v, staging_bf16, b * self.value_h)?; + + // dW for value FC: dW[VH, SH2] += staging_bf16^T @ h_s2 self.launch_dw_only( stream, - scratch_d_h_v, + staging_bf16, // dY [B, VH] — bf16 staging save_h_s2, grad_buf_base + goff_w_v1, grad_buf_base + goff_b_v1, @@ -657,10 +672,10 @@ impl CublasBackward { b, )?; - // dX for value FC: accumulated into scratch_d_h_s2 (beta=1.0 — branches already wrote) + // dX for value FC: accumulated into f32 scratch_d_h_s2 (beta=1.0 — branches already wrote) self.launch_dx_only( stream, - scratch_d_h_v, + staging_bf16, // dY [B, VH] — bf16 staging w_ptrs[4], // W_v1 [VH, SH2] scratch_d_h_s2, self.value_h, @@ -674,28 +689,36 @@ impl CublasBackward { // ══════════════════════════════════════════════════════════════════ // ── Shared layer 2 (ReLU) ───────────────────────────────────── + // ReLU mask on f32 accumulated d_h_s2 self.relu_mask(stream, scratch_d_h_s2, save_h_s2, b * self.shared_h2)?; + // Cast f32 d_h_s2 → bf16 staging for shared layer 2 backward + self.cast_dx_to_staging(stream, scratch_d_h_s2, staging_bf16, b * self.shared_h2)?; + self.backward_fc_layer( stream, - scratch_d_h_s2, + staging_bf16, // dY [B, SH2] — bf16 staging save_h_s1, w_ptrs[2], // W_s2 [SH2, SH1] grad_buf_base + goff_w_s2, grad_buf_base + goff_b_s2, - scratch_d_h_s1, // dX [B, SH1] + scratch_d_h_s1, // dX [B, SH1] — f32 self.shared_h2, self.shared_h1, b, )?; // ── Shared layer 1 (ReLU) — input layer, no dX needed ──────── + // ReLU mask on f32 d_h_s1 self.relu_mask(stream, scratch_d_h_s1, save_h_s1, b * self.shared_h1)?; + // Cast f32 d_h_s1 → bf16 staging for shared layer 1 backward + self.cast_dx_to_staging(stream, scratch_d_h_s1, staging_bf16, b * self.shared_h1)?; + // dW for shared layer 1: states is the input self.backward_fc_layer( stream, - scratch_d_h_s1, + staging_bf16, // dY [B, SH1] — bf16 staging states, w_ptrs[0], // W_s1 [SH1, SD] grad_buf_base + goff_w_s1, @@ -751,6 +774,9 @@ impl CublasBackward { /// `beta=0.0` overwrites dX; `beta=1.0` accumulates into dX. Used to merge /// contributions from multiple branches into the shared `d_h_s2` buffer. /// + /// dX is f32 scratch — uses gemmex_bf16_acc_f32 (bf16 A,B → f32 C). + /// Beta works the same for f32 C as it did for bf16 C. + /// /// Also used by ensemble diversity backward to skip dW/db for value head layers /// (only the upstream gradient d_h_s2 is needed, not the value head weight grads). #[allow(clippy::too_many_arguments)] @@ -765,8 +791,8 @@ impl CublasBackward { batch: usize, beta: f32, ) -> Result<(), MLError> { - // dX[B, in] = dY[B, out] @ W[out, in] — GemmEx BF16 - self.gemmex_bf16( + // dX[B, in] = dY[B, out] @ W[out, in] — GemmEx BF16→F32 + self.gemmex_bf16_acc_f32( cublas_sys::cublasOperation_t::CUBLAS_OP_N, cublas_sys::cublasOperation_t::CUBLAS_OP_N, in_dim as i32, batch as i32, out_dim as i32, @@ -807,6 +833,40 @@ impl CublasBackward { Ok(()) } + + /// Cast f32 dX scratch to bf16 staging buffer. + /// + /// Called once per layer transition in the backward chain. The f32 dX + /// from the current layer is cast to bf16 into `staging_ptr` so it can + /// be used as the bf16 `dY` input for the next layer's dW GemmEx. + /// + /// Grid: `ceil(n / 256)`, Block: 256. + pub fn cast_dx_to_staging( + &self, + stream: &Arc, + dx_f32: u64, + staging_bf16: u64, + n: usize, + ) -> Result<(), MLError> { + let n_i32 = n as i32; + let blocks = ((n + 255) / 256) as u32; + + unsafe { + stream + .launch_builder(&self.f32_to_bf16_cast_kernel) + .arg(&staging_bf16) + .arg(&dx_f32) + .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_cast_kernel: {e}")))?; + } + + Ok(()) + } } // ── Kernel compilation ──────────────────────────────────────────────────────── @@ -826,7 +886,7 @@ static BACKWARD_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/backwar fn compile_backward_kernels( stream: &Arc, -) -> Result<(CudaFunction, CudaFunction), MLError> { +) -> Result<(CudaFunction, CudaFunction, CudaFunction), MLError> { let context = stream.context(); let module = context .load_cubin(BACKWARD_CUBIN.to_vec()) @@ -838,8 +898,11 @@ fn compile_backward_kernels( let bias_grad = module .load_function("bias_grad_reduce_kernel") .map_err(|e| MLError::ModelError(format!("bias_grad_reduce_kernel load: {e}")))?; + let f32_to_bf16_cast = module + .load_function("f32_to_bf16_cast_kernel") + .map_err(|e| MLError::ModelError(format!("f32_to_bf16_cast_kernel load: {e}")))?; - Ok((relu_mask, bias_grad)) + Ok((relu_mask, bias_grad, f32_to_bf16_cast)) } // ── Raw device pointer helpers ──────────────────────────────────────────────── @@ -848,8 +911,15 @@ fn compile_backward_kernels( // can be compiled independently without `pub use`-ing the forward module's // private functions. -/// Extract raw F32 device pointer from a CudaSlice (read-only). -pub(crate) fn raw_f32_ptr(slice: &CudaSlice, stream: &Arc) -> u64 { +/// Extract raw device pointer from a bf16 CudaSlice (read-only). +pub(crate) fn raw_bf16_ptr(slice: &CudaSlice, stream: &Arc) -> u64 { + let (ptr, guard) = slice.device_ptr(stream); + let _no_drop = ManuallyDrop::new(guard); + ptr +} + +/// Extract raw device pointer from an f32 CudaSlice (read-only). +pub(crate) fn raw_f32_ptr(slice: &CudaSlice, stream: &Arc) -> u64 { let (ptr, guard) = slice.device_ptr(stream); let _no_drop = ManuallyDrop::new(guard); ptr @@ -860,30 +930,45 @@ pub(crate) fn raw_f32_ptr(slice: &CudaSlice, stream: &Arc bf16 dY +/// before passing to the next layer's backward_fc_layer. pub fn alloc_backward_scratch( stream: &Arc, config: &GpuDqnTrainConfig, ) -> Result<( - CudaSlice, // d_h_s2 [B, SH2] - CudaSlice, // d_h_s1 [B, SH1] - CudaSlice, // d_h_v [B, VH] - CudaSlice, // d_h_b0 [B, AH] - CudaSlice, // d_h_b1 [B, AH] - CudaSlice, // d_h_b2 [B, AH] + CudaSlice, // d_h_s2 [B, SH2] + CudaSlice, // d_h_s1 [B, SH1] + CudaSlice, // d_h_v [B, VH] + CudaSlice, // d_h_b0 [B, AH] + CudaSlice, // d_h_b1 [B, AH] + CudaSlice, // d_h_b2 [B, AH] + CudaSlice, // dy_bf16_staging (sized for the largest layer) ), MLError> { let b = config.batch_size; - let alloc = |n: usize| -> Result, MLError> { - stream.alloc_zeros::(n) - .map_err(|e| MLError::ModelError(format!("backward scratch alloc [{n}]: {e}"))) + let alloc_f32 = |n: usize| -> Result, MLError> { + stream.alloc_zeros::(n) + .map_err(|e| MLError::ModelError(format!("backward scratch f32 alloc [{n}]: {e}"))) }; + // Staging buffer sized for the widest layer that needs bf16 casting. + // max(B*SH2, B*SH1, B*VH, B*AH) + let max_staging = b * config.shared_h2 + .max(config.shared_h1) + .max(config.value_h) + .max(config.adv_h); + let staging = stream.alloc_zeros::(max_staging) + .map_err(|e| MLError::ModelError(format!("backward staging alloc [{max_staging}]: {e}")))?; + Ok(( - alloc(b * config.shared_h2)?, - alloc(b * config.shared_h1)?, - alloc(b * config.value_h)?, - alloc(b * config.adv_h)?, - alloc(b * config.adv_h)?, - alloc(b * config.adv_h)?, + alloc_f32(b * config.shared_h2)?, + alloc_f32(b * config.shared_h1)?, + alloc_f32(b * config.value_h)?, + alloc_f32(b * config.adv_h)?, + alloc_f32(b * config.adv_h)?, + alloc_f32(b * config.adv_h)?, + staging, )) } diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index fc85d0baa..1a4c4e57a 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -56,7 +56,7 @@ use crate::MLError; use super::gpu_attention::GpuAttention; use super::gpu_weights::{DuelingWeightSet, BranchingWeightSet}; use super::batched_forward::{CublasForward, bf16_weight_ptrs_from_base}; -use super::batched_backward::{CublasBackward, alloc_backward_scratch, raw_f32_ptr as bw_raw_ptr}; +use super::batched_backward::{CublasBackward, alloc_backward_scratch, raw_f32_ptr as bw_raw_f32_ptr, raw_bf16_ptr as bw_raw_bf16_ptr}; // ── Precompiled cubins (build.rs → include_bytes! → ZERO runtime nvcc) ────── static DQN_UTILITY_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/dqn_utility_kernels.cubin")); @@ -603,16 +603,20 @@ pub struct GpuDqnTrainer { // Scratch buffers for inter-layer gradient propagation during cuBLAS backward. // Sized for the widest layer at each point in the network. - /// Accumulated d_h_s2 from value head + all 3 branch heads: [B, SH2] - bw_d_h_s2: CudaSlice, - /// Upstream gradient for shared layer 1: [B, SH1] - bw_d_h_s1: CudaSlice, - /// Upstream gradient for value FC: [B, VH] - bw_d_h_v: CudaSlice, - /// Branch FC upstream gradients (one per branch): [B, AH] - bw_d_h_b0: CudaSlice, - bw_d_h_b1: CudaSlice, - bw_d_h_b2: CudaSlice, + // All f32 to prevent bf16 overflow in the backward chain. + /// Accumulated d_h_s2 from value head + all 3 branch heads: [B, SH2] — f32 + bw_d_h_s2: CudaSlice, + /// Upstream gradient for shared layer 1: [B, SH1] — f32 + bw_d_h_s1: CudaSlice, + /// Upstream gradient for value FC: [B, VH] — f32 + bw_d_h_v: CudaSlice, + /// Branch FC upstream gradients (one per branch): [B, AH] — f32 + bw_d_h_b0: CudaSlice, + bw_d_h_b1: CudaSlice, + bw_d_h_b2: CudaSlice, + /// Shared bf16 staging buffer for f32→bf16 cast at layer transitions. + /// Sized for the largest layer: max(B*SH2, B*SH1, B*VH, B*AH). + bw_dy_bf16_staging: CudaSlice, // ── Expected Q-value kernel (ad-hoc validation, not captured in CUDA Graph) ─ /// Converts C51 value+advantage logits → expected Q-values (validation path). @@ -658,11 +662,25 @@ impl GpuDqnTrainer { &self.save_h_s2 } - /// Backward gradient w.r.t. h_s2 (trunk activation). - pub fn bw_d_h_s2_buf(&self) -> &CudaSlice { + /// Backward gradient w.r.t. h_s2 (trunk activation) — f32. + pub fn bw_d_h_s2_buf(&self) -> &CudaSlice { &self.bw_d_h_s2 } + /// Cast bw_d_h_s2 (f32) → bf16 staging buffer and return a reference to it. + /// + /// Used by the attention backward kernel which expects bf16 input. + pub fn bw_d_h_s2_as_bf16(&self) -> Result<&CudaSlice, MLError> { + let n = self.config.batch_size * self.config.shared_h2; + self.cublas_backward.cast_dx_to_staging( + &self.stream, + self.ptrs.bw_d_h_s2, + bw_raw_bf16_ptr(&self.bw_dy_bf16_staging, &self.stream), + n, + )?; + Ok(&self.bw_dy_bf16_staging) + } + /// Shared trunk hidden layer 2 dimension. pub fn shared_h2(&self) -> usize { self.config.shared_h2 @@ -788,20 +806,29 @@ impl GpuDqnTrainer { } } - // ── 2. Copy IQN d_h_s2 → bw_d_h_s2 scratch ────────────────────── - // d_h_s2 is bf16 (from IQN head) — use bf16 byte size, not f32. + // ── 2. Cast IQN d_h_s2 (bf16) → bw_d_h_s2 (f32) ───────────────── + // IQN head produces bf16 gradient. Cast to f32 for the backward scratch. { - let n_bytes = b * sh2 * std::mem::size_of::(); let src = iqn_d_h_s2.raw_ptr(); let dst = self.ptrs.bw_d_h_s2; + let n_elems = (b * sh2) as i32; + let blocks = ((b * sh2 + 255) / 256) as u32; unsafe { - cudarc::driver::result::memcpy_dtod_async( - dst, src, n_bytes, self.stream.cu_stream() - ).map_err(|e| MLError::ModelError(format!("IQN d_h_s2 DtoD: {e}")))?; + self.stream + .launch_builder(&self.bf16_to_f32_kernel) + .arg(&src) + .arg(&dst) + .arg(&n_elems) + .launch(LaunchConfig { + grid_dim: (blocks, 1, 1), + block_dim: (256, 1, 1), + shared_mem_bytes: 0, + }) + .map_err(|e| MLError::ModelError(format!("IQN d_h_s2 bf16→f32: {e}")))?; } } - // ── 3. ReLU mask: bw_d_h_s2 *= (save_h_s2 > 0) ────────────────── + // ── 3. ReLU mask: bw_d_h_s2 *= (save_h_s2 > 0) (f32 dx, bf16 act) ── { let d_ptr = self.ptrs.bw_d_h_s2; let act_ptr = self.ptrs.save_h_s2; @@ -822,11 +849,15 @@ impl GpuDqnTrainer { } } - // ── 4. Backward FC layer 2: h_s1 → h_s2 (into SCRATCH) ────────── + // ── 4. Cast f32 d_h_s2 → bf16 staging, then backward FC layer 2 ── // Computes dW_s2, db_s2 into scratch (iqn_trunk_m), d_h_s1 into bw_d_h_s1. { - let dy = bw_raw_ptr(&self.bw_d_h_s2, &self.stream); - let x = bw_raw_ptr(&self.save_h_s1, &self.stream); + let staging = bw_raw_bf16_ptr(&self.bw_dy_bf16_staging, &self.stream); + self.cublas_backward.cast_dx_to_staging( + &self.stream, self.ptrs.bw_d_h_s2, staging, b * sh2, + )?; + + let x = bw_raw_bf16_ptr(&self.save_h_s1, &self.stream); let param_sizes = compute_param_sizes(&self.config); let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, ¶m_sizes); let w = w_ptrs[2]; // w_s2 @@ -836,14 +867,14 @@ impl GpuDqnTrainer { let dw = scratch_base + (w_s1_n + b_s1_n) as u64 * f32_u; // goff_w_s2 in scratch let db = dw + w_s2_n as u64 * f32_u; // goff_b_s2 in scratch - let dx = bw_raw_ptr(&self.bw_d_h_s1, &self.stream); + let dx = bw_raw_f32_ptr(&self.bw_d_h_s1, &self.stream); self.cublas_backward.backward_fc_layer( - &self.stream, dy, x, w, dw, db, dx, sh2, sh1, b, + &self.stream, staging, x, w, dw, db, dx, sh2, sh1, b, )?; } - // ── 5. ReLU mask: bw_d_h_s1 *= (save_h_s1 > 0) ────────────────── + // ── 5. ReLU mask: bw_d_h_s1 *= (save_h_s1 > 0) (f32 dx, bf16 act) ── { let d_ptr = self.ptrs.bw_d_h_s1; let act_ptr = self.ptrs.save_h_s1; @@ -864,11 +895,15 @@ impl GpuDqnTrainer { } } - // ── 6. Backward FC layer 1: states → h_s1 (into SCRATCH) ──────── + // ── 6. Cast f32 d_h_s1 → bf16 staging, then backward FC layer 1 ── // Computes dW_s1, db_s1 into scratch. dx=0 (skip input gradient). { - let dy = bw_raw_ptr(&self.bw_d_h_s1, &self.stream); - let x = bw_raw_ptr(&self.states_buf, &self.stream); + let staging = bw_raw_bf16_ptr(&self.bw_dy_bf16_staging, &self.stream); + self.cublas_backward.cast_dx_to_staging( + &self.stream, self.ptrs.bw_d_h_s1, staging, b * sh1, + )?; + + let x = bw_raw_bf16_ptr(&self.states_buf, &self.stream); let param_sizes = compute_param_sizes(&self.config); let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, ¶m_sizes); let w = w_ptrs[0]; // w_s1 @@ -879,7 +914,7 @@ impl GpuDqnTrainer { let db = scratch_base + w_s1_n as u64 * f32_u; // goff_b_s1 in scratch self.cublas_backward.backward_fc_layer( - &self.stream, dy, x, w, dw, db, 0, sh1, sd, b, + &self.stream, staging, x, w, dw, db, 0, sh1, sd, b, )?; } @@ -982,9 +1017,10 @@ impl GpuDqnTrainer { } } - // ── 2. Backward value output layer: d_logits -> d_h_v ─────────────── + // ── 2. Backward value output layer: d_logits -> d_h_v (f32) ──────── // d_logits [B, NA] x W_v2^T [VH, NA] -> d_h_v [B, VH] // Only upstream gradient (dX) is needed -- skip dW/db for value head. + // dX output is f32 (gemmex_bf16_acc_f32). { let param_sizes = compute_param_sizes(&self.config); let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, ¶m_sizes); @@ -999,7 +1035,7 @@ impl GpuDqnTrainer { )?; } - // ── 3. ReLU mask: d_h_v *= (save_h_v > 0) ───────────────────────── + // ── 3. ReLU mask: d_h_v *= (save_h_v > 0) (f32 dx, bf16 act) ────── { let d_ptr = self.ptrs.bw_d_h_v; let act_ptr = self.ptrs.save_h_v; @@ -1020,25 +1056,28 @@ impl GpuDqnTrainer { } } - // ── 4. Backward value FC layer: d_h_v -> d_h_s2 ──────────────────── - // d_h_v [B, VH] x W_v1^T [SH2, VH] -> d_h_s2 [B, SH2] + // ── 4. Cast f32 d_h_v → bf16 staging, then backward value FC → d_h_s2 (f32) ── // Only upstream gradient (dX) is needed -- skip dW/db for value head. { + let staging = bw_raw_bf16_ptr(&self.bw_dy_bf16_staging, &self.stream); + self.cublas_backward.cast_dx_to_staging( + &self.stream, self.ptrs.bw_d_h_v, staging, b * vh, + )?; + let param_sizes = compute_param_sizes(&self.config); let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, ¶m_sizes); let w_v1 = w_ptrs[4]; // W_v1 [VH, SH2] - let dy = self.ptrs.bw_d_h_v; let dx = self.ptrs.bw_d_h_s2; // launch_dx_only: computes only dX = dY @ W^T (no dW/db) self.cublas_backward.launch_dx_only( - &self.stream, dy, w_v1, dx, + &self.stream, staging, w_v1, dx, vh, sh2, b, 0.0, )?; } - // ── 5. ReLU mask: d_h_s2 *= (save_h_s2 > 0) ────────────────────── + // ── 5. ReLU mask: d_h_s2 *= (save_h_s2 > 0) (f32 dx, bf16 act) ─── { let d_ptr = self.ptrs.bw_d_h_s2; let act_ptr = self.ptrs.save_h_s2; @@ -1059,10 +1098,14 @@ impl GpuDqnTrainer { } } - // ── 6. Backward FC layer 2: h_s1 -> h_s2 (into SCRATCH) ──────────── + // ── 6. Cast f32 d_h_s2 → bf16 staging, then backward FC layer 2 ─── { - let dy = bw_raw_ptr(&self.bw_d_h_s2, &self.stream); - let x = bw_raw_ptr(&self.save_h_s1, &self.stream); + let staging = bw_raw_bf16_ptr(&self.bw_dy_bf16_staging, &self.stream); + self.cublas_backward.cast_dx_to_staging( + &self.stream, self.ptrs.bw_d_h_s2, staging, b * sh2, + )?; + + let x = bw_raw_bf16_ptr(&self.save_h_s1, &self.stream); let param_sizes = compute_param_sizes(&self.config); let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, ¶m_sizes); let w = w_ptrs[2]; // w_s2 @@ -1072,14 +1115,14 @@ impl GpuDqnTrainer { let dw = scratch_base + (w_s1_n + b_s1_n) as u64 * f32_u; let db = dw + w_s2_n as u64 * f32_u; - let dx = bw_raw_ptr(&self.bw_d_h_s1, &self.stream); + let dx = bw_raw_f32_ptr(&self.bw_d_h_s1, &self.stream); self.cublas_backward.backward_fc_layer( - &self.stream, dy, x, w, dw, db, dx, sh2, sh1, b, + &self.stream, staging, x, w, dw, db, dx, sh2, sh1, b, )?; } - // ── 7. ReLU mask: d_h_s1 *= (save_h_s1 > 0) ────────────────────── + // ── 7. ReLU mask: d_h_s1 *= (save_h_s1 > 0) (f32 dx, bf16 act) ─── { let d_ptr = self.ptrs.bw_d_h_s1; let act_ptr = self.ptrs.save_h_s1; @@ -1100,10 +1143,14 @@ impl GpuDqnTrainer { } } - // ── 8. Backward FC layer 1: states -> h_s1 (into SCRATCH) ────────── + // ── 8. Cast f32 d_h_s1 → bf16 staging, then backward FC layer 1 ─── { - let dy = bw_raw_ptr(&self.bw_d_h_s1, &self.stream); - let x = bw_raw_ptr(&self.states_buf, &self.stream); + let staging = bw_raw_bf16_ptr(&self.bw_dy_bf16_staging, &self.stream); + self.cublas_backward.cast_dx_to_staging( + &self.stream, self.ptrs.bw_d_h_s1, staging, b * sh1, + )?; + + let x = bw_raw_bf16_ptr(&self.states_buf, &self.stream); let param_sizes = compute_param_sizes(&self.config); let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, ¶m_sizes); let w = w_ptrs[0]; // w_s1 @@ -1114,7 +1161,7 @@ impl GpuDqnTrainer { let db = scratch_base + w_s1_n as u64 * f32_u; self.cublas_backward.backward_fc_layer( - &self.stream, dy, x, w, dw, db, 0, sh1, sd, b, + &self.stream, staging, x, w, dw, db, 0, sh1, sd, b, )?; } @@ -1388,6 +1435,7 @@ impl GpuDqnTrainer { // Run full backward pass with CQL logit gradients (bf16 staging) into ISOLATED scratch buffer. // Produces CQL parameter gradients WITHOUT mixing with C51's grad_buf. + let staging = self.bw_dy_bf16_staging.raw_ptr(); self.cublas_backward.backward_full( &self.stream, d_val_bf16, @@ -1399,6 +1447,7 @@ impl GpuDqnTrainer { self.cql_grad_scratch.raw_ptr(), scratch_d_h_s2, scratch_d_h_s1, scratch_d_h_v, &[scratch_d_h_b0, scratch_d_h_b1, scratch_d_h_b2], + staging, ).map_err(|e| MLError::ModelError(format!("CQL backward_full: {e}")))?; Ok(true) @@ -2145,7 +2194,7 @@ impl GpuDqnTrainer { // ── Backward scratch buffers ──────────────────────────────── // Pre-allocate inter-layer gradient buffers for the cuBLAS backward // pass. These are separate from the activation saves used in forward. - let (bw_d_h_s2, bw_d_h_s1, bw_d_h_v, bw_d_h_b0, bw_d_h_b1, bw_d_h_b2) = + let (bw_d_h_s2, bw_d_h_s1, bw_d_h_v, bw_d_h_b0, bw_d_h_b1, bw_d_h_b2, bw_dy_bf16_staging) = alloc_backward_scratch(&stream, &config) .map_err(|e| MLError::ModelError(format!("backward scratch alloc: {e}")))?; @@ -2350,6 +2399,7 @@ impl GpuDqnTrainer { bw_d_h_b0, bw_d_h_b1, bw_d_h_b2, + bw_dy_bf16_staging, expected_q_kernel, q_stats_kernel, q_stats_buf, @@ -3852,24 +3902,25 @@ impl GpuDqnTrainer { let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, ¶m_sizes); let grad_base = self.grad_buf.raw_ptr(); // f32 buffer — use raw_ptr directly - let states_ptr = bw_raw_ptr(&self.states_buf, &self.stream); - let h_s1_ptr = bw_raw_ptr(&self.save_h_s1, &self.stream); - let h_s2_ptr = bw_raw_ptr(&self.save_h_s2, &self.stream); - let h_v_ptr = bw_raw_ptr(&self.save_h_v, &self.stream); - let h_b0_ptr = bw_raw_ptr(&self.save_h_b0, &self.stream); - let h_b1_ptr = bw_raw_ptr(&self.save_h_b1, &self.stream); - let h_b2_ptr = bw_raw_ptr(&self.save_h_b2, &self.stream); + let states_ptr = bw_raw_bf16_ptr(&self.states_buf, &self.stream); + let h_s1_ptr = bw_raw_bf16_ptr(&self.save_h_s1, &self.stream); + let h_s2_ptr = bw_raw_bf16_ptr(&self.save_h_s2, &self.stream); + let h_v_ptr = bw_raw_bf16_ptr(&self.save_h_v, &self.stream); + let h_b0_ptr = bw_raw_bf16_ptr(&self.save_h_b0, &self.stream); + let h_b1_ptr = bw_raw_bf16_ptr(&self.save_h_b1, &self.stream); + let h_b2_ptr = bw_raw_bf16_ptr(&self.save_h_b2, &self.stream); - let d_h_s2_ptr = bw_raw_ptr(&self.bw_d_h_s2, &self.stream); - let d_h_s1_ptr = bw_raw_ptr(&self.bw_d_h_s1, &self.stream); - let d_h_v_ptr = bw_raw_ptr(&self.bw_d_h_v, &self.stream); - let d_h_b0_ptr = bw_raw_ptr(&self.bw_d_h_b0, &self.stream); - let d_h_b1_ptr = bw_raw_ptr(&self.bw_d_h_b1, &self.stream); - let d_h_b2_ptr = bw_raw_ptr(&self.bw_d_h_b2, &self.stream); + let d_h_s2_ptr = bw_raw_f32_ptr(&self.bw_d_h_s2, &self.stream); + let d_h_s1_ptr = bw_raw_f32_ptr(&self.bw_d_h_s1, &self.stream); + let d_h_v_ptr = bw_raw_f32_ptr(&self.bw_d_h_v, &self.stream); + let d_h_b0_ptr = bw_raw_f32_ptr(&self.bw_d_h_b0, &self.stream); + let d_h_b1_ptr = bw_raw_f32_ptr(&self.bw_d_h_b1, &self.stream); + let d_h_b2_ptr = bw_raw_f32_ptr(&self.bw_d_h_b2, &self.stream); // dL/d_logits from bf16 staging (cast from f32 by cast_d_logits_to_bf16) - let d_value_logits_ptr = bw_raw_ptr(&self.d_value_logits_bf16, &self.stream); - let d_adv_logits_ptr = bw_raw_ptr(&self.d_adv_logits_bf16, &self.stream); + let d_value_logits_ptr = bw_raw_bf16_ptr(&self.d_value_logits_bf16, &self.stream); + let d_adv_logits_ptr = bw_raw_bf16_ptr(&self.d_adv_logits_bf16, &self.stream); + let staging_ptr = bw_raw_bf16_ptr(&self.bw_dy_bf16_staging, &self.stream); let na = self.config.num_atoms; let bf16_size = std::mem::size_of::() as u64; @@ -3892,6 +3943,7 @@ impl GpuDqnTrainer { d_h_s1_ptr, d_h_v_ptr, &[d_h_b0_ptr, d_h_b1_ptr, d_h_b2_ptr], + staging_ptr, )?; Ok(()) diff --git a/crates/ml/src/cuda_pipeline/relu_mask_kernel.cu b/crates/ml/src/cuda_pipeline/relu_mask_kernel.cu index e2463525f..8f0c9136f 100644 --- a/crates/ml/src/cuda_pipeline/relu_mask_kernel.cu +++ b/crates/ml/src/cuda_pipeline/relu_mask_kernel.cu @@ -1,18 +1,20 @@ /** - * Standalone ReLU mask kernel for IQN trunk gradient. + * Standalone ReLU mask kernel for IQN trunk gradient (f32 dX path). * - * dx[i] *= (activation[i] > 0.0f) + * dx[i] = (activation[i] > 0.0f) ? dx[i] : 0.0f + * + * dx is f32 (backward dX scratch), activation is bf16 (saved forward output). * * Launch config: grid=(ceil(n/256), 1, 1), block=(256, 1, 1). */ extern "C" __global__ -void relu_mask_standalone(__nv_bfloat16* __restrict__ dx, +void relu_mask_standalone(float* __restrict__ dx, const __nv_bfloat16* __restrict__ activation, int n) { int i = blockIdx.x * blockDim.x + threadIdx.x; if (i >= n) return; - __nv_bfloat16 act = activation[i]; - if (!(act > bf16_zero())) dx[i] = bf16_zero(); + float act_f = (float)activation[i]; + if (act_f <= 0.0f) dx[i] = 0.0f; } diff --git a/crates/ml/src/trainers/dqn/fused_training.rs b/crates/ml/src/trainers/dqn/fused_training.rs index 2ebfa8d5c..c20551ccc 100644 --- a/crates/ml/src/trainers/dqn/fused_training.rs +++ b/crates/ml/src/trainers/dqn/fused_training.rs @@ -643,8 +643,10 @@ impl FusedTrainingCtx { .map_err(|e| anyhow::anyhow!("Attention forward: {e}"))?; // Attention backward: compute d_params from bw_d_h_s2 - let d_h_s2 = self.trainer.bw_d_h_s2_buf(); - attn.backward(d_h_s2, self.batch_size) + // Cast f32 bw_d_h_s2 → bf16 staging (attention kernel expects bf16) + let d_h_s2_bf16 = self.trainer.bw_d_h_s2_as_bf16() + .map_err(|e| anyhow::anyhow!("bw_d_h_s2 bf16 cast: {e}"))?; + attn.backward(d_h_s2_bf16, self.batch_size) .map_err(|e| anyhow::anyhow!("Attention backward: {e}"))?; // Attention Adam: update attention weights