feat: add f32 dY backward chain — no bf16 precision loss
New methods: backward_fc_layer_f32dy, backward_fc_layer_lda_f32dy, launch_dw_only_f32dy, launch_dx_only_f32dy, launch_bias_grad_f32. Uses gemmex_f32b_bf16a_acc_f32 (f32 dY + bf16 activations → f32 output). backward_full now passes f32 dX directly between layers — no bf16 cast. Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
@@ -129,6 +129,10 @@ pub struct CublasBackward {
|
||||
/// `f32_to_bf16_cast_kernel(dst, src, n)` — cast f32 dX scratch to bf16 staging.
|
||||
f32_to_bf16_cast_kernel: CudaFunction,
|
||||
|
||||
/// `bias_grad_reduce_f32_kernel(dy, db, out_dim, batch_size)` — reduce f32 dY over batch.
|
||||
/// Variant of `bias_grad_kernel` that reads f32 dY instead of bf16.
|
||||
bias_grad_f32_kernel: CudaFunction,
|
||||
|
||||
// ── Network dimensions (baked at construction) ──
|
||||
batch_size: usize,
|
||||
state_dim: usize,
|
||||
@@ -168,13 +172,14 @@ impl CublasBackward {
|
||||
}
|
||||
|
||||
// ── Compile helper kernels ──────────────────────────────────
|
||||
let (relu_mask_kernel, bias_grad_kernel, f32_to_bf16_cast_kernel) = compile_backward_kernels(stream)?;
|
||||
let (relu_mask_kernel, bias_grad_kernel, f32_to_bf16_cast_kernel, bias_grad_f32_kernel) = compile_backward_kernels(stream)?;
|
||||
|
||||
Ok(Self {
|
||||
handle: SendSyncCublasHandle(raw_handle),
|
||||
relu_mask_kernel,
|
||||
bias_grad_kernel,
|
||||
f32_to_bf16_cast_kernel,
|
||||
bias_grad_f32_kernel,
|
||||
batch_size: config.batch_size,
|
||||
state_dim: config.state_dim,
|
||||
state_dim_padded: (config.state_dim + 127) & !127,
|
||||
@@ -293,6 +298,48 @@ impl CublasBackward {
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// GemmEx with f32 B (dY) + bf16 A (activations/weights) → f32 C.
|
||||
/// Full f32 precision for the backward chain — no bf16 truncation on dY.
|
||||
/// Used instead of gemmex_bf16_acc_f32 to prevent gradient precision loss
|
||||
/// at large batch sizes with mean-reduced gradients.
|
||||
fn gemmex_f32b_bf16a_acc_f32(
|
||||
&self,
|
||||
transa: cublas_sys::cublasOperation_t,
|
||||
transb: cublas_sys::cublasOperation_t,
|
||||
m: i32, n: i32, k: i32,
|
||||
alpha: f32,
|
||||
a: u64, lda: i32, // bf16 (activations or weights)
|
||||
b_ptr: u64, ldb: i32, // f32 (upstream gradient dY)
|
||||
beta: f32,
|
||||
c: u64, ldc: i32, // f32 (output: dW or dX)
|
||||
label: &str,
|
||||
) -> Result<(), MLError> {
|
||||
unsafe {
|
||||
let status = cublas_sys::cublasGemmEx(
|
||||
self.handle.0,
|
||||
transa, transb,
|
||||
m, n, k,
|
||||
&alpha as *const f32 as *const std::ffi::c_void,
|
||||
a as *const std::ffi::c_void,
|
||||
cublas_sys::cudaDataType_t::CUDA_R_16BF, // A: bf16 (activations)
|
||||
lda,
|
||||
b_ptr as *const std::ffi::c_void,
|
||||
cublas_sys::cudaDataType_t::CUDA_R_32F, // B: f32 (dY gradient)
|
||||
ldb,
|
||||
&beta as *const f32 as *const std::ffi::c_void,
|
||||
c as *mut std::ffi::c_void,
|
||||
cublas_sys::cudaDataType_t::CUDA_R_32F, // C: f32 (output)
|
||||
ldc,
|
||||
cublas_sys::cublasComputeType_t::CUBLAS_COMPUTE_32F, // Full f32 compute
|
||||
cublas_sys::cublasGemmAlgo_t::CUBLAS_GEMM_DEFAULT_TENSOR_OP,
|
||||
);
|
||||
if status != cublas_sys::cublasStatus_t::CUBLAS_STATUS_SUCCESS {
|
||||
return Err(MLError::ModelError(format!("cublasGemmEx backward f32dY {label}: {status:?}")));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ══════════════════════════════════════════════════════════════════════════
|
||||
// Public building blocks
|
||||
// ══════════════════════════════════════════════════════════════════════════
|
||||
@@ -503,8 +550,9 @@ impl CublasBackward {
|
||||
scratch_d_h_s1: u64, // [B, SH1] — f32
|
||||
scratch_d_h_v: u64, // [B, VH] — f32
|
||||
scratch_d_h_b: &[u64; 4], // [B, AH] each — f32
|
||||
// Shared bf16 staging buffer (overwritten per layer)
|
||||
staging_bf16: u64, // max(B*SH2, B*SH1, B*VH, B*AH) — bf16
|
||||
// Shared bf16 staging buffer — no longer used by backward_full (f32 chain),
|
||||
// but kept in the signature for API compatibility with external callers.
|
||||
_staging_bf16: u64, // max(B*SH2, B*SH1, B*VH, B*AH) — bf16
|
||||
// #31 Bottleneck: optional f32 dX output for shared layer 1.
|
||||
// When non-zero, computes d_loss/d_bn_concat (upstream gradient for bottleneck).
|
||||
// When 0, dX is not computed (states not trainable — no bottleneck).
|
||||
@@ -659,14 +707,13 @@ impl CublasBackward {
|
||||
b * self.adv_h,
|
||||
)?;
|
||||
|
||||
// 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)?;
|
||||
// f32 dY chain: pass f32 scratch directly — no bf16 cast needed.
|
||||
|
||||
// dW for branch FC: dW[AH, SH2] += staging_bf16^T @ h_s2
|
||||
// dW for branch FC: dW[AH, SH2] += dY^T @ h_s2
|
||||
// (dX is computed separately below to allow accumulation)
|
||||
self.launch_dw_only(
|
||||
self.launch_dw_only_f32dy(
|
||||
stream,
|
||||
staging_bf16, // dY [B, AH] — bf16 staging
|
||||
scratch_d_h_b[d], // dY [B, AH] — f32
|
||||
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]
|
||||
@@ -675,13 +722,12 @@ impl CublasBackward {
|
||||
b,
|
||||
)?;
|
||||
|
||||
// dX for branch FC: d_h_s2 += staging_bf16 @ W_bdk_fc
|
||||
// dX for branch FC: d_h_s2 += dY @ 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(
|
||||
self.launch_dx_only_f32dy(
|
||||
stream,
|
||||
staging_bf16, // dY [B, AH] — bf16 staging
|
||||
scratch_d_h_b[d], // dY [B, AH] — f32
|
||||
w_fc, // W [AH, SH2]
|
||||
scratch_d_h_s2, // dX [B, SH2] — f32
|
||||
self.adv_h, // out_dim
|
||||
@@ -714,13 +760,12 @@ impl CublasBackward {
|
||||
// ReLU mask on f32 dX_v
|
||||
self.relu_mask(stream, scratch_d_h_v, save_h_v, b * self.value_h)?;
|
||||
|
||||
// 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)?;
|
||||
// f32 dY chain: pass f32 scratch directly — no bf16 cast needed.
|
||||
|
||||
// dW for value FC: dW[VH, SH2] += staging_bf16^T @ h_s2
|
||||
self.launch_dw_only(
|
||||
// dW for value FC: dW[VH, SH2] += dY^T @ h_s2
|
||||
self.launch_dw_only_f32dy(
|
||||
stream,
|
||||
staging_bf16, // dY [B, VH] — bf16 staging
|
||||
scratch_d_h_v, // dY [B, VH] — f32
|
||||
save_h_s2,
|
||||
grad_buf_base + goff_w_v1,
|
||||
grad_buf_base + goff_b_v1,
|
||||
@@ -730,9 +775,9 @@ impl CublasBackward {
|
||||
)?;
|
||||
|
||||
// dX for value FC: accumulated into f32 scratch_d_h_s2 (beta=1.0 — branches already wrote)
|
||||
self.launch_dx_only(
|
||||
self.launch_dx_only_f32dy(
|
||||
stream,
|
||||
staging_bf16, // dY [B, VH] — bf16 staging
|
||||
scratch_d_h_v, // dY [B, VH] — f32
|
||||
w_ptrs[4], // W_v1 [VH, SH2]
|
||||
scratch_d_h_s2,
|
||||
self.value_h,
|
||||
@@ -749,12 +794,10 @@ impl CublasBackward {
|
||||
// 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(
|
||||
// f32 dY chain: pass f32 scratch directly — no bf16 cast needed.
|
||||
self.backward_fc_layer_f32dy(
|
||||
stream,
|
||||
staging_bf16, // dY [B, SH2] — bf16 staging
|
||||
scratch_d_h_s2, // dY [B, SH2] — f32
|
||||
save_h_s1,
|
||||
w_ptrs[2], // W_s2 [SH2, SH1]
|
||||
grad_buf_base + goff_w_s2,
|
||||
@@ -769,18 +812,16 @@ impl CublasBackward {
|
||||
// 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)?;
|
||||
|
||||
// f32 dY chain: pass f32 scratch directly — no bf16 cast needed.
|
||||
// dW for shared layer 1: input is states (or bn_concat when bottleneck active).
|
||||
// States buffer has padded stride pad128(state_dim) for CUTLASS K-tile
|
||||
// alignment. Use backward_fc_layer_lda with the padded stride so the
|
||||
// alignment. Use backward_fc_layer_lda_f32dy with the padded stride so the
|
||||
// backward GEMM reads the correct row offsets.
|
||||
// s1_dx_output: 0 = no dX (original, states not trainable),
|
||||
// non-zero = compute dX for bottleneck backward chain rule.
|
||||
self.backward_fc_layer_lda(
|
||||
self.backward_fc_layer_lda_f32dy(
|
||||
stream,
|
||||
staging_bf16, // dY [B, SH1] — bf16 staging
|
||||
scratch_d_h_s1, // dY [B, SH1] — f32
|
||||
states, // X [B, s1_input_dim] — padded stride
|
||||
w_ptrs[0], // W_s1 [SH1, s1_input_dim]
|
||||
grad_buf_base + goff_w_s1,
|
||||
@@ -930,6 +971,197 @@ impl CublasBackward {
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
// ══════════════════════════════════════════════════════════════════════════
|
||||
// F32 dY variants — full precision backward chain (no bf16 truncation)
|
||||
// ══════════════════════════════════════════════════════════════════════════
|
||||
|
||||
/// Launch `bias_grad_reduce_f32_kernel`: db[j] += sum_b(dY[b * out_dim + j]).
|
||||
/// Variant that reads f32 dY instead of bf16 — used by the f32 backward chain.
|
||||
fn launch_bias_grad_f32(
|
||||
&self,
|
||||
stream: &Arc<CudaStream>,
|
||||
dy: u64, // f32 pointer
|
||||
db: u64,
|
||||
out_dim: usize,
|
||||
batch: usize,
|
||||
) -> Result<(), MLError> {
|
||||
let out_dim_i32 = out_dim as i32;
|
||||
let batch_i32 = batch as i32;
|
||||
let blocks = ((out_dim + 255) / 256) as u32;
|
||||
|
||||
unsafe {
|
||||
stream
|
||||
.launch_builder(&self.bias_grad_f32_kernel)
|
||||
.arg(&dy)
|
||||
.arg(&db)
|
||||
.arg(&out_dim_i32)
|
||||
.arg(&batch_i32)
|
||||
.launch(LaunchConfig {
|
||||
grid_dim: (blocks, 1, 1),
|
||||
block_dim: (256, 1, 1),
|
||||
shared_mem_bytes: 0,
|
||||
})
|
||||
.map_err(|e| MLError::ModelError(format!("bias_grad_reduce_f32_kernel: {e}")))?;
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Backward through one fully-connected layer with f32 dY (no bf16 truncation).
|
||||
///
|
||||
/// Identical to `backward_fc_layer` but uses `gemmex_f32b_bf16a_acc_f32`
|
||||
/// (f32 B + bf16 A → f32 C) and `launch_bias_grad_f32` so the upstream
|
||||
/// gradient stays in f32 throughout the entire backward chain.
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub fn backward_fc_layer_f32dy(
|
||||
&self,
|
||||
stream: &Arc<CudaStream>,
|
||||
dy: u64, // f32 pointer (not bf16!)
|
||||
x: u64, // bf16 (saved activations)
|
||||
w: u64, // bf16 (weights)
|
||||
dw: u64, // f32 (grad_buf)
|
||||
db: u64, // f32 (grad_buf)
|
||||
dx: u64, // f32 (scratch) — 0 means skip
|
||||
out_dim: usize,
|
||||
in_dim: usize,
|
||||
batch: usize,
|
||||
) -> Result<(), MLError> {
|
||||
// ── Weight gradient: dW[out, in] += dY^T @ X ──
|
||||
// A=X (bf16), B=dY (f32) → C=dW (f32)
|
||||
self.gemmex_f32b_bf16a_acc_f32(
|
||||
cublas_sys::cublasOperation_t::CUBLAS_OP_N,
|
||||
cublas_sys::cublasOperation_t::CUBLAS_OP_T,
|
||||
in_dim as i32, out_dim as i32, batch as i32,
|
||||
1.0, x, in_dim as i32,
|
||||
dy, out_dim as i32,
|
||||
1.0, dw, in_dim as i32,
|
||||
"fc_dW_f32dy",
|
||||
)?;
|
||||
|
||||
// ── Bias gradient: db[out] += sum_b(dY[b, out]) ──
|
||||
self.launch_bias_grad_f32(stream, dy, db, out_dim, batch)?;
|
||||
|
||||
// ── Upstream gradient: dX[B, in] = dY[B, out] @ W[out, in] ──
|
||||
// A=W (bf16), B=dY (f32) → C=dX (f32)
|
||||
if dx != 0 {
|
||||
self.gemmex_f32b_bf16a_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,
|
||||
1.0, w, in_dim as i32,
|
||||
dy, out_dim as i32,
|
||||
0.0, dx, in_dim as i32,
|
||||
"fc_dX_f32dy",
|
||||
)?;
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Like `backward_fc_layer_f32dy` but with a custom leading dimension for X.
|
||||
/// Used when X has a padded row stride (e.g. states buffer padded to
|
||||
/// pad128(state_dim) for CUTLASS K-tile).
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub fn backward_fc_layer_lda_f32dy(
|
||||
&self,
|
||||
stream: &Arc<CudaStream>,
|
||||
dy: u64, // f32 pointer (not bf16!)
|
||||
x: u64, // bf16 (saved activations, padded stride)
|
||||
w: u64, // bf16 (weights)
|
||||
dw: u64, // f32 (grad_buf)
|
||||
db: u64, // f32 (grad_buf)
|
||||
dx: u64, // f32 (scratch) — 0 means skip
|
||||
out_dim: usize,
|
||||
in_dim: usize,
|
||||
x_lda: usize,
|
||||
batch: usize,
|
||||
) -> Result<(), MLError> {
|
||||
// Weight gradient: dW[out, in] += dY^T @ X (X has padded stride x_lda)
|
||||
self.gemmex_f32b_bf16a_acc_f32(
|
||||
cublas_sys::cublasOperation_t::CUBLAS_OP_N,
|
||||
cublas_sys::cublasOperation_t::CUBLAS_OP_T,
|
||||
in_dim as i32, out_dim as i32, batch as i32,
|
||||
1.0, x, x_lda as i32,
|
||||
dy, out_dim as i32,
|
||||
1.0, dw, in_dim as i32,
|
||||
"fc_dW_lda_f32dy",
|
||||
)?;
|
||||
|
||||
self.launch_bias_grad_f32(stream, dy, db, out_dim, batch)?;
|
||||
|
||||
if dx != 0 {
|
||||
// dX has stride x_lda (same as input for dimensional consistency).
|
||||
self.gemmex_f32b_bf16a_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,
|
||||
1.0, w, in_dim as i32,
|
||||
dy, out_dim as i32,
|
||||
0.0, dx, x_lda as i32,
|
||||
"fc_dX_lda_f32dy",
|
||||
)?;
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
/// Compute only weight gradient + bias gradient with f32 dY (no upstream dX).
|
||||
///
|
||||
/// F32 dY variant of `launch_dw_only`. Used for branch FC layers where
|
||||
/// the upstream dX is computed separately via `launch_dx_only_f32dy`.
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub(crate) fn launch_dw_only_f32dy(
|
||||
&self,
|
||||
stream: &Arc<CudaStream>,
|
||||
dy: u64, // f32 pointer
|
||||
x: u64, // bf16 (saved activations)
|
||||
dw: u64, // f32 (grad_buf)
|
||||
db: u64, // f32 (grad_buf)
|
||||
out_dim: usize,
|
||||
in_dim: usize,
|
||||
batch: usize,
|
||||
) -> Result<(), MLError> {
|
||||
// dW[out, in] += dY^T @ X — A=X (bf16), B=dY (f32) → f32 grad_buf, beta=1.0
|
||||
self.gemmex_f32b_bf16a_acc_f32(
|
||||
cublas_sys::cublasOperation_t::CUBLAS_OP_N,
|
||||
cublas_sys::cublasOperation_t::CUBLAS_OP_T,
|
||||
in_dim as i32, out_dim as i32, batch as i32,
|
||||
1.0, x, in_dim as i32,
|
||||
dy, out_dim as i32,
|
||||
1.0, dw, in_dim as i32,
|
||||
"dW_only_f32dy",
|
||||
)?;
|
||||
|
||||
self.launch_bias_grad_f32(stream, dy, db, out_dim, batch)
|
||||
}
|
||||
|
||||
/// Compute only upstream gradient dX with f32 dY, with configurable beta.
|
||||
///
|
||||
/// F32 dY variant of `launch_dx_only`. A=W (bf16), B=dY (f32) → f32 dX.
|
||||
#[allow(clippy::too_many_arguments)]
|
||||
pub(crate) fn launch_dx_only_f32dy(
|
||||
&self,
|
||||
_stream: &Arc<CudaStream>,
|
||||
dy: u64, // f32 pointer
|
||||
w: u64, // bf16 (weights)
|
||||
dx: u64, // f32 (output scratch)
|
||||
out_dim: usize,
|
||||
in_dim: usize,
|
||||
batch: usize,
|
||||
beta: f32,
|
||||
) -> Result<(), MLError> {
|
||||
// dX[B, in] = dY[B, out] @ W[out, in] — A=W (bf16), B=dY (f32) → f32
|
||||
self.gemmex_f32b_bf16a_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,
|
||||
1.0, w, in_dim as i32,
|
||||
dy, out_dim as i32,
|
||||
beta, dx, in_dim as i32,
|
||||
"dX_only_f32dy",
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
// ── Kernel compilation ────────────────────────────────────────────────────────
|
||||
@@ -949,7 +1181,7 @@ static BACKWARD_CUBIN: &[u8] = include_bytes!(concat!(env!("OUT_DIR"), "/backwar
|
||||
|
||||
fn compile_backward_kernels(
|
||||
stream: &Arc<CudaStream>,
|
||||
) -> Result<(CudaFunction, CudaFunction, CudaFunction), MLError> {
|
||||
) -> Result<(CudaFunction, CudaFunction, CudaFunction, CudaFunction), MLError> {
|
||||
let context = stream.context();
|
||||
let module = context
|
||||
.load_cubin(BACKWARD_CUBIN.to_vec())
|
||||
@@ -964,8 +1196,11 @@ fn compile_backward_kernels(
|
||||
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}")))?;
|
||||
let bias_grad_f32 = module
|
||||
.load_function("bias_grad_reduce_f32_kernel")
|
||||
.map_err(|e| MLError::ModelError(format!("bias_grad_reduce_f32_kernel load: {e}")))?;
|
||||
|
||||
Ok((relu_mask, bias_grad, f32_to_bf16_cast))
|
||||
Ok((relu_mask, bias_grad, f32_to_bf16_cast, bias_grad_f32))
|
||||
}
|
||||
|
||||
// ── Raw device pointer helpers ────────────────────────────────────────────────
|
||||
|
||||
Reference in New Issue
Block a user