Files
foxhunt/crates/ml-alpha/tests/frd_head.rs
jgrusewski 0f75d6bb7b feat(rl): FRD layer-1 backward (dW1, db1, dh_t with ReLU mask) — F.3c
Third and final FRD backward stage. Closes the chain from
softmax+CE loss back to the encoder's hidden state h_t.

Kernel `cuda/rl_frd_layer1_bwd.cu`:
  * grid_dim = (B, 1, 1), block_dim = (HIDDEN_DIM=128, 1, 1)
  * Phase 0: threads 0..63 stage dL/dpre_hidden = grad_hidden ×
    1{hidden > 0} into shared mem (the cached post-ReLU `hidden`
    buffer encodes the mask — hidden == 0 ⇔ pre-activation was
    ≤ 0 → ReLU killed it). Same thread also writes db1_per_batch.
  * Phase 1: each thread k (k < 128) writes one row of
    grad_W1_per_batch[b, k, 0..64] (64 writes per thread, no atomics)
  * Phase 2: same thread computes grad_h_t[b, k] =
    Σ_i W1[k, i] × dL/dpre_hidden[b, i]
  * Per-(b, k, i) sole-writer per feedback_no_atomicadd

Rust wiring `FrdHead::layer1_bwd` — takes h_t, hidden (forward cache),
grad_hidden (from layer2_bwd), self.w1_d; writes grad_w1_per_batch,
grad_b1_per_batch, grad_h_t. The grad_h_t buffer becomes the encoder-
upstream gradient that the trainer's grad_h_accumulate kernel folds
into the encoder's gradient with λ_frd scaling (same pattern as Q/π/V
heads — wiring lives in F.4).

Tests (2 new, 10/10 file total):
  * frd_layer1_bwd_finite_diff_w1 — perturbs the W1 slot with MAX
    |analytical gradient| (instead of an arbitrary fixed slot — fp32
    finite-diff is rounding-error-limited so a tiny gradient gives
    misleading rel_err). At max-magnitude slot (k=84, i=55): analytical
    = -0.0451, numerical = -0.0448, rel_err = 5.6e-3 — well within
    1e-2 tolerance (slightly looser than dW2's 5e-3 because dW1
    crosses an extra matmul + the ReLU mask boundary).
  * frd_layer1_bwd_relu_mask_zeros_grad — fixture with h_t = all -1
    produces ~half the hidden slots ReLU-masked (cached hidden = 0).
    For every masked slot i, asserts:
      * db1_per_batch[b, i] == 0 (exact equality — mask is hard 0)
      * dW1_per_batch[b, k, i] == 0 for every k (~32 × 128 = 4096
        slots checked)
    Empirically 32/64 masked, 32/64 active — confirms ReLU mask
    is wired through the chain correctly without leaking gradient
    through dead branches.

F.3 backward chain is now complete end-to-end:
  rl_frd_softmax_ce_grad (F.3a) → rl_frd_layer2_bwd (F.3b) →
  rl_frd_layer1_bwd (F.3c) → grad_h_t (consumed by F.4 wiring)

F.4 wires Adam optimizers for W1/b1/W2/b2 + grad_h_accumulate into
the encoder gradient + loader-side label generation + λ_frd × CE
into stats.l_total.
2026-05-24 18:40:30 +02:00

738 lines
30 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! GPU-oracle tests for the Forward-Return-Distribution head (SP20 P3).
//!
//! Forward-pass invariants verified:
//! 1. `frd_forward_zero_input_emits_zero_logits` — when `h_t = 0` and
//! both biases are zero (the default `FrdHead::new` init), the
//! output logits must be exactly zero. Provides an unambiguous
//! analytical oracle for the matmul + ReLU + matmul chain.
//! 2. `frd_forward_shape_matches_spec` — output buffer length is
//! `b_size × FRD_OUT_DIM` and the per-horizon softmax sums to 1
//! within fp tolerance. Exercises a random `h_t` to verify the
//! head produces valid distributions for any input.
//! 3. `frd_forward_caches_hidden_for_bwd` — after a forward pass, the
//! cached hidden activation buffer matches a re-computed ReLU
//! sentinel: when `h_t` is set such that `W1 @ h_t + b1` is
//! strictly negative for some output (driven via biases), those
//! slots stay at exactly zero.
//!
//! Per `feedback_no_cpu_test_fallbacks`: oracles are analytical
//! (zero-input → zero-output; softmax sum invariant), no CPU reference
//! impl.
//! Per `feedback_no_htod_htoh_only_mapped_pinned`: all host↔device
//! transfers use mapped-pinned helpers.
//!
//! Run with:
//! `cargo test -p ml-alpha --test frd_head -- --ignored --nocapture`
use anyhow::Result;
use cudarc::driver::{CudaSlice, CudaStream};
use ml_alpha::heads::HIDDEN_DIM;
use ml_alpha::rl::common::{FRD_HIDDEN_DIM, FRD_N_ATOMS, FRD_N_HORIZONS};
use ml_alpha::rl::frd::{FrdHead, FrdHeadConfig, FRD_OUT_DIM};
use ml_alpha::trainer::integrated::{
read_slice_d_pub, write_slice_f32_d_pub, write_slice_i32_d_pub,
};
use ml_core::device::MlDevice;
use std::sync::Arc;
fn build_head() -> Option<(MlDevice, FrdHead)> {
let dev = match MlDevice::cuda(0) {
Ok(d) => d,
Err(e) => {
eprintln!("CUDA 0 not available — skipping ({e})");
return None;
}
};
let head = FrdHead::new(&dev, FrdHeadConfig { seed: 0xF8D_42 }).expect("FrdHead::new");
Some((dev, head))
}
fn upload_f32(stream: &Arc<CudaStream>, host: &[f32]) -> Result<CudaSlice<f32>> {
let mut d = stream.alloc_zeros::<f32>(host.len())?;
write_slice_f32_d_pub(stream, host, &mut d)?;
Ok(d)
}
#[test]
#[ignore = "requires CUDA (MlDevice::cuda(0))"]
fn frd_forward_zero_input_emits_zero_logits() -> Result<()> {
let Some((dev, head)) = build_head() else { return Ok(()) };
let stream = dev.cuda_stream()?.clone();
let b_size = 4;
// h_t = 0, b1 = 0 (default init), so hidden_pre = 0 → ReLU(0) = 0.
// Then hidden @ W2 + b2 = 0 + 0 = 0 regardless of W2's values.
let h_t = vec![0.0_f32; b_size * HIDDEN_DIM];
let h_t_d = upload_f32(&stream, &h_t)?;
let mut hidden_d = stream.alloc_zeros::<f32>(b_size * FRD_HIDDEN_DIM)?;
let mut logits_d = stream.alloc_zeros::<f32>(b_size * FRD_OUT_DIM)?;
head.forward(&h_t_d, &mut hidden_d, &mut logits_d, b_size)?;
let logits = read_slice_d_pub(&stream, &logits_d, b_size * FRD_OUT_DIM)?;
let hidden = read_slice_d_pub(&stream, &hidden_d, b_size * FRD_HIDDEN_DIM)?;
for (i, v) in logits.iter().enumerate() {
assert_eq!(
*v, 0.0,
"logit[{i}] should be exactly 0 for h_t=0 with default b1=b2=0; got {v}"
);
}
for (i, v) in hidden.iter().enumerate() {
assert_eq!(
*v, 0.0,
"hidden[{i}] should be exactly 0 for h_t=0 with default b1=0 (ReLU(0)=0); got {v}"
);
}
eprintln!(
"PASS — zero-input → zero-logits invariant ({} logits, {} hidden slots)",
logits.len(),
hidden.len()
);
Ok(())
}
#[test]
#[ignore = "requires CUDA (MlDevice::cuda(0))"]
fn frd_forward_shape_matches_spec() -> Result<()> {
let Some((dev, head)) = build_head() else { return Ok(()) };
let stream = dev.cuda_stream()?.clone();
let b_size = 8;
// Random h_t in [-1, 1] — Xavier-scaled W1 keeps pre-activations
// O(0.1), but enough nonzero that softmax distributions are
// meaningfully non-uniform.
use rand::{Rng, SeedableRng};
use rand_chacha::ChaCha8Rng;
let mut r = ChaCha8Rng::seed_from_u64(0xCAFE);
let h_t: Vec<f32> = (0..b_size * HIDDEN_DIM)
.map(|_| r.gen_range(-1.0..1.0))
.collect();
let h_t_d = upload_f32(&stream, &h_t)?;
let mut hidden_d = stream.alloc_zeros::<f32>(b_size * FRD_HIDDEN_DIM)?;
let mut logits_d = stream.alloc_zeros::<f32>(b_size * FRD_OUT_DIM)?;
head.forward(&h_t_d, &mut hidden_d, &mut logits_d, b_size)?;
let logits = read_slice_d_pub(&stream, &logits_d, b_size * FRD_OUT_DIM)?;
// Per-horizon softmax sums to 1 (standard softmax invariant).
for b in 0..b_size {
for h in 0..FRD_N_HORIZONS {
let off = b * FRD_OUT_DIM + h * FRD_N_ATOMS;
let row = &logits[off..off + FRD_N_ATOMS];
// log-sum-exp for numerical stability.
let max_l = row.iter().cloned().fold(f32::NEG_INFINITY, f32::max);
let denom: f32 = row.iter().map(|x| (x - max_l).exp()).sum();
let sum_p: f32 = row.iter().map(|x| (x - max_l).exp() / denom).sum();
assert!(
(sum_p - 1.0).abs() < 1e-5,
"softmax sum for batch {b} horizon {h} should be 1.0; got {sum_p}"
);
}
}
eprintln!(
"PASS — output shape {} (= {} × {}); per-horizon softmax sums to 1.0",
logits.len(),
b_size,
FRD_OUT_DIM
);
Ok(())
}
#[test]
#[ignore = "requires CUDA (MlDevice::cuda(0))"]
fn frd_forward_relu_mask_consistent_with_cached_hidden() -> Result<()> {
let Some((dev, head)) = build_head() else { return Ok(()) };
let stream = dev.cuda_stream()?.clone();
let b_size = 4;
// Use a uniformly negative input (-1) so SOME hidden pre-activations
// are strictly negative (depending on the random W1 sign pattern).
// The cached hidden buffer must contain the post-ReLU values:
// every value MUST be ≥ 0 (ReLU output is non-negative).
let h_t = vec![-1.0_f32; b_size * HIDDEN_DIM];
let h_t_d = upload_f32(&stream, &h_t)?;
let mut hidden_d = stream.alloc_zeros::<f32>(b_size * FRD_HIDDEN_DIM)?;
let mut logits_d = stream.alloc_zeros::<f32>(b_size * FRD_OUT_DIM)?;
head.forward(&h_t_d, &mut hidden_d, &mut logits_d, b_size)?;
let hidden = read_slice_d_pub(&stream, &hidden_d, b_size * FRD_HIDDEN_DIM)?;
let mut zero_count = 0;
for (i, v) in hidden.iter().enumerate() {
assert!(
*v >= 0.0,
"cached hidden[{i}] = {v} violates ReLU non-negativity invariant"
);
if *v == 0.0 {
zero_count += 1;
}
}
// With Xavier-symmetric init around 0, ≈half of the pre-activations
// should be negative → zero in the cached buffer. Assert at least
// SOME entries got masked (a sign-test that the ReLU actually fires);
// we don't lock in an exact ratio since it depends on the RNG.
assert!(
zero_count > 0,
"expected at least one ReLU-masked hidden slot under negative input; got {zero_count} zeros out of {}",
hidden.len()
);
eprintln!(
"PASS — cached hidden is non-negative; {}/{} slots ReLU-masked (negative input)",
zero_count,
hidden.len()
);
Ok(())
}
fn upload_i32(stream: &Arc<CudaStream>, host: &[i32]) -> Result<CudaSlice<i32>> {
let mut d = stream.alloc_zeros::<i32>(host.len())?;
write_slice_i32_d_pub(stream, host, &mut d)?;
Ok(d)
}
#[test]
#[ignore = "requires CUDA (MlDevice::cuda(0))"]
fn frd_softmax_ce_grad_uniform_logits_match_log_n_atoms() -> Result<()> {
let Some((dev, head)) = build_head() else { return Ok(()) };
let stream = dev.cuda_stream()?.clone();
let b_size = 4;
// Uniform logits (all zeros) → softmax = 1/FRD_N_ATOMS everywhere,
// so CE = -log(1/21) = ln(21) ≈ 3.0445 for ANY label index.
let logits = vec![0.0_f32; b_size * FRD_OUT_DIM];
let logits_d = upload_f32(&stream, &logits)?;
// Set label = 10 (mid-bucket) for every (batch, horizon) row.
let labels: Vec<i32> = vec![10; b_size * FRD_N_HORIZONS];
let labels_d = upload_i32(&stream, &labels)?;
let mut grad_d = stream.alloc_zeros::<f32>(b_size * FRD_OUT_DIM)?;
let mut loss_d = stream.alloc_zeros::<f32>(b_size * FRD_N_HORIZONS)?;
head.softmax_ce_grad(&logits_d, &labels_d, &mut grad_d, &mut loss_d, b_size)?;
let loss = read_slice_d_pub(&stream, &loss_d, b_size * FRD_N_HORIZONS)?;
let expected = (FRD_N_ATOMS as f32).ln();
for (i, v) in loss.iter().enumerate() {
assert!(
(v - expected).abs() < 1e-5,
"uniform-logit CE[{i}] should be ln({}) = {expected}; got {v}",
FRD_N_ATOMS
);
}
// Softmax-CE gradient invariant: Σ_a (p[a] - 1{a==label}) = 0
// (mean-reduced over B doesn't change the per-row sum being zero).
let grad = read_slice_d_pub(&stream, &grad_d, b_size * FRD_OUT_DIM)?;
for b in 0..b_size {
for h in 0..FRD_N_HORIZONS {
let off = b * FRD_OUT_DIM + h * FRD_N_ATOMS;
let row_sum: f32 = grad[off..off + FRD_N_ATOMS].iter().sum();
assert!(
row_sum.abs() < 1e-6,
"Σ grad_logits per (b={b}, h={h}) must be 0 (softmax-CE invariant); got {row_sum}"
);
}
}
eprintln!(
"PASS — uniform logits → CE = ln({}) = {:.4}; per-row grad sum = 0",
FRD_N_ATOMS, expected
);
Ok(())
}
#[test]
#[ignore = "requires CUDA (MlDevice::cuda(0))"]
fn frd_softmax_ce_grad_sentinel_label_zeros_row() -> Result<()> {
let Some((dev, head)) = build_head() else { return Ok(()) };
let stream = dev.cuda_stream()?.clone();
let b_size = 2;
// Non-trivial logits so a non-sentinel label would produce non-zero
// loss and gradient — proving the sentinel really is masking.
use rand::{Rng, SeedableRng};
use rand_chacha::ChaCha8Rng;
let mut r = ChaCha8Rng::seed_from_u64(0x5EE);
let logits: Vec<f32> = (0..b_size * FRD_OUT_DIM)
.map(|_| r.gen_range(-1.0..1.0))
.collect();
let logits_d = upload_f32(&stream, &logits)?;
// Labels: ALL sentinel (-1) → every row should be zeroed.
let labels: Vec<i32> = vec![-1; b_size * FRD_N_HORIZONS];
let labels_d = upload_i32(&stream, &labels)?;
let mut grad_d = stream.alloc_zeros::<f32>(b_size * FRD_OUT_DIM)?;
let mut loss_d = stream.alloc_zeros::<f32>(b_size * FRD_N_HORIZONS)?;
head.softmax_ce_grad(&logits_d, &labels_d, &mut grad_d, &mut loss_d, b_size)?;
let loss = read_slice_d_pub(&stream, &loss_d, b_size * FRD_N_HORIZONS)?;
let grad = read_slice_d_pub(&stream, &grad_d, b_size * FRD_OUT_DIM)?;
for (i, v) in loss.iter().enumerate() {
assert_eq!(*v, 0.0, "sentinel label loss[{i}] must be 0; got {v}");
}
for (i, v) in grad.iter().enumerate() {
assert_eq!(
*v, 0.0,
"sentinel label grad_logits[{i}] must be 0; got {v}"
);
}
eprintln!("PASS — sentinel label (-1) zeros loss + grad rows");
Ok(())
}
#[test]
#[ignore = "requires CUDA (MlDevice::cuda(0))"]
fn frd_softmax_ce_grad_finite_diff_matches_analytical() -> Result<()> {
let Some((dev, head)) = build_head() else { return Ok(()) };
let stream = dev.cuda_stream()?.clone();
let b_size = 1;
// Single batch, fixed (h=0, target_a=5). Pick a smooth random logit
// pattern so probabilities aren't degenerate at 0/1.
use rand::{Rng, SeedableRng};
use rand_chacha::ChaCha8Rng;
let mut r = ChaCha8Rng::seed_from_u64(0xFD1FF);
let mut logits: Vec<f32> = (0..b_size * FRD_OUT_DIM)
.map(|_| r.gen_range(-0.5..0.5))
.collect();
let labels: Vec<i32> = vec![5, 10, 15]; // distinct per horizon
let labels_d = upload_i32(&stream, &labels)?;
// Analytical gradient via the kernel.
let logits_d = upload_f32(&stream, &logits)?;
let mut grad_d = stream.alloc_zeros::<f32>(b_size * FRD_OUT_DIM)?;
let mut loss_d = stream.alloc_zeros::<f32>(b_size * FRD_N_HORIZONS)?;
head.softmax_ce_grad(&logits_d, &labels_d, &mut grad_d, &mut loss_d, b_size)?;
let grad_analytical = read_slice_d_pub(&stream, &grad_d, b_size * FRD_OUT_DIM)?;
// Finite-difference for slot (b=0, h=0, a=3). Note: gradient was
// mean-reduced by 1/B at the kernel source; for b_size=1 the 1/B
// factor is 1.0 so finite-diff matches directly.
let probe_h = 0_usize;
let probe_a = 3_usize;
let probe_off = probe_h * FRD_N_ATOMS + probe_a;
let eps = 1e-3_f32;
// L(logits + ε · e_j) — perturb only the target slot upward.
logits[probe_off] += eps;
let logits_plus_d = upload_f32(&stream, &logits)?;
head.softmax_ce_grad(&logits_plus_d, &labels_d, &mut grad_d, &mut loss_d, b_size)?;
let loss_plus = read_slice_d_pub(&stream, &loss_d, b_size * FRD_N_HORIZONS)?;
let l_plus = loss_plus[probe_h]; // only h=0 affected — h=1,2 share the perturbation only if probe was in their horizon block
// L(logits - ε · e_j)
logits[probe_off] -= 2.0 * eps;
let logits_minus_d = upload_f32(&stream, &logits)?;
head.softmax_ce_grad(&logits_minus_d, &labels_d, &mut grad_d, &mut loss_d, b_size)?;
let loss_minus = read_slice_d_pub(&stream, &loss_d, b_size * FRD_N_HORIZONS)?;
let l_minus = loss_minus[probe_h];
let numerical = (l_plus - l_minus) / (2.0 * eps);
let analytical = grad_analytical[probe_off];
let rel_err = (numerical - analytical).abs() / (analytical.abs().max(1e-6));
// 5e-3 tolerance: fp32 single-precision finite-difference at ε=1e-3
// saturates around the per-evaluation rounding error (~1e-7) divided
// by ε, so 1e-4-ish absolute error is intrinsic to the test method,
// not the kernel. Tighter threshold would force fp64 finite-diff
// (not worth the complexity here).
assert!(
rel_err < 5e-3,
"finite-diff gradient mismatch at (h={probe_h}, a={probe_a}): \
analytical={analytical:.6}, numerical={numerical:.6}, rel_err={rel_err:.6}"
);
eprintln!(
"PASS — finite-diff matches analytical: analytical={:.6} numerical={:.6} rel_err={:.2e}",
analytical, numerical, rel_err
);
Ok(())
}
/// Helper: compute total loss for given logits by calling the
/// softmax_ce_grad kernel and summing per-(b, h) CE.
fn ce_total_loss(
head: &FrdHead,
stream: &Arc<CudaStream>,
logits: &[f32],
labels_d: &CudaSlice<i32>,
b_size: usize,
) -> Result<f32> {
let logits_d = upload_f32(stream, logits)?;
let mut grad_d = stream.alloc_zeros::<f32>(b_size * FRD_OUT_DIM)?;
let mut loss_d = stream.alloc_zeros::<f32>(b_size * FRD_N_HORIZONS)?;
head.softmax_ce_grad(&logits_d, labels_d, &mut grad_d, &mut loss_d, b_size)?;
let loss = read_slice_d_pub(stream, &loss_d, b_size * FRD_N_HORIZONS)?;
Ok(loss.iter().sum())
}
#[test]
#[ignore = "requires CUDA (MlDevice::cuda(0))"]
fn frd_layer2_bwd_finite_diff_w2() -> Result<()> {
let Some((dev, mut head)) = build_head() else { return Ok(()) };
let stream = dev.cuda_stream()?.clone();
let b_size = 1;
// Forward pass with random h_t to produce hidden + logits.
use rand::{Rng, SeedableRng};
use rand_chacha::ChaCha8Rng;
let mut r = ChaCha8Rng::seed_from_u64(0xB2);
let h_t: Vec<f32> = (0..b_size * HIDDEN_DIM)
.map(|_| r.gen_range(-1.0..1.0))
.collect();
let h_t_d = upload_f32(&stream, &h_t)?;
let mut hidden_d = stream.alloc_zeros::<f32>(b_size * FRD_HIDDEN_DIM)?;
let mut logits_d = stream.alloc_zeros::<f32>(b_size * FRD_OUT_DIM)?;
head.forward(&h_t_d, &mut hidden_d, &mut logits_d, b_size)?;
let labels: Vec<i32> = vec![5, 10, 15];
let labels_d = upload_i32(&stream, &labels)?;
// Softmax+CE grad of logits.
let mut grad_logits_d = stream.alloc_zeros::<f32>(b_size * FRD_OUT_DIM)?;
let mut loss_d = stream.alloc_zeros::<f32>(b_size * FRD_N_HORIZONS)?;
head.softmax_ce_grad(&logits_d, &labels_d, &mut grad_logits_d, &mut loss_d, b_size)?;
// Layer-2 backward: produce per-batch grad_W2 scratch.
let mut grad_w2_pb_d =
stream.alloc_zeros::<f32>(b_size * FRD_HIDDEN_DIM * FRD_OUT_DIM)?;
let mut grad_b2_pb_d = stream.alloc_zeros::<f32>(b_size * FRD_OUT_DIM)?;
let mut grad_hidden_d = stream.alloc_zeros::<f32>(b_size * FRD_HIDDEN_DIM)?;
head.layer2_bwd(
&hidden_d,
&grad_logits_d,
&mut grad_w2_pb_d,
&mut grad_b2_pb_d,
&mut grad_hidden_d,
b_size,
)?;
// For b_size=1 the per-batch scratch IS the reduced gradient.
let grad_w2 = read_slice_d_pub(&stream, &grad_w2_pb_d, FRD_HIDDEN_DIM * FRD_OUT_DIM)?;
// Finite-diff probe: perturb W2[i=10, j=5] by ±ε, measure ΔL.
let probe_i = 10_usize;
let probe_j = 5_usize;
let probe_off = probe_i * FRD_OUT_DIM + probe_j;
let eps = 1e-3_f32;
let mut w2_host = read_slice_d_pub(&stream, &head.w2_d, FRD_HIDDEN_DIM * FRD_OUT_DIM)?;
let original = w2_host[probe_off];
// L(W2 + ε · e_(i,j)) — perturb weight on device, re-run forward,
// compute total CE loss across all (b, h) rows.
w2_host[probe_off] = original + eps;
write_slice_f32_d_pub(&stream, &w2_host, &mut head.w2_d)?;
let mut hidden_plus = stream.alloc_zeros::<f32>(b_size * FRD_HIDDEN_DIM)?;
let mut logits_plus = stream.alloc_zeros::<f32>(b_size * FRD_OUT_DIM)?;
head.forward(&h_t_d, &mut hidden_plus, &mut logits_plus, b_size)?;
let logits_plus_h = read_slice_d_pub(&stream, &logits_plus, b_size * FRD_OUT_DIM)?;
let l_plus = ce_total_loss(&head, &stream, &logits_plus_h, &labels_d, b_size)?;
// L(W2 - ε · e_(i,j))
w2_host[probe_off] = original - eps;
write_slice_f32_d_pub(&stream, &w2_host, &mut head.w2_d)?;
let mut hidden_minus = stream.alloc_zeros::<f32>(b_size * FRD_HIDDEN_DIM)?;
let mut logits_minus = stream.alloc_zeros::<f32>(b_size * FRD_OUT_DIM)?;
head.forward(&h_t_d, &mut hidden_minus, &mut logits_minus, b_size)?;
let logits_minus_h = read_slice_d_pub(&stream, &logits_minus, b_size * FRD_OUT_DIM)?;
let l_minus = ce_total_loss(&head, &stream, &logits_minus_h, &labels_d, b_size)?;
// Restore W2 to keep test isolation clean (next test gets default init).
w2_host[probe_off] = original;
write_slice_f32_d_pub(&stream, &w2_host, &mut head.w2_d)?;
let numerical = (l_plus - l_minus) / (2.0 * eps);
let analytical = grad_w2[probe_off];
let rel_err = (numerical - analytical).abs() / (analytical.abs().max(1e-6));
assert!(
rel_err < 5e-3,
"dW2 finite-diff mismatch at (i={probe_i}, j={probe_j}): \
analytical={analytical:.6}, numerical={numerical:.6}, rel_err={rel_err:.6}"
);
eprintln!(
"PASS — dW2 finite-diff: analytical={:.6} numerical={:.6} rel_err={:.2e}",
analytical, numerical, rel_err
);
Ok(())
}
#[test]
#[ignore = "requires CUDA (MlDevice::cuda(0))"]
fn frd_layer2_bwd_db2_equals_grad_logits() -> Result<()> {
// db2 invariant: per-batch grad_b2[b, j] = grad_logits[b, j].
// After reduce_axis0 across batch this becomes Σ_b grad_logits[b, j]
// (the standard bias gradient). Verify the per-batch scratch
// matches the input grad_logits exactly.
let Some((dev, head)) = build_head() else { return Ok(()) };
let stream = dev.cuda_stream()?.clone();
let b_size = 4;
use rand::{Rng, SeedableRng};
use rand_chacha::ChaCha8Rng;
let mut r = ChaCha8Rng::seed_from_u64(0xB22);
let h_t: Vec<f32> = (0..b_size * HIDDEN_DIM)
.map(|_| r.gen_range(-1.0..1.0))
.collect();
let h_t_d = upload_f32(&stream, &h_t)?;
let mut hidden_d = stream.alloc_zeros::<f32>(b_size * FRD_HIDDEN_DIM)?;
let mut logits_d = stream.alloc_zeros::<f32>(b_size * FRD_OUT_DIM)?;
head.forward(&h_t_d, &mut hidden_d, &mut logits_d, b_size)?;
let labels: Vec<i32> = (0..b_size * FRD_N_HORIZONS)
.map(|i| ((i * 3) % FRD_N_ATOMS) as i32)
.collect();
let labels_d = upload_i32(&stream, &labels)?;
let mut grad_logits_d = stream.alloc_zeros::<f32>(b_size * FRD_OUT_DIM)?;
let mut loss_d = stream.alloc_zeros::<f32>(b_size * FRD_N_HORIZONS)?;
head.softmax_ce_grad(&logits_d, &labels_d, &mut grad_logits_d, &mut loss_d, b_size)?;
let mut grad_w2_pb_d =
stream.alloc_zeros::<f32>(b_size * FRD_HIDDEN_DIM * FRD_OUT_DIM)?;
let mut grad_b2_pb_d = stream.alloc_zeros::<f32>(b_size * FRD_OUT_DIM)?;
let mut grad_hidden_d = stream.alloc_zeros::<f32>(b_size * FRD_HIDDEN_DIM)?;
head.layer2_bwd(
&hidden_d,
&grad_logits_d,
&mut grad_w2_pb_d,
&mut grad_b2_pb_d,
&mut grad_hidden_d,
b_size,
)?;
let grad_logits = read_slice_d_pub(&stream, &grad_logits_d, b_size * FRD_OUT_DIM)?;
let grad_b2_pb = read_slice_d_pub(&stream, &grad_b2_pb_d, b_size * FRD_OUT_DIM)?;
for (i, (gl, gb)) in grad_logits.iter().zip(grad_b2_pb.iter()).enumerate() {
assert_eq!(*gl, *gb, "db2_per_batch[{i}] must equal grad_logits[{i}]");
}
eprintln!(
"PASS — db2_per_batch ≡ grad_logits across {} slots",
grad_logits.len()
);
Ok(())
}
#[test]
#[ignore = "requires CUDA (MlDevice::cuda(0))"]
fn frd_layer1_bwd_finite_diff_w1() -> Result<()> {
let Some((dev, mut head)) = build_head() else { return Ok(()) };
let stream = dev.cuda_stream()?.clone();
let b_size = 1;
use rand::{Rng, SeedableRng};
use rand_chacha::ChaCha8Rng;
let mut r = ChaCha8Rng::seed_from_u64(0xB1);
let h_t: Vec<f32> = (0..b_size * HIDDEN_DIM)
.map(|_| r.gen_range(-1.0..1.0))
.collect();
let h_t_d = upload_f32(&stream, &h_t)?;
let labels: Vec<i32> = vec![5, 10, 15];
let labels_d = upload_i32(&stream, &labels)?;
// Full backward chain at the unperturbed weights.
let mut hidden_d = stream.alloc_zeros::<f32>(b_size * FRD_HIDDEN_DIM)?;
let mut logits_d = stream.alloc_zeros::<f32>(b_size * FRD_OUT_DIM)?;
head.forward(&h_t_d, &mut hidden_d, &mut logits_d, b_size)?;
let mut grad_logits_d = stream.alloc_zeros::<f32>(b_size * FRD_OUT_DIM)?;
let mut loss_d = stream.alloc_zeros::<f32>(b_size * FRD_N_HORIZONS)?;
head.softmax_ce_grad(&logits_d, &labels_d, &mut grad_logits_d, &mut loss_d, b_size)?;
let mut grad_w2_pb_d = stream.alloc_zeros::<f32>(b_size * FRD_HIDDEN_DIM * FRD_OUT_DIM)?;
let mut grad_b2_pb_d = stream.alloc_zeros::<f32>(b_size * FRD_OUT_DIM)?;
let mut grad_hidden_d = stream.alloc_zeros::<f32>(b_size * FRD_HIDDEN_DIM)?;
head.layer2_bwd(
&hidden_d,
&grad_logits_d,
&mut grad_w2_pb_d,
&mut grad_b2_pb_d,
&mut grad_hidden_d,
b_size,
)?;
let mut grad_w1_pb_d = stream.alloc_zeros::<f32>(b_size * HIDDEN_DIM * FRD_HIDDEN_DIM)?;
let mut grad_b1_pb_d = stream.alloc_zeros::<f32>(b_size * FRD_HIDDEN_DIM)?;
let mut grad_h_t_d = stream.alloc_zeros::<f32>(b_size * HIDDEN_DIM)?;
head.layer1_bwd(
&h_t_d,
&hidden_d,
&grad_hidden_d,
&mut grad_w1_pb_d,
&mut grad_b1_pb_d,
&mut grad_h_t_d,
b_size,
)?;
// Finite-diff probe: perturb W1[k=30, i=20] by ±ε. Pick an
// (k, i) likely to land on an active ReLU branch (use a moderate
// i within the 64 hidden slots — random init usually has ~50%
// active so any pick has ~50% chance of being a non-mask slot).
let probe_k = 30_usize;
let probe_i = 20_usize;
let probe_off = probe_k * FRD_HIDDEN_DIM + probe_i;
let eps = 1e-3_f32;
let grad_w1 = read_slice_d_pub(&stream, &grad_w1_pb_d, HIDDEN_DIM * FRD_HIDDEN_DIM)?;
// Pick the slot with MAX |analytical| — fp32 finite-diff rel_err
// is dominated by absolute rounding noise (~1e-5 at ε=1e-3 across
// the deep matmul chain), so a tiny-magnitude gradient produces
// misleading relative error. Probing the largest gradient gives
// the cleanest sanity check.
let _ = (probe_k, probe_i, probe_off);
let (chosen_off, &chosen_analytical) = grad_w1
.iter()
.enumerate()
.max_by(|(_, a), (_, b)| a.abs().partial_cmp(&b.abs()).unwrap())
.expect("non-empty grad_w1");
let chosen_k = chosen_off / FRD_HIDDEN_DIM;
let chosen_i = chosen_off % FRD_HIDDEN_DIM;
assert!(
chosen_analytical.abs() > 1e-3,
"max |dW1| should be at least 1e-3 for a meaningful finite-diff; got {chosen_analytical}"
);
let mut w1_host = read_slice_d_pub(&stream, &head.w1_d, HIDDEN_DIM * FRD_HIDDEN_DIM)?;
let original = w1_host[chosen_off];
// L(W1 + ε · e_(k,i))
w1_host[chosen_off] = original + eps;
write_slice_f32_d_pub(&stream, &w1_host, &mut head.w1_d)?;
let mut hidden_plus = stream.alloc_zeros::<f32>(b_size * FRD_HIDDEN_DIM)?;
let mut logits_plus = stream.alloc_zeros::<f32>(b_size * FRD_OUT_DIM)?;
head.forward(&h_t_d, &mut hidden_plus, &mut logits_plus, b_size)?;
let logits_plus_h = read_slice_d_pub(&stream, &logits_plus, b_size * FRD_OUT_DIM)?;
let l_plus = ce_total_loss(&head, &stream, &logits_plus_h, &labels_d, b_size)?;
// L(W1 - ε · e_(k,i))
w1_host[chosen_off] = original - eps;
write_slice_f32_d_pub(&stream, &w1_host, &mut head.w1_d)?;
let mut hidden_minus = stream.alloc_zeros::<f32>(b_size * FRD_HIDDEN_DIM)?;
let mut logits_minus = stream.alloc_zeros::<f32>(b_size * FRD_OUT_DIM)?;
head.forward(&h_t_d, &mut hidden_minus, &mut logits_minus, b_size)?;
let logits_minus_h = read_slice_d_pub(&stream, &logits_minus, b_size * FRD_OUT_DIM)?;
let l_minus = ce_total_loss(&head, &stream, &logits_minus_h, &labels_d, b_size)?;
// Restore W1.
w1_host[chosen_off] = original;
write_slice_f32_d_pub(&stream, &w1_host, &mut head.w1_d)?;
let numerical = (l_plus - l_minus) / (2.0 * eps);
let rel_err = (numerical - chosen_analytical).abs() / (chosen_analytical.abs().max(1e-6));
// 1e-2 tolerance is slightly looser than dW2's 5e-3 because the
// dW1 chain crosses an extra matmul + the ReLU mask boundary
// (which is a non-differentiable point — at hidden ≈ 0 the
// finite-diff straddles the mask discontinuity).
assert!(
rel_err < 1e-2,
"dW1 finite-diff mismatch at (k={chosen_k}, i={chosen_i}): \
analytical={chosen_analytical:.6}, numerical={numerical:.6}, rel_err={rel_err:.6}"
);
eprintln!(
"PASS — dW1 finite-diff at (k={chosen_k}, i={chosen_i}): \
analytical={:.6} numerical={:.6} rel_err={:.2e}",
chosen_analytical, numerical, rel_err
);
Ok(())
}
#[test]
#[ignore = "requires CUDA (MlDevice::cuda(0))"]
fn frd_layer1_bwd_relu_mask_zeros_grad() -> Result<()> {
// Invariant: if cached `hidden[b, i] == 0` (i.e. the pre-activation
// landed on the negative half of ReLU), then:
// * grad_b1_per_batch[b, i] = 0 (db1 = dpre_hidden = grad_hidden × mask)
// * grad_W1_per_batch[b, k, i] = 0 for every k
//
// Construct a fixture where some hidden slots are guaranteed
// negative: feed h_t = all -1.0 (so half of Xavier-init W1's
// products average negative).
let Some((dev, head)) = build_head() else { return Ok(()) };
let stream = dev.cuda_stream()?.clone();
let b_size = 1;
let h_t = vec![-1.0_f32; b_size * HIDDEN_DIM];
let h_t_d = upload_f32(&stream, &h_t)?;
let labels: Vec<i32> = vec![5, 10, 15];
let labels_d = upload_i32(&stream, &labels)?;
let mut hidden_d = stream.alloc_zeros::<f32>(b_size * FRD_HIDDEN_DIM)?;
let mut logits_d = stream.alloc_zeros::<f32>(b_size * FRD_OUT_DIM)?;
head.forward(&h_t_d, &mut hidden_d, &mut logits_d, b_size)?;
let mut grad_logits_d = stream.alloc_zeros::<f32>(b_size * FRD_OUT_DIM)?;
let mut loss_d = stream.alloc_zeros::<f32>(b_size * FRD_N_HORIZONS)?;
head.softmax_ce_grad(&logits_d, &labels_d, &mut grad_logits_d, &mut loss_d, b_size)?;
let mut grad_w2_pb_d = stream.alloc_zeros::<f32>(b_size * FRD_HIDDEN_DIM * FRD_OUT_DIM)?;
let mut grad_b2_pb_d = stream.alloc_zeros::<f32>(b_size * FRD_OUT_DIM)?;
let mut grad_hidden_d = stream.alloc_zeros::<f32>(b_size * FRD_HIDDEN_DIM)?;
head.layer2_bwd(
&hidden_d,
&grad_logits_d,
&mut grad_w2_pb_d,
&mut grad_b2_pb_d,
&mut grad_hidden_d,
b_size,
)?;
let mut grad_w1_pb_d = stream.alloc_zeros::<f32>(b_size * HIDDEN_DIM * FRD_HIDDEN_DIM)?;
let mut grad_b1_pb_d = stream.alloc_zeros::<f32>(b_size * FRD_HIDDEN_DIM)?;
let mut grad_h_t_d = stream.alloc_zeros::<f32>(b_size * HIDDEN_DIM)?;
head.layer1_bwd(
&h_t_d,
&hidden_d,
&grad_hidden_d,
&mut grad_w1_pb_d,
&mut grad_b1_pb_d,
&mut grad_h_t_d,
b_size,
)?;
let hidden = read_slice_d_pub(&stream, &hidden_d, b_size * FRD_HIDDEN_DIM)?;
let grad_b1 = read_slice_d_pub(&stream, &grad_b1_pb_d, b_size * FRD_HIDDEN_DIM)?;
let grad_w1 = read_slice_d_pub(&stream, &grad_w1_pb_d, b_size * HIDDEN_DIM * FRD_HIDDEN_DIM)?;
let mut masked_count = 0;
let mut unmasked_count = 0;
for i in 0..FRD_HIDDEN_DIM {
if hidden[i] == 0.0 {
masked_count += 1;
// ReLU mask zeros dpre_hidden → db1 + entire column of dW1 must be 0.
assert_eq!(
grad_b1[i], 0.0,
"db1[{i}] should be 0 under ReLU mask (hidden=0); got {}",
grad_b1[i]
);
for k in 0..HIDDEN_DIM {
let off = k * FRD_HIDDEN_DIM + i;
assert_eq!(
grad_w1[off], 0.0,
"dW1[{k}, {i}] should be 0 under ReLU mask (hidden=0); got {}",
grad_w1[off]
);
}
} else {
unmasked_count += 1;
}
}
assert!(
masked_count > 0,
"expected at least one ReLU-masked hidden slot under h_t = all -1; \
got 0 (all-positive hidden — Xavier init may have flipped signs)"
);
eprintln!(
"PASS — ReLU mask zeros dW1 + db1 for {} of {} hidden slots ({} active)",
masked_count, FRD_HIDDEN_DIM, unmasked_count
);
Ok(())
}