feat(bf16): f32 backward dX scratch + bf16 staging — eliminates dX truncation

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<half::bf16> → CudaSlice<f32>
- 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) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-03-29 12:17:04 +02:00
parent 05a184168f
commit 875f263ec9
5 changed files with 290 additions and 130 deletions

View File

@@ -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) */

View File

@@ -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<CudaStream>,
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<CudaStream>,
) -> 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<half::bf16>, stream: &Arc<CudaStream>) -> u64 {
/// Extract raw device pointer from a bf16 CudaSlice (read-only).
pub(crate) fn raw_bf16_ptr(slice: &CudaSlice<half::bf16>, stream: &Arc<CudaStream>) -> 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<f32>, stream: &Arc<CudaStream>) -> 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<half::bf16>, stream: &Arc<CudaStream
/// Allocate zero-initialised scratch buffers for the cuBLAS backward pass.
///
/// Called from `GpuDqnTrainer::new()` alongside the other buffer allocations.
/// Returns `(d_h_s2, d_h_s1, d_h_v, d_h_b0, d_h_b1, d_h_b2)`.
/// Returns `(d_h_s2, d_h_s1, d_h_v, d_h_b0, d_h_b1, d_h_b2, dy_bf16_staging)`.
///
/// The dX scratch buffers are **f32** to prevent bf16 overflow in the backward
/// chain. A shared bf16 staging buffer is used to cast f32 dX -> bf16 dY
/// before passing to the next layer's backward_fc_layer.
pub fn alloc_backward_scratch(
stream: &Arc<CudaStream>,
config: &GpuDqnTrainConfig,
) -> Result<(
CudaSlice<half::bf16>, // d_h_s2 [B, SH2]
CudaSlice<half::bf16>, // d_h_s1 [B, SH1]
CudaSlice<half::bf16>, // d_h_v [B, VH]
CudaSlice<half::bf16>, // d_h_b0 [B, AH]
CudaSlice<half::bf16>, // d_h_b1 [B, AH]
CudaSlice<half::bf16>, // d_h_b2 [B, AH]
CudaSlice<f32>, // d_h_s2 [B, SH2]
CudaSlice<f32>, // d_h_s1 [B, SH1]
CudaSlice<f32>, // d_h_v [B, VH]
CudaSlice<f32>, // d_h_b0 [B, AH]
CudaSlice<f32>, // d_h_b1 [B, AH]
CudaSlice<f32>, // d_h_b2 [B, AH]
CudaSlice<half::bf16>, // dy_bf16_staging (sized for the largest layer)
), MLError> {
let b = config.batch_size;
let alloc = |n: usize| -> Result<CudaSlice<half::bf16>, MLError> {
stream.alloc_zeros::<half::bf16>(n)
.map_err(|e| MLError::ModelError(format!("backward scratch alloc [{n}]: {e}")))
let alloc_f32 = |n: usize| -> Result<CudaSlice<f32>, MLError> {
stream.alloc_zeros::<f32>(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::<half::bf16>(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,
))
}

View File

@@ -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<half::bf16>,
/// Upstream gradient for shared layer 1: [B, SH1]
bw_d_h_s1: CudaSlice<half::bf16>,
/// Upstream gradient for value FC: [B, VH]
bw_d_h_v: CudaSlice<half::bf16>,
/// Branch FC upstream gradients (one per branch): [B, AH]
bw_d_h_b0: CudaSlice<half::bf16>,
bw_d_h_b1: CudaSlice<half::bf16>,
bw_d_h_b2: CudaSlice<half::bf16>,
// 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<f32>,
/// Upstream gradient for shared layer 1: [B, SH1] — f32
bw_d_h_s1: CudaSlice<f32>,
/// Upstream gradient for value FC: [B, VH] — f32
bw_d_h_v: CudaSlice<f32>,
/// Branch FC upstream gradients (one per branch): [B, AH] — f32
bw_d_h_b0: CudaSlice<f32>,
bw_d_h_b1: CudaSlice<f32>,
bw_d_h_b2: CudaSlice<f32>,
/// 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<half::bf16>,
// ── 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<half::bf16> {
/// Backward gradient w.r.t. h_s2 (trunk activation) — f32.
pub fn bw_d_h_s2_buf(&self) -> &CudaSlice<f32> {
&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<half::bf16>, 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::<half::bf16>();
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, &param_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, &param_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, &param_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, &param_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, &param_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, &param_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, &param_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::<half::bf16>() 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(())

View File

@@ -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;
}

View File

@@ -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