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:
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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],
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 -- */
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user