//! Finite-difference validation of cfc_step_backward + multi_horizon_heads_backward. //! //! Per `feedback_no_cpu_test_fallbacks.md`: the kernel is the reference. //! We perturb a single input, re-run the GPU forward, and compare the //! resulting finite-difference gradient to the analytic backward. use approx::assert_relative_eq; use ml_alpha::cfc::step::{cfc_step_backward_gpu, cfc_step_gpu, CfcWeights}; use ml_alpha::heads::{ multi_horizon_heads_backward_gpu, multi_horizon_heads_gpu, HeadsWeights, HIDDEN_DIM, N_HORIZONS, }; 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_heads(seed: u64) -> HeadsWeights { let mut r = ChaCha8Rng::seed_from_u64(seed); HeadsWeights { w: (0..N_HORIZONS * HIDDEN_DIM).map(|_| r.gen_range(-0.1..0.1)).collect(), b: (0..N_HORIZONS).map(|_| r.gen_range(-0.01..0.01)).collect(), } } fn rand_cfc(seed: u64, n_in: usize, n_hid: usize) -> CfcWeights { let mut r = ChaCha8Rng::seed_from_u64(seed); 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: (0..n_hid).map(|_| r.gen_range(-0.01..0.01)).collect(), tau: (0..n_hid).map(|_| r.gen_range(0.05..2.0)).collect(), n_in, n_hid, } } #[test] fn heads_backward_grad_h_matches_finite_diff() { let dev = test_device(); let w = rand_heads(0x1111); let mut r = ChaCha8Rng::seed_from_u64(0x2222); let h: Vec = (0..HIDDEN_DIM).map(|_| r.gen_range(-1.0..1.0)).collect(); let probs = multi_horizon_heads_gpu(&dev, &w, &h).unwrap(); // Surrogate loss L = sum_k probs[k] -> dL/dp = 1 for all k. let grad_probs = [1.0f32; N_HORIZONS]; let (_, _, grad_h_analytic) = multi_horizon_heads_backward_gpu(&dev, &w, &probs, &h, &grad_probs).unwrap(); let eps = 1e-3_f32; for &i in &[0usize, 32, 64, 96, 127] { let mut h_up = h.clone(); h_up[i] += eps; let mut h_dn = h.clone(); h_dn[i] -= eps; let p_up = multi_horizon_heads_gpu(&dev, &w, &h_up).unwrap(); let p_dn = multi_horizon_heads_gpu(&dev, &w, &h_dn).unwrap(); let l_up: f32 = p_up.iter().sum(); let l_dn: f32 = p_dn.iter().sum(); let fd = (l_up - l_dn) / (2.0 * eps); assert_relative_eq!(grad_h_analytic[i], fd, epsilon = 1e-3, max_relative = 1e-2); } } #[test] fn heads_backward_grad_b_matches_finite_diff() { let dev = test_device(); let w = rand_heads(0x3333); let mut r = ChaCha8Rng::seed_from_u64(0x4444); let h: Vec = (0..HIDDEN_DIM).map(|_| r.gen_range(-1.0..1.0)).collect(); let probs = multi_horizon_heads_gpu(&dev, &w, &h).unwrap(); let grad_probs = [1.0f32; N_HORIZONS]; let (_, grad_b_analytic, _) = multi_horizon_heads_backward_gpu(&dev, &w, &probs, &h, &grad_probs).unwrap(); let eps = 1e-3_f32; for k in 0..N_HORIZONS { let mut w_up = w.clone(); w_up.b[k] += eps; let mut w_dn = w.clone(); w_dn.b[k] -= eps; let p_up = multi_horizon_heads_gpu(&dev, &w_up, &h).unwrap(); let p_dn = multi_horizon_heads_gpu(&dev, &w_dn, &h).unwrap(); let l_up: f32 = p_up.iter().sum(); let l_dn: f32 = p_dn.iter().sum(); let fd = (l_up - l_dn) / (2.0 * eps); assert_relative_eq!(grad_b_analytic[k], fd, epsilon = 1e-3, max_relative = 1e-2); let _ = w; // suppress let _ = w_up; let _ = w_dn; } let _ = w; } #[test] fn cfc_backward_grad_b_matches_finite_diff() { let dev = test_device(); let n_in = 8; let n_hid = 16; let w = rand_cfc(0x5555, n_in, n_hid); let mut r = ChaCha8Rng::seed_from_u64(0x6666); let x: Vec = (0..n_in).map(|_| r.gen_range(-1.0..1.0)).collect(); let h_old: Vec = (0..n_hid).map(|_| r.gen_range(-1.0..1.0)).collect(); let dt_s = 0.02_f32; // Surrogate loss L = sum_i h_new[i]. let grad_h_new = vec![1.0f32; n_hid]; let (_, _, grad_b_analytic, _, _) = cfc_step_backward_gpu(&dev, &w, &x, &h_old, &grad_h_new, dt_s).unwrap(); let eps = 1e-3_f32; for &i in &[0usize, 4, 8, 15] { let mut w_up = w.clone(); w_up.b[i] += eps; let mut w_dn = w.clone(); w_dn.b[i] -= eps; let h_up = cfc_step_gpu(&dev, &w_up, &x, &h_old, dt_s).unwrap(); let h_dn = cfc_step_gpu(&dev, &w_dn, &x, &h_old, dt_s).unwrap(); let l_up: f32 = h_up.iter().sum(); let l_dn: f32 = h_dn.iter().sum(); let fd = (l_up - l_dn) / (2.0 * eps); assert_relative_eq!(grad_b_analytic[i], fd, epsilon = 5e-3, max_relative = 5e-2); } } #[test] fn cfc_backward_grad_h_old_matches_finite_diff() { let dev = test_device(); let n_in = 8; let n_hid = 16; let w = rand_cfc(0x7777, n_in, n_hid); let mut r = ChaCha8Rng::seed_from_u64(0x8888); let x: Vec = (0..n_in).map(|_| r.gen_range(-1.0..1.0)).collect(); let h_old: Vec = (0..n_hid).map(|_| r.gen_range(-1.0..1.0)).collect(); let dt_s = 0.02_f32; let grad_h_new = vec![1.0f32; n_hid]; let (_, _, _, grad_h_old_analytic, _) = cfc_step_backward_gpu(&dev, &w, &x, &h_old, &grad_h_new, dt_s).unwrap(); let eps = 1e-3_f32; for &i in &[0usize, 4, 8, 15] { let mut h_up = h_old.clone(); h_up[i] += eps; let mut h_dn = h_old.clone(); h_dn[i] -= eps; let h_new_up = cfc_step_gpu(&dev, &w, &x, &h_up, dt_s).unwrap(); let h_new_dn = cfc_step_gpu(&dev, &w, &x, &h_dn, dt_s).unwrap(); let l_up: f32 = h_new_up.iter().sum(); let l_dn: f32 = h_new_dn.iter().sum(); let fd = (l_up - l_dn) / (2.0 * eps); assert_relative_eq!(grad_h_old_analytic[i], fd, epsilon = 5e-3, max_relative = 5e-2); } }