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:
@@ -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) */
|
||||
|
||||
@@ -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,
|
||||
))
|
||||
}
|
||||
|
||||
@@ -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, ¶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::<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(())
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user