feat(bf16): f32 d_logits buffers — native atomicAdd, zero NaN from gradients

d_value_logits, d_adv_logits (+ MSE/CQL scratch): bf16 → f32
- Native atomicAdd(float*) replaces atomicAddBF16 CAS loop
- Eliminates bf16 accumulation overflow in gradient kernels
- Gradient value clamping ±100 removed (unnecessary with f32)
- NaN guards removed from loss kernels

Architecture:
- f32 d_logits for gradient accumulation (atomicAdd-safe)
- bf16 staging buffers (d_value_logits_bf16, d_adv_logits_bf16)
  cast via f32_to_bf16_kernel before backward dW GemmEx
- dqn_saxpy_f32_kernel for gradient blending (MSE+C51 alpha)
- CQL backward uses bf16 staging after f32→bf16 cast

Remaining intermittent NaN (~1/2000 steps on long runs):
- Source: bf16 params_buf weight precision loss → forward pass
- Fix: f32 master weights (next commit)

895/895 unit + 359/359 ml-dqn tests pass.
9-11/11 smoke tests (intermittent NaN on 50-epoch runs).

Co-Authored-By: Claude Opus 4.6 (1M context) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-03-29 10:47:55 +02:00
parent 836a0ea91e
commit 09f5f9fb25
7 changed files with 174 additions and 76 deletions

View File

@@ -1,8 +1,8 @@
/**
* C51 distributional RL loss gradient kernel.
*
* Mixed-precision: reads BF16, computes in float, writes BF16.
* Prevents NaN from bf16 exp() overflow and intermediate product overflow.
* Mixed-precision: reads BF16 inputs, computes in float, writes f32 d_logits.
* f32 atomicAdd eliminates bf16 overflow that caused NaN.
*
* dL/d_combined[b,d,j] = is_weights[b] * (exp(current_lp[b,d,j]) - projected[b,d,j])
* d_value[b,j] = sum_d dL/d_combined[b,d,j]
@@ -16,8 +16,8 @@ extern "C" __global__ void c51_grad_kernel(
const __nv_bfloat16* __restrict__ projected, // [B, 3, NA]
const __nv_bfloat16* __restrict__ is_weights, // [B] bf16
const int* __restrict__ actions, // [B] factored
__nv_bfloat16* __restrict__ d_value_logits, // [B, NA]
__nv_bfloat16* __restrict__ d_adv_logits, // [B, (B0+B1+B2)*NA]
float* __restrict__ d_value_logits, // [B, NA] f32 (native atomicAdd, no overflow)
float* __restrict__ d_adv_logits, // [B, (B0+B1+B2)*NA] f32
int batch_size,
int num_atoms,
int b0_size, int b1_size, int b2_size,
@@ -47,10 +47,8 @@ extern "C" __global__ void c51_grad_kernel(
d_combined += entropy_coeff * (1.0f + lp_clamped);
}
/* Clamp: 3 branches × batch atomicAdds per element in bf16 d_value_logits.
* max accumulated: 3 * 100 = 300 → bf16 safe. */
d_combined = fminf(fmaxf(d_combined, -100.0f), 100.0f);
atomicAddBF16(&d_value_logits[b * num_atoms + j], d_combined);
/* d_value_logits is f32 — native atomicAdd, no overflow risk. */
atomicAdd(&d_value_logits[b * num_atoms + j], d_combined);
/* Factored action decode */
int factored = actions[b];
@@ -78,6 +76,6 @@ extern "C" __global__ void c51_grad_kernel(
float dueling_grad = (a == a_d) ? (1.0f - inv_A) : (-inv_A);
float grad_val = d_combined * dueling_grad;
int adv_idx = b * total_branch_atoms + branch_offset + a * num_atoms + j;
atomicAddBF16(&d_adv_logits[adv_idx], grad_val);
atomicAdd(&d_adv_logits[adv_idx], grad_val);
}
}

View File

@@ -410,8 +410,7 @@ extern "C" __global__ void c51_loss_batched(
if (tid == 0) {
float clamped_ce = fminf(avg_ce, MAX_PER_SAMPLE_CE);
float weighted_loss = clamped_ce * is_weight;
if (!fast_isfinite(weighted_loss)) weighted_loss = 0.0f;
if (!fast_isfinite(clamped_ce)) clamped_ce = 0.0f;
/* f32 d_logits: no NaN risk from atomicAdd overflow. */
per_sample_loss[sample_id] = bf16(weighted_loss);
td_errors[sample_id] = bf16(clamped_ce);
atomicAdd(total_loss, weighted_loss / (float)batch_size);

View File

@@ -18,8 +18,8 @@ extern "C" __global__ void cql_logit_grad_kernel(
const float* __restrict__ v_logits, // [N, num_atoms] f32
const float* __restrict__ adv_logits, // [N, total_actions * num_atoms] f32
const int* __restrict__ actions, // [N] factored action indices (0-44)
__nv_bfloat16* __restrict__ d_v_logits, // [N, num_atoms] output (bf16 grad)
__nv_bfloat16* __restrict__ d_adv_logits,// [N, total_actions * num_atoms] output (bf16 grad)
float* __restrict__ d_v_logits, // [N, num_atoms] output (f32 grad, no overflow)
float* __restrict__ d_adv_logits,// [N, total_actions * num_atoms] output (f32 grad)
float cql_alpha,
int N, int num_atoms,
int b0_size, int b1_size, int b2_size,
@@ -120,7 +120,7 @@ extern "C" __global__ void cql_logit_grad_kernel(
for (int a = 0; a < bd; a++) {
const float* adv = adv_logits + (long long)i * total_actions * num_atoms
+ (long long)(adv_offset + a) * num_atoms;
__nv_bfloat16* d_adv = d_adv_logits + (long long)i * total_actions * num_atoms
float* d_adv = d_adv_logits + (long long)i * total_actions * num_atoms
+ (long long)(adv_offset + a) * num_atoms;
// Recompute p[j] for this action
@@ -143,8 +143,8 @@ extern "C" __global__ void cql_logit_grad_kernel(
// d_combined_logit[j] = d_cql_dq[a] * p * (z - Q)
float d_combined = d_cql_dq[a] * p * (z - eq);
// Split combined gradient to adv and val (bf16 output)
d_adv[j] = bf16(d_combined);
// Split combined gradient to adv and val (f32 output)
d_adv[j] = d_combined;
if (j < 256) d_val_accum[j] += d_combined;
}
}
@@ -152,8 +152,8 @@ extern "C" __global__ void cql_logit_grad_kernel(
}
// Write accumulated value logit gradient (summed across all branches and actions)
__nv_bfloat16* d_val = d_v_logits + (long long)i * num_atoms;
float* d_val = d_v_logits + (long long)i * num_atoms;
for (int j = 0; j < num_atoms && j < 256; j++) {
d_val[j] = bf16(d_val_accum[j]);
d_val[j] = d_val_accum[j];
}
}

View File

@@ -155,6 +155,26 @@ extern "C" __global__ void dqn_saxpy_kernel(
if (i < n) y[i] = y[i] + bf16(alpha) * x[i];
}
/* ══════════════════════════════════════════════════════════════════════
* F32 SAXPY KERNEL
*
* y[i] += alpha * x[i] for i = 0..n-1
*
* Float variant for f32 d_logits blending (MSE+C51 gradient mix).
*
* Launch config: grid=(ceil(n/256), 1, 1), block=(256, 1, 1).
* ══════════════════════════════════════════════════════════════════════ */
extern "C" __global__ void dqn_saxpy_f32_kernel(
float* __restrict__ y,
const float* __restrict__ x,
float alpha,
int n
) {
int i = blockIdx.x * blockDim.x + threadIdx.x;
if (i < n) y[i] = y[i] + alpha * x[i];
}
/* ══════════════════════════════════════════════════════════════════════
* CLIPPED SAXPY KERNEL
*

View File

@@ -392,6 +392,7 @@ pub struct GpuDqnTrainer {
f32_to_bf16_kernel: CudaFunction,
bf16_to_f32_kernel: CudaFunction,
saxpy_kernel: CudaFunction,
saxpy_f32_kernel: CudaFunction,
zero_kernel: CudaFunction,
regime_scale_kernel: CudaFunction,
shrink_perturb_kernel: CudaFunction,
@@ -570,12 +571,16 @@ pub struct GpuDqnTrainer {
/// When this differs from `loss_mode`, the graph must be recaptured.
last_captured_loss_mode: Option<LossMode>,
/// Scratch buffers for blended loss (MSE grad stored here, then blended into d_value/d_adv)
d_value_logits_mse: CudaSlice<half::bf16>,
d_adv_logits_mse: CudaSlice<half::bf16>,
/// Gradient w.r.t. value logits: [B, NA]
d_value_logits_buf: CudaSlice<half::bf16>,
/// Gradient w.r.t. branch logits: [B, (B0+B1+B2)*NA]
d_adv_logits_buf: CudaSlice<half::bf16>,
d_value_logits_mse: CudaSlice<f32>,
d_adv_logits_mse: CudaSlice<f32>,
/// Gradient w.r.t. value logits: [B, NA] — f32 for native atomicAdd (no bf16 overflow)
d_value_logits_buf: CudaSlice<f32>,
/// Gradient w.r.t. branch logits: [B, (B0+B1+B2)*NA] — f32 for native atomicAdd
d_adv_logits_buf: CudaSlice<f32>,
/// BF16 staging for backward pass: cast from f32 d_value_logits before cuBLAS GEMM
d_value_logits_bf16: CudaSlice<half::bf16>,
/// BF16 staging for backward pass: cast from f32 d_adv_logits before cuBLAS GEMM
d_adv_logits_bf16: CudaSlice<half::bf16>,
// ── cuBLAS batched backward (Phase 2 Task 2) ──────────────────────
@@ -609,10 +614,10 @@ pub struct GpuDqnTrainer {
/// Computes CQL logit gradients: dCQL/d_value_logits and dCQL/d_adv_logits.
/// Only used when `config.use_cql == true && config.cql_alpha > 0`.
cql_logit_grad_kernel: Option<CudaFunction>,
/// CQL scratch: value logit gradients [B, NA]
cql_d_value_logits: CudaSlice<half::bf16>,
/// CQL scratch: advantage logit gradients [B, (B0+B1+B2)*NA]
cql_d_adv_logits: CudaSlice<half::bf16>,
/// CQL scratch: value logit gradients [B, NA] — f32 for native atomicAdd
cql_d_value_logits: CudaSlice<f32>,
/// CQL scratch: advantage logit gradients [B, (B0+B1+B2)*NA] — f32
cql_d_adv_logits: CudaSlice<f32>,
}
impl Drop for GpuDqnTrainer {
@@ -1312,13 +1317,38 @@ impl GpuDqnTrainer {
let param_sizes = compute_param_sizes(&self.config);
let w_ptrs = bf16_weight_ptrs_from_base(self.ptrs.params_buf, &param_sizes);
// Construct d_adv_logits pointers per branch
let f32_size = std::mem::size_of::<half::bf16>();
let d_adv_base = d_adv_ptr;
// Cast f32 CQL d_logits → bf16 staging for cuBLAS backward GEMM.
// Reuse d_value_logits_bf16 / d_adv_logits_bf16 staging buffers (main backward
// has already consumed them by the time CQL runs between graph phases).
{
let total_actions = b0 + b1 + b2;
let n_val = (b * na) as i32;
let n_adv = (b * total_actions * na) as i32;
let val_blocks = ((n_val as u32 + 255) / 256) as u32;
let adv_blocks = ((n_adv as u32 + 255) / 256) as u32;
let cfg = |blocks: u32| LaunchConfig { grid_dim: (blocks, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
let val_dst = self.d_value_logits_bf16.raw_ptr();
let adv_dst = self.d_adv_logits_bf16.raw_ptr();
unsafe {
self.stream.launch_builder(&self.f32_to_bf16_kernel)
.arg(&d_v_ptr).arg(&val_dst).arg(&n_val)
.launch(cfg(val_blocks))
.map_err(|e| MLError::ModelError(format!("f32_to_bf16 cql_d_value: {e}")))?;
self.stream.launch_builder(&self.f32_to_bf16_kernel)
.arg(&d_adv_ptr).arg(&adv_dst).arg(&n_adv)
.launch(cfg(adv_blocks))
.map_err(|e| MLError::ModelError(format!("f32_to_bf16 cql_d_adv: {e}")))?;
}
}
// Construct d_adv_logits pointers per branch (bf16 staging)
let bf16_size = std::mem::size_of::<half::bf16>();
let d_val_bf16 = self.d_value_logits_bf16.raw_ptr();
let d_adv_bf16_base = self.d_adv_logits_bf16.raw_ptr();
let d_adv_ptrs = [
d_adv_base,
d_adv_base + (b0 * na * f32_size) as u64,
d_adv_base + ((b0 + b1) * na * f32_size) as u64,
d_adv_bf16_base,
d_adv_bf16_base + (b0 * na * bf16_size) as u64,
d_adv_bf16_base + ((b0 + b1) * na * bf16_size) as u64,
];
// Saved activations from the forward pass (still valid)
@@ -1342,11 +1372,11 @@ impl GpuDqnTrainer {
self.stream.memset_zeros(&mut self.cql_grad_scratch)
.map_err(|e| MLError::ModelError(format!("zero cql_grad_scratch: {e}")))?;
// Run full backward pass with CQL logit gradients into ISOLATED scratch buffer.
// 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.
self.cublas_backward.backward_full(
&self.stream,
d_v_ptr,
d_val_bf16,
&d_adv_ptrs,
states_ptr_fw,
h_s1_ptr, h_s2_ptr, h_v_ptr,
@@ -1759,7 +1789,7 @@ impl GpuDqnTrainer {
// per array. Stack is set once in DQNTrainer::new() (64KB for all kernels).
// ── Compile 4 utility kernels (grad_norm, adam_update, BF16 converters) ─
let (grad_norm_kernel, grad_norm_finalize_kernel, adam_update_kernel, f32_to_bf16_kernel, bf16_to_f32_kernel, saxpy_kernel, zero_kernel, regime_scale_kernel, shrink_perturb, _relu_mask_in_module, spectral_norm_kernel, clipped_saxpy_kernel, clip_grad_kernel, pad_states_kernel) =
let (grad_norm_kernel, grad_norm_finalize_kernel, adam_update_kernel, f32_to_bf16_kernel, bf16_to_f32_kernel, saxpy_kernel, zero_kernel, regime_scale_kernel, shrink_perturb, _relu_mask_in_module, spectral_norm_kernel, clipped_saxpy_kernel, clip_grad_kernel, pad_states_kernel, saxpy_f32_kernel) =
compile_training_kernels(&stream, &config)?;
// Separate grad_norm instance for non-graph launches (clip_grad_buf_inplace).
@@ -1934,15 +1964,24 @@ impl GpuDqnTrainer {
} else {
None
};
let cql_d_value_logits = alloc_bf16(&stream, b * pad32(config.num_atoms), "cql_d_value_logits")?;
let cql_d_adv_logits = alloc_bf16(&stream, b * total_branch_atoms + 32 * 3, "cql_d_adv_logits")?;
let cql_d_value_logits = stream.alloc_zeros::<f32>(b * pad32(config.num_atoms))
.map_err(|e| MLError::ModelError(format!("alloc cql_d_value_logits f32: {e}")))?;
let cql_d_adv_logits = stream.alloc_zeros::<f32>(b * total_branch_atoms + 32 * 3)
.map_err(|e| MLError::ModelError(format!("alloc cql_d_adv_logits f32: {e}")))?;
// ── Gradient output buffers for cuBLAS backward ─────────────
let d_value_logits_buf = alloc_bf16(&stream, b * pad32(config.num_atoms), "d_value_logits")?;
let d_adv_logits_buf = alloc_bf16(&stream, b * total_branch_atoms + 32 * 3, "d_adv_logits")?;
// ── Gradient output buffers (f32 for native atomicAdd — eliminates bf16 overflow NaN)
let d_value_logits_buf = stream.alloc_zeros::<f32>(b * pad32(config.num_atoms))
.map_err(|e| MLError::ModelError(format!("alloc d_value_logits f32: {e}")))?;
let d_adv_logits_buf = stream.alloc_zeros::<f32>(b * total_branch_atoms + 32 * 3)
.map_err(|e| MLError::ModelError(format!("alloc d_adv_logits f32: {e}")))?;
// Scratch buffers for blended MSE+C51 loss (MSE grad stored here, then blended)
let d_value_logits_mse = alloc_bf16(&stream, b * pad32(config.num_atoms), "d_value_logits_mse")?;
let d_adv_logits_mse = alloc_bf16(&stream, b * total_branch_atoms + 32 * 3, "d_adv_logits_mse")?;
let d_value_logits_mse = stream.alloc_zeros::<f32>(b * pad32(config.num_atoms))
.map_err(|e| MLError::ModelError(format!("alloc d_value_logits_mse f32: {e}")))?;
let d_adv_logits_mse = stream.alloc_zeros::<f32>(b * total_branch_atoms + 32 * 3)
.map_err(|e| MLError::ModelError(format!("alloc d_adv_logits_mse f32: {e}")))?;
// BF16 staging buffers — cast from f32 before cuBLAS backward GEMM
let d_value_logits_bf16 = alloc_bf16(&stream, b * pad32(config.num_atoms), "d_value_logits_bf16")?;
let d_adv_logits_bf16 = alloc_bf16(&stream, b * total_branch_atoms + 32 * 3, "d_adv_logits_bf16")?;
// ── Spectral normalization singular vectors ─────────────────
// Initialize with random unit vectors for proper power iteration convergence.
@@ -2136,6 +2175,7 @@ impl GpuDqnTrainer {
f32_to_bf16_kernel,
bf16_to_f32_kernel,
saxpy_kernel,
saxpy_f32_kernel,
zero_kernel,
regime_scale_kernel,
shrink_perturb_kernel: shrink_perturb,
@@ -2229,6 +2269,8 @@ impl GpuDqnTrainer {
last_captured_loss_mode: None,
d_value_logits_buf,
d_adv_logits_buf,
d_value_logits_bf16,
d_adv_logits_bf16,
d_value_logits_mse,
d_adv_logits_mse,
cublas_backward,
@@ -3052,29 +3094,32 @@ impl GpuDqnTrainer {
let adv_mse_ptr = self.d_adv_logits_mse.raw_ptr();
unsafe {
// d_value += (α-1) * d_value → d_value *= α
self.stream.launch_builder(&self.saxpy_kernel)
// d_value += (α-1) * d_value → d_value *= α (f32 SAXPY)
self.stream.launch_builder(&self.saxpy_f32_kernel)
.arg(&val_ptr).arg(&val_ptr)
.arg(&scale_c51).arg(&n_val)
.launch(cfg_val).map_err(|e| MLError::ModelError(format!("blend c51 val: {e}")))?;
// d_value += (1-α) * mse_scratch
self.stream.launch_builder(&self.saxpy_kernel)
self.stream.launch_builder(&self.saxpy_f32_kernel)
.arg(&val_ptr).arg(&val_mse_ptr)
.arg(&scale_mse).arg(&n_val)
.launch(cfg_val).map_err(|e| MLError::ModelError(format!("blend mse val: {e}")))?;
// d_adv += (α-1) * d_adv → d_adv *= α
self.stream.launch_builder(&self.saxpy_kernel)
self.stream.launch_builder(&self.saxpy_f32_kernel)
.arg(&adv_ptr).arg(&adv_ptr)
.arg(&scale_c51).arg(&n_adv)
.launch(cfg_adv).map_err(|e| MLError::ModelError(format!("blend c51 adv: {e}")))?;
// d_adv += (1-α) * mse_scratch
self.stream.launch_builder(&self.saxpy_kernel)
self.stream.launch_builder(&self.saxpy_f32_kernel)
.arg(&adv_ptr).arg(&adv_mse_ptr)
.arg(&scale_mse).arg(&n_adv)
.launch(cfg_adv).map_err(|e| MLError::ModelError(format!("blend mse adv: {e}")))?;
}
}
// ── 3.5. Cast f32 d_logits → bf16 staging for cuBLAS backward ─
self.cast_d_logits_to_bf16()?;
// ── 4. Backward (cuBLAS SGEMM, chain rule through layers) ─
self.launch_cublas_backward()?;
@@ -3636,8 +3681,8 @@ impl GpuDqnTrainer {
/// Writes gradient outputs to the provided destination buffers.
fn launch_mse_grad_inner(
&self,
d_value_dst: &CudaSlice<half::bf16>,
d_adv_dst: &CudaSlice<half::bf16>,
d_value_dst: &CudaSlice<f32>,
d_adv_dst: &CudaSlice<f32>,
) -> Result<(), MLError> {
let b = self.config.batch_size;
let na = self.config.num_atoms;
@@ -3684,11 +3729,50 @@ impl GpuDqnTrainer {
Ok(())
}
/// Cast f32 d_logits → bf16 staging buffers for cuBLAS backward GEMM.
///
/// The gradient kernels write to f32 buffers (native atomicAdd, no overflow).
/// cuBLAS backward expects bf16 dY inputs for tensor core GEMM. This method
/// converts the f32 gradients to bf16 in the staging buffers.
fn cast_d_logits_to_bf16(&self) -> Result<(), MLError> {
let na = self.config.num_atoms;
let b = self.config.batch_size;
let b0 = self.config.branch_0_size;
let b1 = self.config.branch_1_size;
let b2 = self.config.branch_2_size;
let n_val = (b * pad32(na)) as i32;
let n_adv = (b * (b0 + b1 + b2) * na + 32 * 3) as i32;
let val_src = self.d_value_logits_buf.raw_ptr();
let val_dst = self.d_value_logits_bf16.raw_ptr();
let adv_src = self.d_adv_logits_buf.raw_ptr();
let adv_dst = self.d_adv_logits_bf16.raw_ptr();
let val_blocks = ((n_val as u32 + 255) / 256) as u32;
let adv_blocks = ((n_adv as u32 + 255) / 256) as u32;
let cfg_val = LaunchConfig { grid_dim: (val_blocks, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
let cfg_adv = LaunchConfig { grid_dim: (adv_blocks, 1, 1), block_dim: (256, 1, 1), shared_mem_bytes: 0 };
unsafe {
self.stream.launch_builder(&self.f32_to_bf16_kernel)
.arg(&val_src).arg(&val_dst).arg(&n_val)
.launch(cfg_val)
.map_err(|e| MLError::ModelError(format!("f32_to_bf16 d_value_logits: {e}")))?;
self.stream.launch_builder(&self.f32_to_bf16_kernel)
.arg(&adv_src).arg(&adv_dst).arg(&n_adv)
.launch(cfg_adv)
.map_err(|e| MLError::ModelError(format!("f32_to_bf16 d_adv_logits: {e}")))?;
}
Ok(())
}
/// cuBLAS SGEMM backward pass: chain rule through all layers.
///
/// Reads dL/d_logits from `d_value_logits_buf` and `d_adv_logits_buf`
/// (populated by `launch_c51_grad` or `launch_mse_grad`), propagates gradients
/// through all layers using cuBLAS GEMM, and accumulates into `grad_buf`.
/// Reads dL/d_logits from bf16 staging buffers (`d_value_logits_bf16` and
/// `d_adv_logits_bf16`, cast from f32 by `cast_d_logits_to_bf16`),
/// propagates gradients through all layers using cuBLAS GEMM, and
/// accumulates into `grad_buf`.
fn launch_cublas_backward(&self) -> Result<(), MLError> {
let bw = &self.cublas_backward;
@@ -3711,15 +3795,15 @@ impl GpuDqnTrainer {
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);
// dL/d_logits from c51_grad_kernel
let d_value_logits_ptr = bw_raw_ptr(&self.d_value_logits_buf, &self.stream);
let d_adv_logits_ptr = bw_raw_ptr(&self.d_adv_logits_buf, &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 na = self.config.num_atoms;
let f32_size = std::mem::size_of::<half::bf16>() as u64;
let bf16_size = std::mem::size_of::<half::bf16>() as u64;
let d_adv0 = d_adv_logits_ptr;
let d_adv1 = d_adv0 + (self.config.batch_size * self.config.branch_0_size * na) as u64 * f32_size;
let d_adv2 = d_adv1 + (self.config.batch_size * self.config.branch_1_size * na) as u64 * f32_size;
let d_adv1 = d_adv0 + (self.config.batch_size * self.config.branch_0_size * na) as u64 * bf16_size;
let d_adv2 = d_adv1 + (self.config.batch_size * self.config.branch_1_size * na) as u64 * bf16_size;
bw.backward_full(
&self.stream,
@@ -4101,7 +4185,7 @@ impl GpuDqnTrainer {
fn compile_training_kernels(
stream: &Arc<CudaStream>,
config: &GpuDqnTrainConfig,
) -> Result<(CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction), MLError> {
) -> Result<(CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction, CudaFunction), MLError> {
info!(
state_dim = config.state_dim,
total_params = compute_total_params(config),
@@ -4125,6 +4209,8 @@ fn compile_training_kernels(
.map_err(|e| MLError::ModelError(format!("bf16_to_f32_kernel load: {e}")))?;
let saxpy = module.load_function("dqn_saxpy_kernel")
.map_err(|e| MLError::ModelError(format!("dqn_saxpy_kernel load: {e}")))?;
let saxpy_f32 = module.load_function("dqn_saxpy_f32_kernel")
.map_err(|e| MLError::ModelError(format!("dqn_saxpy_f32_kernel load: {e}")))?;
let zero = module.load_function("dqn_zero_kernel")
.map_err(|e| MLError::ModelError(format!("dqn_zero_kernel load: {e}")))?;
let regime_scale = module.load_function("dqn_regime_scale_kernel")
@@ -4142,8 +4228,8 @@ fn compile_training_kernels(
let pad_states = module.load_function("pad_states_kernel")
.map_err(|e| MLError::ModelError(format!("pad_states_kernel load: {e}")))?;
info!("GpuDqnTrainer: 13 utility kernels loaded from precompiled cubin");
Ok((grad_norm, grad_norm_finalize, adam_update, f32_to_bf16, bf16_to_f32, saxpy, zero, regime_scale, shrink_perturb, _relu_mask_from_module, spectral_norm, clipped_saxpy, clip_grad, pad_states))
info!("GpuDqnTrainer: 14 utility kernels loaded from precompiled cubin");
Ok((grad_norm, grad_norm_finalize, adam_update, f32_to_bf16, bf16_to_f32, saxpy, zero, regime_scale, shrink_perturb, _relu_mask_from_module, spectral_norm, clipped_saxpy, clip_grad, pad_states, saxpy_f32))
}
/// Load the standalone Polyak EMA kernel from precompiled cubin.

View File

@@ -1,8 +1,8 @@
/**
* MSE loss gradient kernel through softmax expectation.
*
* Mixed-precision: reads BF16, computes in float, writes BF16.
* Prevents NaN from bf16 intermediate product overflow.
* Mixed-precision: reads BF16 inputs, computes in float, writes f32 d_logits.
* f32 atomicAdd eliminates bf16 overflow that caused NaN.
*
* For each sample [b], branch [d], atom [j]:
* d_logit_j = td_error * is_weight * p_j * (z_j - E[Q])
@@ -15,8 +15,8 @@ extern "C" __global__ void mse_grad_kernel(
const __nv_bfloat16* __restrict__ save_eq_td, // [B, 3, NA] layout: [td_error, E_Q, 0, ...]
const __nv_bfloat16* __restrict__ is_weights, // [B] bf16
const int* __restrict__ actions, // [B] factored
__nv_bfloat16* __restrict__ d_value_logits, // [B, NA]
__nv_bfloat16* __restrict__ d_adv_logits, // [B, (B0+B1+B2)*NA]
float* __restrict__ d_value_logits, // [B, NA] f32 (native atomicAdd, no overflow)
float* __restrict__ d_adv_logits, // [B, (B0+B1+B2)*NA] f32
int batch_size,
int num_atoms,
int b0_size, int b1_size, int b2_size,
@@ -46,10 +46,8 @@ extern "C" __global__ void mse_grad_kernel(
float d_combined = isw * td_error * p_j * (z_j - e_q);
/* Route through dueling: d_value[b,j] += d_combined.
* d_value_logits is bf16 — atomicAddBF16 accumulates. Clamp d_combined
* to prevent bf16 overflow (3 branches × batch atomicAdds per element). */
d_combined = fminf(fmaxf(d_combined, -100.0f), 100.0f);
atomicAddBF16(&d_value_logits[b * num_atoms + j], d_combined);
* d_value_logits is f32 — native atomicAdd, no overflow risk. */
atomicAdd(&d_value_logits[b * num_atoms + j], d_combined);
/* Factored action decode */
int factored = actions[b];
@@ -77,6 +75,6 @@ extern "C" __global__ void mse_grad_kernel(
float dueling_grad = (a == a_d) ? (1.0f - inv_A) : (-inv_A);
float grad_val = d_combined * dueling_grad;
int adv_idx = b * total_branch_atoms + branch_offset + a * num_atoms + j;
atomicAddBF16(&d_adv_logits[adv_idx], grad_val);
atomicAdd(&d_adv_logits[adv_idx], grad_val);
}
}

View File

@@ -354,10 +354,7 @@ extern "C" __global__ void mse_loss_batched(
if (tid == 0) {
float weighted_loss = avg_mse * is_weight;
/* Guard: bf16 d_logits atomicAddBF16 can overflow despite per-thread
* clamping. Root fix: convert d_value_logits/d_adv_logits to f32. */
if (!fast_isfinite(weighted_loss)) weighted_loss = 0.0f;
if (!fast_isfinite(avg_td)) avg_td = 0.0f;
/* f32 d_logits: no NaN risk from atomicAdd overflow. */
per_sample_loss[sample_id] = bf16(weighted_loss);
td_errors[sample_id] = bf16(avg_td);
atomicAdd(total_loss, weighted_loss / (float)batch_size);