From 5328f0e33b4be497a7fd56f85498d7c3cdd9d8ea Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sat, 28 Mar 2026 22:21:42 +0100 Subject: [PATCH] fix(bf16): f32 total_loss_buf + training guard raw ptr interface MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - total_loss_buf: CudaSlice → CudaSlice (native atomicAdd, eliminates atomicAddBF16 CAS loop as potential NaN source) - Loss kernels: float* total_loss + atomicAdd (was atomicAddBF16) - Training guard: const float* loss_scalar (reads f32 directly) - Guard check_and_accumulate: takes u64 raw ptrs (type-agnostic) - All callers pass .raw_ptr() — works for both fused (f32) and non-fused (bf16) paths - Readback: reads 4 bytes f32 for loss (was 2 bytes bf16) NaN persists: the per-sample loss computation in the loss kernel produces NaN for specific samples despite float arithmetic and ±500 activation clamping. The NaN is within the softmax/expected-Q/TD-error chain, not from the accumulator. Next step: add in-kernel NaN detection to pinpoint the exact computation step. Co-Authored-By: Claude Opus 4.6 (1M context) --- .../ml/src/cuda_pipeline/c51_loss_kernel.cu | 4 +-- .../ml/src/cuda_pipeline/gpu_dqn_trainer.rs | 31 ++++++++-------- .../src/cuda_pipeline/gpu_training_guard.rs | 11 +++--- .../ml/src/cuda_pipeline/mse_loss_kernel.cu | 4 +-- .../cuda_pipeline/training_guard_kernel.cu | 7 ++-- crates/ml/src/trainers/dqn/fused_training.rs | 4 +-- .../ml/src/trainers/dqn/trainer/train_step.rs | 10 +++--- .../src/trainers/dqn/trainer/training_loop.rs | 35 ++++++++----------- 8 files changed, 49 insertions(+), 57 deletions(-) diff --git a/crates/ml/src/cuda_pipeline/c51_loss_kernel.cu b/crates/ml/src/cuda_pipeline/c51_loss_kernel.cu index e1ba12620..311fc5204 100644 --- a/crates/ml/src/cuda_pipeline/c51_loss_kernel.cu +++ b/crates/ml/src/cuda_pipeline/c51_loss_kernel.cu @@ -183,7 +183,7 @@ extern "C" __global__ void c51_loss_batched( __nv_bfloat16* __restrict__ per_sample_loss, __nv_bfloat16* __restrict__ td_errors, - __nv_bfloat16* __restrict__ total_loss, + float* __restrict__ total_loss, /* [1] float accumulator (native atomicAdd) */ __nv_bfloat16* __restrict__ save_current_lp, __nv_bfloat16* __restrict__ save_projected, @@ -413,6 +413,6 @@ extern "C" __global__ void c51_loss_batched( float weighted_loss = clamped_ce * is_weight; per_sample_loss[sample_id] = bf16(weighted_loss); td_errors[sample_id] = bf16(clamped_ce); - atomicAddBF16(total_loss, bf16(weighted_loss / (float)batch_size)); + atomicAdd(total_loss, weighted_loss / (float)batch_size); } } diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index db0b0b083..c732869f7 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -448,7 +448,7 @@ pub struct GpuDqnTrainer { // ── Forward output buffers ────────────────────────────────────── per_sample_loss_buf: CudaSlice, // [B] td_errors_buf: CudaSlice, // [B] - pub(crate) total_loss_buf: CudaSlice, // [1] + pub(crate) total_loss_buf: CudaSlice, // [1] float accumulator (native atomicAdd) // ── Forward-only Q-value output ───────────────────────────────── q_out_buf: CudaSlice, // [B, TOTAL_ACTIONS(11)] @@ -1774,7 +1774,8 @@ impl GpuDqnTrainer { // ── Allocate forward output buffers ───────────────────────── let per_sample_loss_buf = alloc_bf16(&stream, b, "per_sample_loss")?; let td_errors_buf = alloc_bf16(&stream, b, "td_errors")?; - let total_loss_buf = alloc_bf16(&stream, 1, "total_loss")?; + let total_loss_buf = stream.alloc_zeros::(1) + .map_err(|e| MLError::ModelError(format!("alloc total_loss_f32: {e}")))?; let total_actions = config.branch_0_size + config.branch_1_size + config.branch_2_size; let q_out_buf = alloc_bf16(&stream, b * total_actions, "q_out")?; @@ -2352,22 +2353,21 @@ impl GpuDqnTrainer { unsafe { cudarc::driver::sys::cuStreamSynchronize(self.stream.cu_stream()); } - let bf16_size = std::mem::size_of::(); - let mut loss_bf16 = [half::bf16::ZERO; 1]; + let mut loss_f32 = [0.0_f32; 1]; let mut norm_bf16 = [half::bf16::ZERO; 1]; unsafe { cudarc::driver::sys::cuMemcpyDtoH_v2( - loss_bf16.as_mut_ptr().cast(), - self.ptrs.total_loss_buf, bf16_size, + loss_f32.as_mut_ptr().cast(), + self.ptrs.total_loss_buf, std::mem::size_of::(), ); cudarc::driver::sys::cuMemcpyDtoH_v2( norm_bf16.as_mut_ptr().cast(), - self.ptrs.grad_norm_buf, bf16_size, + self.ptrs.grad_norm_buf, std::mem::size_of::(), ); } Ok(FusedTrainScalars { - total_loss: loss_bf16[0].to_f32(), - grad_norm: norm_bf16[0].to_f32(), // bf16 L2 norm from finalize kernel + total_loss: loss_f32[0], + grad_norm: norm_bf16[0].to_f32(), }) } @@ -2487,20 +2487,19 @@ impl GpuDqnTrainer { unsafe { cudarc::driver::sys::cuStreamSynchronize(self.stream.cu_stream()); } - let bf16_size = std::mem::size_of::(); - let mut loss_bf16 = [half::bf16::ZERO; 1]; + let mut loss_f32 = [0.0_f32; 1]; let mut norm_bf16 = [half::bf16::ZERO; 1]; unsafe { cudarc::driver::sys::cuMemcpyDtoH_v2( - loss_bf16.as_mut_ptr().cast(), - self.total_loss_buf.raw_ptr(), bf16_size, + loss_f32.as_mut_ptr().cast(), + self.total_loss_buf.raw_ptr(), std::mem::size_of::(), ); cudarc::driver::sys::cuMemcpyDtoH_v2( norm_bf16.as_mut_ptr().cast(), - self.grad_norm_buf.raw_ptr(), bf16_size, + self.grad_norm_buf.raw_ptr(), std::mem::size_of::(), ); - } // gpu-exit: 2 scalar readbacks (4 bytes total, bf16) - self.scalars_readback_host = [loss_bf16[0].to_f32(), norm_bf16[0].to_f32()]; + } + self.scalars_readback_host = [loss_f32[0], norm_bf16[0].to_f32()]; Ok(FusedTrainScalars { total_loss: self.scalars_readback_host[0], diff --git a/crates/ml/src/cuda_pipeline/gpu_training_guard.rs b/crates/ml/src/cuda_pipeline/gpu_training_guard.rs index d61472b10..7fbbd111e 100644 --- a/crates/ml/src/cuda_pipeline/gpu_training_guard.rs +++ b/crates/ml/src/cuda_pipeline/gpu_training_guard.rs @@ -267,10 +267,12 @@ impl GpuTrainingGuard { /// Returns the safety flags and scalar values from the *previous* step /// (one-step delay due to double-buffering). On the very first call, /// returns safe defaults (no halts, zero loss/grad_norm). + /// loss_ptr: device pointer to a SINGLE f32 scalar (total_loss_buf) + /// grad_norm_ptr: device pointer to a SINGLE bf16 scalar (grad_norm_buf) pub fn check_and_accumulate( &mut self, - loss_gpu: &CudaSlice, - grad_norm_gpu: &CudaSlice, + loss_ptr: u64, + grad_norm_ptr: u64, clip_threshold: f32, collapse_threshold: f32, warmup: bool, @@ -316,14 +318,13 @@ impl GpuTrainingGuard { }; // Pass raw device pointers to bypass cudarc event tracking on graph-captured buffers. - let loss_ptr = loss_gpu.raw_ptr(); - let grad_ptr = grad_norm_gpu.raw_ptr(); + // loss_ptr and grad_norm_ptr are already raw u64 device addresses. let acc_ptr = self.acc_buf.raw_ptr(); unsafe { self.stream .launch_builder(&self.fused_check_accum_func) .arg(&loss_ptr) - .arg(&grad_ptr) + .arg(&grad_norm_ptr) .arg(&write_dev_ptr) .arg(&acc_ptr) .arg(&clip_threshold) diff --git a/crates/ml/src/cuda_pipeline/mse_loss_kernel.cu b/crates/ml/src/cuda_pipeline/mse_loss_kernel.cu index efbda9f15..d6cc3191c 100644 --- a/crates/ml/src/cuda_pipeline/mse_loss_kernel.cu +++ b/crates/ml/src/cuda_pipeline/mse_loss_kernel.cu @@ -133,7 +133,7 @@ extern "C" __global__ void mse_loss_batched( /* ── Outputs ──────────────────────────────────────────────────── */ __nv_bfloat16* __restrict__ per_sample_loss, /* [B] IS-weighted loss per sample */ __nv_bfloat16* __restrict__ td_errors, /* [B] unweighted, for PER priority update */ - __nv_bfloat16* __restrict__ total_loss, /* [1] batch mean loss (BF16 atomicAddBF16) */ + float* __restrict__ total_loss, /* [1] float accumulator (native atomicAdd) */ /* ── Saved tensors for backward pass ─────────────────────────── */ __nv_bfloat16* __restrict__ save_current_lp, /* [B, NUM_BRANCHES, num_atoms] online probs */ @@ -357,6 +357,6 @@ extern "C" __global__ void mse_loss_batched( float weighted_loss = avg_mse * is_weight; per_sample_loss[sample_id] = bf16(weighted_loss); td_errors[sample_id] = bf16(avg_td); - atomicAddBF16(total_loss, bf16(weighted_loss / (float)batch_size)); + atomicAdd(total_loss, weighted_loss / (float)batch_size); } } diff --git a/crates/ml/src/cuda_pipeline/training_guard_kernel.cu b/crates/ml/src/cuda_pipeline/training_guard_kernel.cu index 42c75cca2..378f59748 100644 --- a/crates/ml/src/cuda_pipeline/training_guard_kernel.cu +++ b/crates/ml/src/cuda_pipeline/training_guard_kernel.cu @@ -38,16 +38,15 @@ /* [2] step_count */ /* ------------------------------------------------------------------ */ extern "C" __global__ void training_guard_check_and_accumulate( - const __nv_bfloat16* __restrict__ loss_scalar, /* GPU-resident scalar */ - const __nv_bfloat16* __restrict__ grad_norm_scalar, /* GPU-resident scalar */ + const float* __restrict__ loss_scalar, /* GPU-resident f32 scalar */ + const __nv_bfloat16* __restrict__ grad_norm_scalar, /* GPU-resident bf16 scalar */ float* output, /* pinned host buffer (7 floats) */ __nv_bfloat16* acc_buf, /* device accumulator (3 bf16) */ float clip_threshold, float collapse_threshold, int warmup ) { - /* Read BF16 scalars, cast to F32 for NaN/Inf detection (no isnan on bf16) */ - float loss = (float)loss_scalar[0]; + float loss = *loss_scalar; /* native f32 read */ float grad_norm = (float)grad_norm_scalar[0]; /* -- Guard check -- */ diff --git a/crates/ml/src/trainers/dqn/fused_training.rs b/crates/ml/src/trainers/dqn/fused_training.rs index 538ffabf7..699bd9859 100644 --- a/crates/ml/src/trainers/dqn/fused_training.rs +++ b/crates/ml/src/trainers/dqn/fused_training.rs @@ -1165,9 +1165,9 @@ impl FusedTrainingCtx { self.trainer.set_c51_alpha(alpha); } - /// Return a reference to the GPU-resident total_loss scalar (bf16). + /// Return a reference to the GPU-resident total_loss scalar (f32). /// Written by the CUDA graph's loss kernel — valid after `replay_forward()`. - pub(crate) fn loss_gpu_buf(&self) -> &cudarc::driver::CudaSlice { + pub(crate) fn loss_gpu_buf(&self) -> &cudarc::driver::CudaSlice { &self.trainer.total_loss_buf } diff --git a/crates/ml/src/trainers/dqn/trainer/train_step.rs b/crates/ml/src/trainers/dqn/trainer/train_step.rs index aff5f4bc7..47122e6a4 100644 --- a/crates/ml/src/trainers/dqn/trainer/train_step.rs +++ b/crates/ml/src/trainers/dqn/trainer/train_step.rs @@ -125,12 +125,10 @@ impl DQNTrainer { let fused = self.fused_ctx.as_mut() .ok_or_else(|| anyhow::anyhow!("fused_ctx required for training guard"))?; - let loss_buf = fused.loss_gpu_buf(); - let grad_buf = fused.grad_norm_gpu_buf(); let result = guard .check_and_accumulate( - loss_buf, - grad_buf, + fused.loss_gpu_buf().raw_ptr(), + fused.grad_norm_gpu_buf().raw_ptr(), 1e6_f32, // loss clip threshold grad_collapse_threshold, !past_warmup, @@ -342,8 +340,8 @@ impl DQNTrainer { let grad_slice = grad_sl_r .map_err(|e| anyhow::anyhow!("GPU guard accum grad CudaSlice: {e}"))?; let guard_result = guard.check_and_accumulate( - &loss_slice, - &grad_slice, + loss_slice.raw_ptr(), + grad_slice.raw_ptr(), 1e6_f32, grad_collapse_threshold, !past_warmup, diff --git a/crates/ml/src/trainers/dqn/trainer/training_loop.rs b/crates/ml/src/trainers/dqn/trainer/training_loop.rs index 164eff302..def91ede8 100644 --- a/crates/ml/src/trainers/dqn/trainer/training_loop.rs +++ b/crates/ml/src/trainers/dqn/trainer/training_loop.rs @@ -1265,27 +1265,22 @@ impl DQNTrainer { if let Some(ref mut guard) = self.training_guard { // Read loss/grad directly from fused trainer's GPU buffers. // GpuTrainResult returns hardcoded zeros per-step (no sync). - let gr = if let Some(ref fused) = self.fused_ctx { - guard.check_and_accumulate( - fused.loss_gpu_buf(), - fused.grad_norm_gpu_buf(), - 1e6_f32, - guard_collapse_thresh, - !guard_past_warmup, - ).map_err(|e| anyhow::anyhow!("guard check: {e}"))? + let (loss_raw, grad_raw) = if let Some(ref fused) = self.fused_ctx { + (fused.loss_gpu_buf().raw_ptr(), fused.grad_norm_gpu_buf().raw_ptr()) } else { - let loss_slice = _gpu_result.loss_cuda_slice() - .map_err(|e| anyhow::anyhow!("guard loss CudaSlice: {e}"))?; - let grad_slice = _gpu_result.grad_norm_cuda_slice() - .map_err(|e| anyhow::anyhow!("guard grad CudaSlice: {e}"))?; - guard.check_and_accumulate( - &loss_slice, - &grad_slice, - 1e6_f32, - guard_collapse_thresh, - !guard_past_warmup, - ).map_err(|e| anyhow::anyhow!("guard check: {e}"))? + let ls = _gpu_result.loss_cuda_slice() + .map_err(|e| anyhow::anyhow!("guard loss: {e}"))?; + let gs = _gpu_result.grad_norm_cuda_slice() + .map_err(|e| anyhow::anyhow!("guard grad: {e}"))?; + (ls.raw_ptr(), gs.raw_ptr()) }; + let gr = guard.check_and_accumulate( + loss_raw, + grad_raw, + 1e6_f32, + guard_collapse_thresh, + !guard_past_warmup, + ).map_err(|e| anyhow::anyhow!("guard check: {e}"))?; if gr.halt_nan { return Err(anyhow::anyhow!( "NaN/Inf at step {}: loss={}, grad={}", @@ -1391,7 +1386,7 @@ impl DQNTrainer { let grad_sl = grad_sl_r .map_err(|e| anyhow::anyhow!("accum grad CudaSlice: {e}"))?; let gr = guard.check_and_accumulate( - &loss_sl, &grad_sl, 1e6_f32, + loss_sl.raw_ptr(), grad_sl.raw_ptr(), 1e6_f32, guard_collapse_thresh, !guard_past_warmup, ).map_err(|e| anyhow::anyhow!("guard accum step: {e}"))?; if gr.halt_nan {