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.
738 lines
30 KiB
Rust
738 lines
30 KiB
Rust
//! 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(())
|
||
}
|