//! Milestone-1 gate for the real selective-scan Mamba (Phase-3 spec). //! //! These tests prove the `MambaBlock` forward is now a genuine S6 //! selective-state-space model — not the prior passthrough stub: //! - **Liveness**: every weight that was *dead* in the stub //! (`conv1d_weight`, `A_log`, `dt_proj`, `x_proj`, `D`) now changes //! the output when perturbed. //! - **Causality**: output at step `l` depends only on inputs `≤ l` //! (no future leakage through the causal conv + scan). //! - **Determinism**: a seeded block is bit-stable across builds. //! - **Scan dynamics**: a non-trivial input produces non-constant, //! finite output with real temporal mixing. use std::collections::HashMap; use rtx_tensor::{Device, Tensor}; use rtx_transformers::layers::mamba::{MambaBlock, MambaConfig}; const D_MODEL: usize = 8; const D_STATE: usize = 16; const D_CONV: usize = 4; const L: usize = 6; fn cfg() -> MambaConfig { MambaConfig::new(D_MODEL, D_STATE, D_CONV) } fn fwd(block: &MambaBlock, x: &Tensor) -> Vec { block .forward(x) .expect("forward") .output .to_vec() .expect("to_vec") } fn max_abs_diff(a: &[f32], b: &[f32]) -> f32 { a.iter() .zip(b) .map(|(x, y)| (x - y).abs()) .fold(0.0, f32::max) } /// Clone all persistence tensors into a map, replacing `target` with a /// copy whose every element is shifted by `delta`. fn perturbed_map( block: &MambaBlock, target: &str, delta: f32, dev: &Device, ) -> HashMap { let mut map = HashMap::new(); for (name, t) in block.persistence_tensors() { if name == target { let mut v = t.to_vec().expect("to_vec"); for x in v.iter_mut() { *x += delta; } let pt = Tensor::from_vec(v, t.shape().dims(), dev).expect("from_vec"); map.insert(name.to_string(), pt); } else { map.insert(name.to_string(), t.clone()); } } map } /// A seeded block with `dt_bias` overridden to 0 ⇒ `Δ = softplus(0) ≈ /// 0.69`. The default init uses a deliberately small `Δ≈0.01` (the scan /// is near-identity — correct for training stability), which makes the /// scan-only weights' influence vanish below the f32 floor. This widens /// `Δ` so the scan genuinely contributes to the output, letting us /// observe each weight's effect. fn active_block(cfg: &MambaConfig, dev: &Device, seed: u64) -> MambaBlock { let base = MambaBlock::new_seeded(cfg.clone(), dev, seed).expect("base"); let d = cfg.get_d_inner(); let mut map = HashMap::new(); for (name, t) in base.persistence_tensors() { if name == "dt_bias" { map.insert( name.to_string(), Tensor::from_vec(vec![0.0f32; d], &[d], dev).expect("dtb"), ); } else { map.insert(name.to_string(), t.clone()); } } MambaBlock::from_persistence_tensors(cfg.clone(), dev, map).expect("rebuild") } #[test] fn every_formerly_dead_weight_is_live() { let dev = Device::cpu(); let block = active_block(&cfg(), &dev, 42); let x = Tensor::randn_seeded(&[1, L, D_MODEL], &dev, 7).expect("x"); let y0 = fwd(&block, &x); // In the stub these were all discarded; each must now move the output. for name in ["conv1d_weight", "A_log", "dt_proj", "x_proj", "D"] { let map = perturbed_map(&block, name, 0.5, &dev); let b2 = MambaBlock::from_persistence_tensors(cfg(), &dev, map).expect("rebuild"); let y2 = fwd(&b2, &x); let diff = max_abs_diff(&y0, &y2); assert!( diff > 1e-4, "weight `{name}` is DEAD — perturbing it left the output unchanged (max_abs_diff={diff:.2e})" ); } } #[test] fn forward_is_causal() { let dev = Device::cpu(); let block = MambaBlock::new_seeded(cfg(), &dev, 11).expect("block"); let base = Tensor::randn_seeded(&[1, L, D_MODEL], &dev, 3).expect("x"); let y0 = fwd(&block, &base); // Perturb ONLY the last timestep. let mut xv = base.to_vec().expect("to_vec"); for m in 0..D_MODEL { xv[(L - 1) * D_MODEL + m] += 1.0; } let x2 = Tensor::from_vec(xv, &[1, L, D_MODEL], &dev).expect("x2"); let y2 = fwd(&block, &x2); // Outputs for steps 0..L-1 must be identical (no future leakage); // the last step must change. let prefix0 = &y0[..(L - 1) * D_MODEL]; let prefix2 = &y2[..(L - 1) * D_MODEL]; assert!( max_abs_diff(prefix0, prefix2) < 1e-6, "non-causal: changing input[L-1] altered an earlier output" ); let last0 = &y0[(L - 1) * D_MODEL..]; let last2 = &y2[(L - 1) * D_MODEL..]; assert!( max_abs_diff(last0, last2) > 1e-4, "last output did not respond to the last input" ); } #[test] fn seeded_forward_is_deterministic() { let dev = Device::cpu(); let a = MambaBlock::new_seeded(cfg(), &dev, 99).expect("a"); let b = MambaBlock::new_seeded(cfg(), &dev, 99).expect("b"); let x = Tensor::randn_seeded(&[1, L, D_MODEL], &dev, 5).expect("x"); assert_eq!( fwd(&a, &x), fwd(&b, &x), "same (config, seed) must be bit-exact" ); } /// Build a seeded block whose every weight is scaled by `scale`, so the /// SSM operates in a well-conditioned regime (small `A_log` ⇒ `A≈-1` ⇒ /// the scan neither saturates nor underflows), giving a meaningful /// finite-difference signal. fn scaled_block(cfg: &MambaConfig, dev: &Device, seed: u64, scale: f32) -> MambaBlock { let base = MambaBlock::new_seeded(cfg.clone(), dev, seed).expect("base"); let mut map = HashMap::new(); for (name, t) in base.persistence_tensors() { let mut v = t.to_vec().expect("to_vec"); for x in v.iter_mut() { *x *= scale; } map.insert( name.to_string(), Tensor::from_vec(v, t.shape().dims(), dev).expect("from_vec"), ); } MambaBlock::from_persistence_tensors(cfg.clone(), dev, map).expect("rebuild") } fn loss(block: &MambaBlock, x: &Tensor, dov: &[f32]) -> f32 { let out = block .forward(x) .expect("fwd") .output .to_vec() .expect("to_vec"); out.iter().zip(dov).map(|(o, g)| o * g).sum::() } fn perturbed_one( block: &MambaBlock, name: &str, idx: usize, eps: f32, cfg: &MambaConfig, dev: &Device, ) -> MambaBlock { let mut map = HashMap::new(); for (nm, t) in block.persistence_tensors() { let mut v = t.to_vec().expect("to_vec"); if nm == name { v[idx] += eps; } map.insert( nm.to_string(), Tensor::from_vec(v, t.shape().dims(), dev).expect("from_vec"), ); } MambaBlock::from_persistence_tensors(cfg.clone(), dev, map).expect("rebuild") } #[test] fn analytic_gradients_match_finite_differences() { let dev = Device::cpu(); // Small, well-conditioned instance for a clean f32 FD signal. let cfg = MambaConfig::new(4, 4, 3); let block = scaled_block(&cfg, &dev, 21, 0.2); let ll = 4usize; let dmodel = 4usize; // Input and a fixed upstream gradient d_out (so L = Σ d_out·out). let xvec: Vec = (0..ll * dmodel) .map(|i| 0.5 * ((i as f32 * 0.37).sin())) .collect(); let x = Tensor::from_vec(xvec, &[1, ll, dmodel], &dev).expect("x"); let dov: Vec = (0..ll * dmodel) .map(|i| ((i as f32 * 0.91).cos())) .collect(); let grads = block .backward( &x, &Tensor::from_vec(dov.clone(), &[1, ll, dmodel], &dev).expect("dov"), ) .expect("backward"); let eps = 1e-2f32; let mut checked = 0; for name in [ "in_proj", "conv1d_weight", "conv1d_bias", "A_log", "x_proj", "dt_proj", "dt_bias", "D", "out_proj", ] { let g = grads .get(name) .unwrap_or_else(|| panic!("missing grad {name}")) .to_vec() .expect("g"); let len = g.len(); // Sample up to 4 spread-out indices per parameter. let idxs: Vec = if len <= 4 { (0..len).collect() } else { vec![0, len / 4, len / 2, (3 * len) / 4] }; for &idx in &idxs { let bp = perturbed_one(&block, name, idx, eps, &cfg, &dev); let bm = perturbed_one(&block, name, idx, -eps, &cfg, &dev); let fd = (loss(&bp, &x, &dov) - loss(&bm, &x, &dov)) / (2.0 * eps); let an = g[idx]; let tol = 5e-3 + 5e-2 * fd.abs(); assert!( (an - fd).abs() <= tol, "grad mismatch for `{name}`[{idx}]: analytic={an:.6} finite-diff={fd:.6} (tol {tol:.4})" ); checked += 1; } } assert!( checked >= 30, "expected to check ≥30 grad elements, got {checked}" ); } fn named_param(block: &MambaBlock, name: &str) -> Vec { block .persistence_tensors() .into_iter() .find(|(nm, _)| *nm == name) .map(|(_, t)| t.to_vec().expect("to_vec")) .expect("param") } /// Milestone-3: a self-contained Adam loop drives the loss down and /// actually moves the backbone weights — i.e. the block is *trainable*. /// The target is a teacher block's output on the same input (a learnable /// regression at matched capacity), so a real gradient signal must shrink /// the loss substantially. #[test] fn training_loop_reduces_loss_and_moves_weights() { let dev = Device::cpu(); let cfg = MambaConfig::new(4, 4, 3); let ll = 4usize; let dmodel = 4usize; let teacher = scaled_block(&cfg, &dev, 99, 0.2); let xvec: Vec = (0..ll * dmodel) .map(|i| 0.5 * ((i as f32 * 0.41).sin())) .collect(); let x = Tensor::from_vec(xvec, &[1, ll, dmodel], &dev).expect("x"); let target = teacher .forward(&x) .expect("teacher") .output .to_vec() .expect("t"); let mut block = scaled_block(&cfg, &dev, 5, 0.2); // conv1d_weight is a seed-varying backbone weight (A_log/dt_bias/D // are deterministic init, identical between teacher and student, so // they carry little gradient here) — it must move if the SSM block // (not just a head) is being trained. let cw_before = named_param(&block, "conv1d_weight"); let (lr, b1, b2, eps) = (0.03f32, 0.9f32, 0.999f32, 1e-8f32); let mut state: HashMap, Vec)> = HashMap::new(); let mut first_loss = 0.0f32; let mut last_loss = 0.0f32; for step in 1..=200i32 { let out = block.forward(&x).expect("fwd").output.to_vec().expect("o"); let ne = out.len(); let mut dov = vec![0.0f32; ne]; let mut loss = 0.0f32; for i in 0..ne { let e = out[i] - target[i]; loss += e * e; dov[i] = 2.0 * e / (ne as f32); } loss /= ne as f32; if step == 1 { first_loss = loss; } last_loss = loss; let d_out = Tensor::from_vec(dov, &[1, ll, dmodel], &dev).expect("dov"); let grads = block.backward(&x, &d_out).expect("backward"); let mut map = HashMap::new(); for (name, t) in block.persistence_tensors() { let p = t.to_vec().expect("p"); let g = grads.get(name).expect("g").to_vec().expect("gv"); let entry = state .entry(name.to_string()) .or_insert_with(|| (vec![0.0f32; p.len()], vec![0.0f32; p.len()])); let (m, v) = entry; let mut np = vec![0.0f32; p.len()]; for i in 0..p.len() { m[i] = b1 * m[i] + (1.0 - b1) * g[i]; v[i] = b2 * v[i] + (1.0 - b2) * g[i] * g[i]; let mhat = m[i] / (1.0 - b1.powi(step)); let vhat = v[i] / (1.0 - b2.powi(step)); np[i] = p[i] - lr * mhat / (vhat.sqrt() + eps); } map.insert( name.to_string(), Tensor::from_vec(np, t.shape().dims(), &dev).expect("np"), ); } block = MambaBlock::from_persistence_tensors(cfg.clone(), &dev, map).expect("rebuild"); } assert!( last_loss < 0.5 * first_loss, "training did not reduce loss enough: {first_loss:.5} → {last_loss:.5}" ); let cw_after = named_param(&block, "conv1d_weight"); assert!( max_abs_diff(&cw_before, &cw_after) > 1e-3, "backbone weight conv1d_weight did not move during training" ); } #[test] fn output_is_finite_and_non_constant() { let dev = Device::cpu(); let block = MambaBlock::new_seeded(cfg(), &dev, 1).expect("block"); let x = Tensor::randn_seeded(&[1, L, D_MODEL], &dev, 2).expect("x"); let y = fwd(&block, &x); assert!(y.iter().all(|v| v.is_finite()), "non-finite output"); let spread = max_abs_diff(&y, &vec![y[0]; y.len()]); assert!( spread > 1e-4, "output is suspiciously constant ({spread:.2e})" ); }