fix(bf16): CRITICAL — hardcoded *4 byte offsets for branch logit pointers

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::<half::bf16>() (= 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) <noreply@anthropic.com>
This commit is contained in:
jgrusewski
2026-03-28 17:50:53 +01:00
parent ff4ab9f9a0
commit 62b57c537e

View File

@@ -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::<half::bf16>()) as u64;
let on_b2_ptr = on_b1_ptr + (b * b1 * na * std::mem::size_of::<half::bf16>()) 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::<half::bf16>()) as u64;
let tg_b2_ptr = tg_b1_ptr + (b * b1 * na * std::mem::size_of::<half::bf16>()) 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::<half::bf16>()) as u64;
let on_next_b2_ptr = on_next_b1_ptr + (b * b1 * na * std::mem::size_of::<half::bf16>()) 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::<half::bf16>()) as u64;
let on_b2_ptr = on_b1_ptr + (b * b1 * na * std::mem::size_of::<half::bf16>()) 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::<half::bf16>()) as u64;
let tg_b2_ptr = tg_b1_ptr + (b * b1 * na * std::mem::size_of::<half::bf16>()) 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::<half::bf16>()) as u64;
let on_next_b2_ptr = on_next_b1_ptr + (b * b1 * na * std::mem::size_of::<half::bf16>()) as u64;
// N-step returns: use gamma^n for the Bellman projection.
let gamma = self.config.gamma.powi(self.config.n_steps as i32);