Files
foxhunt/crates/ml-alpha/tests/cfc_step_bit_equiv.rs
jgrusewski f927469ed3 feat(ml-alpha): cfc_step kernel + invariant-based GPU validation
Hasani 2022 closed-form CfC recurrence; one thread per hidden unit, no
atomicAdd. Tests assert algebraic invariants (dt=0 -> identity, zero
weights -> h_old * decay, large tau -> h_old preserved, output bound).

Also removes src/cfc/oracle.rs and replaces snap_feature bit-equiv
test with property assertions per feedback_no_cpu_test_fallbacks.md.
CPU mirrors are bug-locks; validation is now via known synthetic
inputs + analytical relations on the GPU output.

12 tests pass on local sm_86 (7 snap_feature invariants + 5 cfc_step
invariants).

Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
2026-05-16 21:50:20 +02:00

136 lines
4.8 KiB
Rust

//! CfC step GPU invariants.
//!
//! Per `feedback_no_cpu_test_fallbacks.md`: no CPU oracle. Validation is
//! via analytically-known invariants of the closed-form recurrence:
//!
//! h_new[i] = h_old[i] * decay[i] + (1 - decay[i]) * tanh(W_in·x + W_rec·h_old + b)
//! decay[i] = exp(-dt_s / tau[i])
//!
//! Key invariants tested:
//! - dt = 0 → decay = 1 → h_new = h_old (identity)
//! - tau very large with dt small → decay ≈ 1 → h_new ≈ h_old
//! - Zero weights + zero bias → pre = 0 → tanh(0) = 0 → h_new = h_old · decay
//! - Output is bounded by max(|h_old|, 1) since tanh ∈ [-1, 1]
use approx::assert_relative_eq;
use ml_alpha::cfc::step::{cfc_step_gpu, CfcWeights};
use ml_core::device::MlDevice;
use rand::{Rng, SeedableRng};
use rand_chacha::ChaCha8Rng;
fn test_device() -> MlDevice {
MlDevice::cuda(0).expect("CUDA 0 required for ml-alpha tests")
}
fn rand_weights(seed: u64, n_in: usize, n_hid: usize) -> CfcWeights {
let mut r = ChaCha8Rng::seed_from_u64(seed);
let w_in: Vec<f32> = (0..n_hid * n_in).map(|_| r.gen_range(-0.1..0.1)).collect();
let w_rec: Vec<f32> = (0..n_hid * n_hid).map(|_| r.gen_range(-0.05..0.05)).collect();
let b: Vec<f32> = (0..n_hid).map(|_| r.gen_range(-0.01..0.01)).collect();
let tau: Vec<f32> = (0..n_hid).map(|_| r.gen_range(0.05..2.0)).collect();
CfcWeights {
w_in,
w_rec,
b,
tau,
n_in,
n_hid,
}
}
#[test]
fn cfc_step_zero_dt_returns_h_old() {
let dev = test_device();
let w = rand_weights(0xDEAD_BEEF, 32, 128);
let h_old: Vec<f32> = (0..w.n_hid).map(|i| 0.01 * i as f32).collect();
let x = vec![0.0; w.n_in];
let h_gpu = cfc_step_gpu(&dev, &w, &x, &h_old, 0.0).expect("gpu");
for i in 0..w.n_hid {
assert_relative_eq!(h_gpu[i], h_old[i], epsilon = 1e-6);
}
}
#[test]
fn cfc_step_zero_weights_decays_h_old() {
// Zero W_in, W_rec, b → pre = 0 → tanh(0) = 0 → h_new = h_old · decay.
let dev = test_device();
let n_in = 32;
let n_hid = 64;
let w = CfcWeights {
w_in: vec![0.0; n_hid * n_in],
w_rec: vec![0.0; n_hid * n_hid],
b: vec![0.0; n_hid],
tau: vec![1.0; n_hid],
n_in,
n_hid,
};
let h_old: Vec<f32> = (0..n_hid).map(|i| 0.5 + 0.01 * i as f32).collect();
let dt_s = 0.5_f32;
let expected_decay = (-dt_s / 1.0_f32).exp();
let h_gpu = cfc_step_gpu(&dev, &w, &vec![0.0; n_in], &h_old, dt_s).expect("gpu");
for i in 0..n_hid {
let expected = h_old[i] * expected_decay;
assert_relative_eq!(h_gpu[i], expected, epsilon = 1e-5);
}
}
#[test]
fn cfc_step_large_tau_preserves_h_old() {
// tau >> dt → decay ≈ 1 → h_new ≈ h_old.
let dev = test_device();
let n_in = 16;
let n_hid = 32;
let mut r = ChaCha8Rng::seed_from_u64(7);
let w = CfcWeights {
w_in: (0..n_hid * n_in).map(|_| r.gen_range(-0.1..0.1)).collect(),
w_rec: (0..n_hid * n_hid).map(|_| r.gen_range(-0.05..0.05)).collect(),
b: vec![0.0; n_hid],
tau: vec![10_000.0; n_hid], // huge
n_in,
n_hid,
};
let h_old: Vec<f32> = (0..n_hid).map(|_| r.gen_range(-0.5..0.5)).collect();
let x: Vec<f32> = (0..n_in).map(|_| r.gen_range(-0.5..0.5)).collect();
let h_gpu = cfc_step_gpu(&dev, &w, &x, &h_old, 0.001).expect("gpu");
for i in 0..n_hid {
// decay = exp(-0.001/10000) ≈ 1 - 1e-7 ≈ 1
assert_relative_eq!(h_gpu[i], h_old[i], epsilon = 1e-3);
}
}
#[test]
fn cfc_step_output_bounded_by_one_plus_h_old_norm() {
// tanh ∈ [-1, 1], decay ∈ [0, 1], so |h_new[i]| ≤ |h_old[i]| · 1 + 1 · 1
let dev = test_device();
let w = rand_weights(0xCFC_5EED, 32, 128);
let mut r = ChaCha8Rng::seed_from_u64(42);
let x: Vec<f32> = (0..w.n_in).map(|_| r.gen_range(-1.0..1.0)).collect();
let h_old: Vec<f32> = (0..w.n_hid).map(|_| r.gen_range(-1.0..1.0)).collect();
let h_gpu = cfc_step_gpu(&dev, &w, &x, &h_old, 0.02).expect("gpu");
for i in 0..w.n_hid {
let bound = h_old[i].abs() + 1.0;
assert!(
h_gpu[i].abs() <= bound + 1e-5,
"h_new[{i}]={} exceeds |h_old|+1 bound {}",
h_gpu[i],
bound
);
}
}
#[test]
fn cfc_step_runs_on_small_hidden() {
// n_hid < 128 — exercises the grid-dim ceiling. Just assert the call
// succeeds and outputs are finite.
let dev = test_device();
let w = rand_weights(0x1234, 8, 16);
let mut r = ChaCha8Rng::seed_from_u64(99);
let x: Vec<f32> = (0..w.n_in).map(|_| r.gen_range(-1.0..1.0)).collect();
let h_old: Vec<f32> = (0..w.n_hid).map(|_| r.gen_range(-1.0..1.0)).collect();
let h_gpu = cfc_step_gpu(&dev, &w, &x, &h_old, 0.01).expect("gpu");
assert_eq!(h_gpu.len(), w.n_hid);
for v in &h_gpu {
assert!(v.is_finite(), "h_new must be finite, got {}", v);
}
}