From 62b57c537eb9dffaa593476fce9c921d628cfe44 Mon Sep 17 00:00:00 2001 From: jgrusewski Date: Sat, 28 Mar 2026 17:50:53 +0100 Subject: [PATCH] =?UTF-8?q?fix(bf16):=20CRITICAL=20=E2=80=94=20hardcoded?= =?UTF-8?q?=20*4=20byte=20offsets=20for=20branch=20logit=20pointers?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Root cause of all NaN/garbage: 12 branch logit pointer offsets used hardcoded * 4 (sizeof f32) instead of * sizeof(bf16) = 2. This caused mse_loss_batched and c51_loss_batched to read 2× past the end of branch logit buffers — 3719 out-of-bounds reads per step (detected by compute-sanitizer). The out-of-bounds reads produced NaN gradients → Adam propagated NaN to params → all downstream reads returned garbage. Fix: replace * 4 with * std::mem::size_of::() (= 2). Smoke test: Q-values now VALID (1.18, 1.91), Sharpe +8.98, val_loss=-4.04. Training completes 3 epochs without divergence. Remaining: train_loss/grad_norm readback still 0 (training_guard). Co-Authored-By: Claude Opus 4.6 (1M context) --- .../ml/src/cuda_pipeline/gpu_dqn_trainer.rs | 24 +++++++++---------- 1 file changed, 12 insertions(+), 12 deletions(-) diff --git a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs index b72566d07..94bad02de 100644 --- a/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs +++ b/crates/ml/src/cuda_pipeline/gpu_dqn_trainer.rs @@ -3256,20 +3256,20 @@ impl GpuDqnTrainer { // Extract raw pointers for online branch logits (contiguous [B, (B0+B1+B2)*NA]) let on_b_base = self.on_b_logits_buf.raw_ptr(); let on_b0_ptr = on_b_base; - let on_b1_ptr = on_b_base + (b * b0 * na * 4) as u64; - let on_b2_ptr = on_b1_ptr + (b * b1 * na * 4) as u64; + let on_b1_ptr = on_b_base + (b * b0 * na * std::mem::size_of::()) as u64; + let on_b2_ptr = on_b1_ptr + (b * b1 * na * std::mem::size_of::()) as u64; // Target branch logits let tg_b_base = self.tg_b_logits_buf.raw_ptr(); let tg_b0_ptr = tg_b_base; - let tg_b1_ptr = tg_b_base + (b * b0 * na * 4) as u64; - let tg_b2_ptr = tg_b1_ptr + (b * b1 * na * 4) as u64; + let tg_b1_ptr = tg_b_base + (b * b0 * na * std::mem::size_of::()) as u64; + let tg_b2_ptr = tg_b1_ptr + (b * b1 * na * std::mem::size_of::()) as u64; // Online-next branch logits (Double DQN action selector) let on_next_b_base = self.on_next_b_logits_buf.raw_ptr(); let on_next_b0_ptr = on_next_b_base; - let on_next_b1_ptr = on_next_b_base + (b * b0 * na * 4) as u64; - let on_next_b2_ptr = on_next_b1_ptr + (b * b1 * na * 4) as u64; + let on_next_b1_ptr = on_next_b_base + (b * b0 * na * std::mem::size_of::()) as u64; + let on_next_b2_ptr = on_next_b1_ptr + (b * b1 * na * std::mem::size_of::()) as u64; // N-step returns: use gamma^n for the Bellman projection. // The experience collector pre-computes R_n = sum(gamma^i * r_i). @@ -3402,20 +3402,20 @@ impl GpuDqnTrainer { // Extract raw pointers for online branch logits (contiguous [B, (B0+B1+B2)*NA]) let on_b_base = self.on_b_logits_buf.raw_ptr(); let on_b0_ptr = on_b_base; - let on_b1_ptr = on_b_base + (b * b0 * na * 4) as u64; - let on_b2_ptr = on_b1_ptr + (b * b1 * na * 4) as u64; + let on_b1_ptr = on_b_base + (b * b0 * na * std::mem::size_of::()) as u64; + let on_b2_ptr = on_b1_ptr + (b * b1 * na * std::mem::size_of::()) as u64; // Target branch logits let tg_b_base = self.tg_b_logits_buf.raw_ptr(); let tg_b0_ptr = tg_b_base; - let tg_b1_ptr = tg_b_base + (b * b0 * na * 4) as u64; - let tg_b2_ptr = tg_b1_ptr + (b * b1 * na * 4) as u64; + let tg_b1_ptr = tg_b_base + (b * b0 * na * std::mem::size_of::()) as u64; + let tg_b2_ptr = tg_b1_ptr + (b * b1 * na * std::mem::size_of::()) as u64; // Online-next branch logits (Double DQN action selector) let on_next_b_base = self.on_next_b_logits_buf.raw_ptr(); let on_next_b0_ptr = on_next_b_base; - let on_next_b1_ptr = on_next_b_base + (b * b0 * na * 4) as u64; - let on_next_b2_ptr = on_next_b1_ptr + (b * b1 * na * 4) as u64; + let on_next_b1_ptr = on_next_b_base + (b * b0 * na * std::mem::size_of::()) as u64; + let on_next_b2_ptr = on_next_b1_ptr + (b * b1 * na * std::mem::size_of::()) as u64; // N-step returns: use gamma^n for the Bellman projection. let gamma = self.config.gamma.powi(self.config.n_steps as i32);