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:
@@ -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);
|
||||
|
||||
Reference in New Issue
Block a user