use super::{
ArdSharing, SaeManifoldRho, profiled_reml_block_efs_log_lambda_steps,
profiled_reml_block_log_lambda_gradient, profiled_reml_criterion,
};
use ndarray::{Array1, arr1};
#[test]
fn block_log_lambda_gradient_matches_central_difference() {
let n_obs = 50usize;
let p_x = 8usize;
let rss_x = 40.0_f64;
let block_rss = [30.0_f64, 12.0, 7.5];
let dims = [5usize, 3, 2];
let log_lambda = [0.2_f64, -0.5, 0.9];
let analytic =
profiled_reml_block_log_lambda_gradient(n_obs, p_x, rss_x, &block_rss, &dims, &log_lambda);
assert_eq!(analytic.len(), block_rss.len());
let h = 1e-6_f64;
for l in 0..block_rss.len() {
let mut plus = log_lambda;
let mut minus = log_lambda;
plus[l] += h;
minus[l] -= h;
let c_plus = profiled_reml_criterion(n_obs, p_x, rss_x, &block_rss, &dims, &plus);
let c_minus = profiled_reml_criterion(n_obs, p_x, rss_x, &block_rss, &dims, &minus);
let fd = (c_plus - c_minus) / (2.0 * h);
let tol = 1e-5 * (1.0 + analytic[l].abs());
assert!(
(analytic[l] - fd).abs() < tol,
"block {l}: analytic {} vs FD {} (diff {})",
analytic[l],
fd,
(analytic[l] - fd).abs()
);
}
}
#[test]
fn gradient_and_efs_step_vanish_at_variance_ratio_fixed_point() {
let n_obs = 64usize;
let p_x = 6usize;
let rss_x = 24.0_f64;
let block_rss = [50.0_f64, 9.0];
let dims = [4usize, 3];
let var_x = rss_x / p_x as f64;
let log_lambda: Vec<f64> = block_rss
.iter()
.zip(dims.iter())
.map(|(&rss, &dim)| (var_x / (rss / dim as f64)).ln())
.collect();
let grad =
profiled_reml_block_log_lambda_gradient(n_obs, p_x, rss_x, &block_rss, &dims, &log_lambda);
for (l, g) in grad.iter().enumerate() {
assert!(
g.abs() < 1e-9,
"block {l} gradient {g} not ~0 at fixed point"
);
}
let steps =
profiled_reml_block_efs_log_lambda_steps(p_x, rss_x, &block_rss, &dims, &log_lambda);
for (l, s) in steps.iter().enumerate() {
assert!(
s.abs() < 1e-12,
"block {l} EFS step {s} not ~0 at fixed point"
);
}
}
#[test]
fn one_efs_step_reaches_the_variance_ratio_root() {
let p_x = 10usize;
let rss_x = 55.0_f64;
let block_rss = [8.0_f64, 40.0];
let dims = [3usize, 5];
let log_lambda = [1.3_f64, -2.1];
let steps =
profiled_reml_block_efs_log_lambda_steps(p_x, rss_x, &block_rss, &dims, &log_lambda);
let var_x = rss_x / p_x as f64;
for l in 0..block_rss.len() {
let root = (var_x / (block_rss[l] / dims[l] as f64)).ln();
assert!(
(log_lambda[l] + steps[l] - root).abs() < 1e-12,
"block {l}: start+step {} != root {root}",
log_lambda[l] + steps[l]
);
}
}
#[test]
fn unidentifiable_block_is_held() {
let p_x = 4usize;
let rss_x = 12.0_f64;
let block_rss = [0.0_f64, 6.0]; let dims = [2usize, 3];
let log_lambda = [0.0_f64, 0.0];
let steps =
profiled_reml_block_efs_log_lambda_steps(p_x, rss_x, &block_rss, &dims, &log_lambda);
assert_eq!(steps[0], 0.0, "unidentifiable block must be held");
assert!(steps[1] != 0.0, "identifiable block must still move");
}
#[test]
fn rho_flat_round_trip_carries_block_coordinates() {
let log_ard = vec![arr1(&[0.1_f64]), arr1(&[0.2_f64])];
let plain = SaeManifoldRho::new(-0.3, 0.5, log_ard.clone());
let plain_len = plain.to_flat().len();
assert_eq!(plain_len, 1 + 2 + 2, "plain PerAtom flat length");
assert_eq!(plain.num_blocks(), 0);
let block = plain.clone().with_log_lambda_block(vec![0.7_f64, -0.4]);
let flat = block.to_flat();
assert_eq!(flat.len(), plain_len + 2, "block coords appended after ARD");
assert_eq!(flat[plain_len], 0.7);
assert_eq!(flat[plain_len + 1], -0.4);
let recovered = block.from_flat(flat.view());
assert_eq!(recovered.log_lambda_block, vec![0.7, -0.4]);
assert_eq!(recovered.log_lambda_sparse, -0.3);
assert_eq!(recovered.log_lambda_smooth, vec![0.5, 0.5]);
let shared =
SaeManifoldRho::new_shared_ard(-0.3, 0.5, log_ard).with_log_lambda_block(vec![1.1_f64]);
assert_eq!(shared.ard_sharing(), ArdSharing::Shared);
let sflat = shared.to_flat();
assert_eq!(sflat.len(), 1 + 2 + 1 + 1);
assert_eq!(sflat[sflat.len() - 1], 1.1);
let srecovered = shared.from_flat(sflat.view());
assert_eq!(srecovered.log_lambda_block, vec![1.1]);
}
#[test]
fn empty_block_flat_is_plain_sae_layout() {
let rho = SaeManifoldRho::new(0.0, 0.0, vec![arr1(&[0.0_f64])]);
let flat: Array1<f64> = rho.to_flat();
assert_eq!(flat.len(), 1 + 1 + 1);
let back = rho.from_flat(flat.view());
assert!(back.log_lambda_block.is_empty());
}