fix(bf16): f32 total_loss_buf + training guard raw ptr interface

- total_loss_buf: CudaSlice<half::bf16> → CudaSlice<f32> (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) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-03-28 22:21:42 +01:00
parent 0af691cae5
commit 5328f0e33b
8 changed files with 49 additions and 57 deletions

View File

@@ -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);
}
}

View File

@@ -448,7 +448,7 @@ pub struct GpuDqnTrainer {
// ── Forward output buffers ──────────────────────────────────────
per_sample_loss_buf: CudaSlice<half::bf16>, // [B]
td_errors_buf: CudaSlice<half::bf16>, // [B]
pub(crate) total_loss_buf: CudaSlice<half::bf16>, // [1]
pub(crate) total_loss_buf: CudaSlice<f32>, // [1] float accumulator (native atomicAdd)
// ── Forward-only Q-value output ─────────────────────────────────
q_out_buf: CudaSlice<half::bf16>, // [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::<f32>(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::<half::bf16>();
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::<f32>(),
);
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::<half::bf16>(),
);
}
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::<half::bf16>();
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::<f32>(),
);
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::<half::bf16>(),
);
} // 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],

View File

@@ -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<half::bf16>,
grad_norm_gpu: &CudaSlice<half::bf16>,
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)

View File

@@ -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);
}
}

View File

@@ -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 -- */

View File

@@ -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<half::bf16> {
pub(crate) fn loss_gpu_buf(&self) -> &cudarc::driver::CudaSlice<f32> {
&self.trainer.total_loss_buf
}

View File

@@ -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,

View File

@@ -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 {