From 09f5f9fb255c01575e220cc37ef590f23c8760eb Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sun, 29 Mar 2026 10:47:55 +0200 Subject: [PATCH] =?UTF-8?q?feat(bf16):=20f32=20d=5Flogits=20buffers=20?= =?UTF-8?q?=E2=80=94=20native=20atomicAdd,=20zero=20NaN=20from=20gradients?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 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) --- .../ml/src/cuda_pipeline/c51_grad_kernel.cu | 16 +- .../ml/src/cuda_pipeline/c51_loss_kernel.cu | 3 +- .../ml/src/cuda_pipeline/cql_grad_kernel.cu | 14 +- .../src/cuda_pipeline/dqn_utility_kernels.cu | 20 ++ .../ml/src/cuda_pipeline/gpu_dqn_trainer.rs | 176 +++++++++++++----- .../ml/src/cuda_pipeline/mse_grad_kernel.cu | 16 +- .../ml/src/cuda_pipeline/mse_loss_kernel.cu | 5 +- 7 files changed, 174 insertions(+), 76 deletions(-) diff --git a/crates/ml/src/cuda_pipeline/c51_grad_kernel.cu b/crates/ml/src/cuda_pipeline/c51_grad_kernel.cu index 876dd0cca..99b31f405 100644 --- a/crates/ml/src/cuda_pipeline/c51_grad_kernel.cu +++ b/crates/ml/src/cuda_pipeline/c51_grad_kernel.cu @@ -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); } } diff --git a/crates/ml/src/cuda_pipeline/c51_loss_kernel.cu b/crates/ml/src/cuda_pipeline/c51_loss_kernel.cu index f5c3199ff..69f09fcbc 100644 --- a/crates/ml/src/cuda_pipeline/c51_loss_kernel.cu +++ b/crates/ml/src/cuda_pipeline/c51_loss_kernel.cu @@ -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); diff --git a/crates/ml/src/cuda_pipeline/cql_grad_kernel.cu b/crates/ml/src/cuda_pipeline/cql_grad_kernel.cu index 41a42fe0b..a800a06e2 100644 --- a/crates/ml/src/cuda_pipeline/cql_grad_kernel.cu +++ b/crates/ml/src/cuda_pipeline/cql_grad_kernel.cu @@ -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]; } } diff --git a/crates/ml/src/cuda_pipeline/dqn_utility_kernels.cu b/crates/ml/src/cuda_pipeline/dqn_utility_kernels.cu index d3591a6d4..35d44cde6 100644 --- a/crates/ml/src/cuda_pipeline/dqn_utility_kernels.cu +++ b/crates/ml/src/cuda_pipeline/dqn_utility_kernels.cu @@ -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 * diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index 6588eb03d..91400400c 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -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, /// Scratch buffers for blended loss (MSE grad stored here, then blended into d_value/d_adv) - d_value_logits_mse: CudaSlice, - d_adv_logits_mse: CudaSlice, - /// Gradient w.r.t. value logits: [B, NA] - d_value_logits_buf: CudaSlice, - /// Gradient w.r.t. branch logits: [B, (B0+B1+B2)*NA] - d_adv_logits_buf: CudaSlice, + d_value_logits_mse: CudaSlice, + d_adv_logits_mse: CudaSlice, + /// Gradient w.r.t. value logits: [B, NA] — f32 for native atomicAdd (no bf16 overflow) + d_value_logits_buf: CudaSlice, + /// Gradient w.r.t. branch logits: [B, (B0+B1+B2)*NA] — f32 for native atomicAdd + d_adv_logits_buf: CudaSlice, + /// BF16 staging for backward pass: cast from f32 d_value_logits before cuBLAS GEMM + d_value_logits_bf16: CudaSlice, + /// BF16 staging for backward pass: cast from f32 d_adv_logits before cuBLAS GEMM + d_adv_logits_bf16: CudaSlice, // ── 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, - /// CQL scratch: value logit gradients [B, NA] - cql_d_value_logits: CudaSlice, - /// CQL scratch: advantage logit gradients [B, (B0+B1+B2)*NA] - cql_d_adv_logits: CudaSlice, + /// CQL scratch: value logit gradients [B, NA] — f32 for native atomicAdd + cql_d_value_logits: CudaSlice, + /// CQL scratch: advantage logit gradients [B, (B0+B1+B2)*NA] — f32 + cql_d_adv_logits: CudaSlice, } 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, ¶m_sizes); - // Construct d_adv_logits pointers per branch - let f32_size = std::mem::size_of::(); - 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::(); + 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::(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::(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::(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::(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::(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::(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, - d_adv_dst: &CudaSlice, + d_value_dst: &CudaSlice, + d_adv_dst: &CudaSlice, ) -> 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::() as u64; + let bf16_size = std::mem::size_of::() 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, 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. diff --git a/crates/ml/src/cuda_pipeline/mse_grad_kernel.cu b/crates/ml/src/cuda_pipeline/mse_grad_kernel.cu index 0b86f13aa..750b01810 100644 --- a/crates/ml/src/cuda_pipeline/mse_grad_kernel.cu +++ b/crates/ml/src/cuda_pipeline/mse_grad_kernel.cu @@ -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); } } diff --git a/crates/ml/src/cuda_pipeline/mse_loss_kernel.cu b/crates/ml/src/cuda_pipeline/mse_loss_kernel.cu index f15fcbe5a..ef6479350 100644 --- a/crates/ml/src/cuda_pipeline/mse_loss_kernel.cu +++ b/crates/ml/src/cuda_pipeline/mse_loss_kernel.cu @@ -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);