use super::prelude::*;
use gam_linalg::utils::{splitmix64, splitmix64_hash};
#[derive(Clone)]
pub struct DeflationSpec {
pub basis: Vec<Array1<f64>>,
}
#[derive(Clone)]
pub struct RationalLogdetPlan {
pub dim: usize,
pub probes: Vec<Array1<f64>>,
pub nodes: Vec<(f64, f64)>,
pub log_center: f64,
pub center: f64,
pub deflation: Option<DeflationSpec>,
}
pub struct RationalLogdetEval {
pub estimate: f64,
pub std_err: f64,
pub shifted_solves: Vec<Vec<Array1<f64>>>,
pub deflation_solves: Vec<Vec<Array1<f64>>>,
pub deflation_basis: Vec<Array1<f64>>,
pub cg_iterations: usize,
}
pub struct RationalLogdetDerivativeBundle {
pub vectors: Vec<Array1<f64>>,
}
impl RationalLogdetDerivativeBundle {
pub fn directional_derivative(
&self,
dmatvec: &(impl Fn(ArrayView1<f64>) -> Array1<f64> + Sync),
) -> Option<f64> {
if self.vectors.is_empty() {
return None;
}
let inv_rank = 1.0 / self.vectors.len() as f64;
let derivative = self
.vectors
.iter()
.map(|vector| vector.dot(&dmatvec(vector.view())))
.sum::<f64>()
* inv_rank;
derivative.is_finite().then_some(derivative)
}
}
impl RationalLogdetPlan {
pub fn build(
dim: usize,
num_probes: usize,
seed: u64,
lambda_min: f64,
lambda_max: f64,
rel_tol: f64,
) -> Option<Self> {
if dim == 0
|| num_probes == 0
|| !(lambda_min.is_finite() && lambda_max.is_finite())
|| lambda_min <= 0.0
|| lambda_max < lambda_min
|| !(rel_tol.is_finite() && rel_tol > 0.0 && rel_tol < 1.0)
{
return None;
}
let mut master = splitmix64_hash(seed);
let probes = rademacher_block(&mut master, num_probes, dim);
let c = (lambda_min * lambda_max).sqrt();
let t_lo = lambda_min * rel_tol;
let t_hi = lambda_max / rel_tol;
let u_of = |t: f64| ((2.0 / std::f64::consts::PI) * (t / c).ln()).asinh();
let u_lo = u_of(t_lo);
let u_hi = u_of(t_hi);
let pole_height = |lam_over_c: f64| -> f64 {
let s = (2.0 / std::f64::consts::PI) * lam_over_c.ln();
std::f64::consts::FRAC_PI_2 / (1.0 + s * s).sqrt()
};
let d_min = pole_height(lambda_min / c)
.min(pole_height(lambda_max / c))
.min(std::f64::consts::FRAC_PI_2);
let h_bound = 2.0 * std::f64::consts::PI * d_min / (1.0f64 / rel_tol).ln();
let steps = (((u_hi - u_lo) / h_bound).ceil() as usize).max(16);
let h = (u_hi - u_lo) / steps as f64;
let mut nodes = Vec::with_capacity(steps + 1);
for s in 0..=steps {
let u = u_lo + h * s as f64;
let t = c * (std::f64::consts::FRAC_PI_2 * u.sinh()).exp();
let w = h * t * std::f64::consts::FRAC_PI_2 * u.cosh();
if t.is_finite() && w.is_finite() && w > 0.0 {
nodes.push((t, w));
}
}
if nodes.is_empty() {
return None;
}
Some(Self {
dim,
probes,
nodes,
log_center: c.ln(),
center: c,
deflation: None,
})
}
pub fn with_deflation(
mut self,
matvec: &(impl Fn(ArrayView1<f64>) -> Array1<f64> + Sync),
rank: usize,
subspace_iters: usize,
seed: u64,
) -> Self {
let basis = build_deflation_basis(matvec, self.dim, rank, subspace_iters, seed);
self.deflation = (!basis.is_empty()).then_some(DeflationSpec { basis });
self
}
pub fn with_two_sided_deflation(
mut self,
matvec: &(impl Fn(ArrayView1<f64>) -> Array1<f64> + Sync),
top_rank: usize,
bottom_rank: usize,
subspace_iters: usize,
seed: u64,
cg: (f64, usize),
) -> Option<Self> {
let (cg_rel_tol, cg_max_iters) = cg;
let mut cols = build_deflation_basis(matvec, self.dim, top_rank, subspace_iters, seed);
cols.extend(build_inverse_deflation_basis(
matvec,
self.dim,
bottom_rank,
subspace_iters,
seed,
cg_rel_tol,
cg_max_iters,
)?);
let basis = orthonormalize(&cols);
self.deflation = (!basis.is_empty()).then_some(DeflationSpec { basis });
Some(self)
}
pub fn evaluate(
&self,
matvec: &(impl Fn(ArrayView1<f64>) -> Array1<f64> + Sync),
cg_rel_tol: f64,
cg_max_iters: usize,
) -> Option<RationalLogdetEval> {
let mut order: Vec<usize> = (0..self.nodes.len()).collect();
order.sort_by(|&a, &b| {
self.nodes[b]
.0
.partial_cmp(&self.nodes[a].0)
.unwrap_or(std::cmp::Ordering::Equal)
});
let basis: &[Array1<f64>] = self
.deflation
.as_ref()
.map(|d| d.basis.as_slice())
.unwrap_or(&[]);
let probes_proj = self.projected_probes(basis);
let (shifted, iters_probe) = solve_shift_ladder(
matvec,
&self.nodes,
&order,
&probes_proj,
cg_rel_tol,
cg_max_iters,
)?;
let (deflation_solves, iters_basis) = if basis.is_empty() {
(Vec::new(), 0)
} else {
solve_shift_ladder(matvec, &self.nodes, &order, basis, cg_rel_tol, cg_max_iters)?
};
self.assemble_eval(
probes_proj,
basis,
shifted,
deflation_solves,
iters_probe + iters_basis,
)
}
fn projected_probes(&self, basis: &[Array1<f64>]) -> Vec<Array1<f64>> {
self.probes
.iter()
.map(|v| {
let mut u = v.clone();
for q in basis {
let c = u.dot(q);
u.scaled_add(-c, q);
}
u
})
.collect()
}
fn assemble_eval(
&self,
probes_proj: Vec<Array1<f64>>,
basis: &[Array1<f64>],
shifted: Vec<Vec<Array1<f64>>>,
deflation_solves: Vec<Vec<Array1<f64>>>,
total_iters: usize,
) -> Option<RationalLogdetEval> {
let m = self.probes.len();
let k = self.dim as f64;
let mut term1 = 0.0_f64;
for (ell, &(t, w)) in self.nodes.iter().enumerate() {
let reference = 1.0 / (self.center + t);
for (i, q) in basis.iter().enumerate() {
term1 += w * (reference - q.dot(&deflation_solves[ell][i]));
}
}
let u_norm_sq: Vec<f64> = probes_proj.iter().map(|u| u.dot(u)).collect();
let mut per_probe = vec![0.0_f64; m];
for (ell, &(t, w)) in self.nodes.iter().enumerate() {
let inv = 1.0 / (self.center + t);
for j in 0..m {
per_probe[j] += w * (u_norm_sq[j] * inv - probes_proj[j].dot(&shifted[ell][j]));
}
}
let term2 = per_probe.iter().sum::<f64>() / m as f64;
let estimate = k * self.log_center + term1 + term2;
let std_err = if m > 1 {
let var = per_probe
.iter()
.map(|e| (e - term2) * (e - term2))
.sum::<f64>()
/ (m as f64 - 1.0);
(var / m as f64).sqrt()
} else {
0.0
};
if !(estimate.is_finite() && std_err.is_finite()) {
return None;
}
Some(RationalLogdetEval {
estimate,
std_err,
shifted_solves: shifted,
deflation_solves,
deflation_basis: basis.to_vec(),
cg_iterations: total_iters,
})
}
pub fn directional_derivative(
&self,
eval: &RationalLogdetEval,
dmatvec: &(impl Fn(ArrayView1<f64>) -> Array1<f64> + Sync),
) -> Option<f64> {
let m = self.probes.len() as f64;
let mut acc_probe = 0.0;
let mut acc_defl = 0.0;
for (ell, &(_, w)) in self.nodes.iter().enumerate() {
for y in &eval.shifted_solves[ell] {
let dy = dmatvec(y.view());
acc_probe += w * y.dot(&dy);
}
if let Some(defl) = eval.deflation_solves.get(ell) {
for y in defl {
let dy = dmatvec(y.view());
acc_defl += w * y.dot(&dy);
}
}
}
let acc = acc_defl + acc_probe / m;
acc.is_finite().then_some(acc)
}
pub fn into_directional_derivative_bundle(
&self,
eval: RationalLogdetEval,
) -> Option<RationalLogdetDerivativeBundle> {
let expected_deflation_nodes =
usize::from(!eval.deflation_basis.is_empty()) * self.nodes.len();
if eval.shifted_solves.len() != self.nodes.len()
|| eval.deflation_solves.len() != expected_deflation_nodes
{
return None;
}
let probe_count = self.probes.len();
if probe_count == 0
|| eval
.shifted_solves
.iter()
.any(|solves| solves.len() != probe_count)
|| eval
.deflation_solves
.iter()
.any(|solves| solves.len() != eval.deflation_basis.len())
{
return None;
}
let term_count = self.nodes.len().checked_mul(
probe_count.checked_add(eval.deflation_basis.len())?,
)?;
if term_count == 0 {
return None;
}
let mut vectors = Vec::with_capacity(term_count);
let rank = term_count as f64;
let probes = probe_count as f64;
let mut deflation_by_node = eval.deflation_solves;
if deflation_by_node.is_empty() {
deflation_by_node.resize_with(self.nodes.len(), Vec::new);
}
for ((mut probe_solves, mut deflation_solves), &(_, weight)) in eval
.shifted_solves
.into_iter()
.zip(deflation_by_node)
.zip(&self.nodes)
{
if !(weight.is_finite() && weight > 0.0) {
return None;
}
let probe_scale = (rank * weight / probes).sqrt();
let deflation_scale = (rank * weight).sqrt();
if !(probe_scale.is_finite() && deflation_scale.is_finite()) {
return None;
}
for mut solve in probe_solves.drain(..) {
if solve.len() != self.dim {
return None;
}
solve *= probe_scale;
vectors.push(solve);
}
for mut solve in deflation_solves.drain(..) {
if solve.len() != self.dim {
return None;
}
solve *= deflation_scale;
vectors.push(solve);
}
}
Some(RationalLogdetDerivativeBundle { vectors })
}
}
fn shifted_cg(
matvec: &(impl Fn(ArrayView1<f64>) -> Array1<f64> + Sync),
t: f64,
b: &Array1<f64>,
y0: &Array1<f64>,
rel_tol: f64,
max_iters: usize,
) -> Option<(Array1<f64>, usize)> {
if !(rel_tol.is_finite() && rel_tol > 0.0) {
return None;
}
let apply = |v: ArrayView1<f64>| -> Array1<f64> {
let mut out = matvec(v);
out.scaled_add(t, &v.to_owned());
out
};
let mut y = y0.clone();
let mut r = b - &apply(y.view());
let b_norm = b.dot(b).sqrt().max(f64::MIN_POSITIVE);
let mut p = r.clone();
let mut rs = r.dot(&r);
if !rs.is_finite() {
return None;
}
let tol = rel_tol * b_norm;
let mut iters = 0usize;
let mut observed_operator_norm = 0.0_f64;
loop {
if rs.sqrt() <= tol {
let true_residual = b - &apply(y.view());
let true_rs = true_residual.dot(&true_residual);
if !true_rs.is_finite() {
return None;
}
let true_residual_norm = true_rs.sqrt();
let y_norm = y.dot(&y).sqrt();
if !y_norm.is_finite() {
return None;
}
let backward_error_certified =
if observed_operator_norm > 0.0 && y_norm > 0.0 {
let log_operator_solution = observed_operator_norm.ln() + y_norm.ln();
let log_rhs = b_norm.ln();
let log_scale = log_operator_solution.max(log_rhs);
let log_denominator = log_scale
+ ((log_operator_solution - log_scale).exp()
+ (log_rhs - log_scale).exp())
.ln();
true_residual_norm.ln() - log_denominator <= rel_tol.ln()
} else {
false
};
if true_residual_norm <= tol || backward_error_certified {
return Some((y, iters));
}
if iters >= max_iters {
return None;
}
r = true_residual;
rs = true_rs;
p = r.clone();
}
if iters >= max_iters {
return None;
}
let ap = apply(p.view());
let denom = p.dot(&ap);
if !(denom.is_finite() && denom > 0.0) {
return None;
}
let p_norm_sq = p.dot(&p);
if !(p_norm_sq.is_finite() && p_norm_sq > 0.0) {
return None;
}
let rayleigh = denom / p_norm_sq;
if rayleigh.is_finite() {
observed_operator_norm = observed_operator_norm.max(rayleigh);
}
let alpha = rs / denom;
y.scaled_add(alpha, &p);
r.scaled_add(-alpha, &ap);
let rs_new = r.dot(&r);
if !rs_new.is_finite() {
return None;
}
p = &r + &(&p * (rs_new / rs));
rs = rs_new;
iters += 1;
}
}
fn orthonormalize(cols: &[Array1<f64>]) -> Vec<Array1<f64>> {
let mut out: Vec<Array1<f64>> = Vec::with_capacity(cols.len());
for col in cols {
let mut v = col.clone();
let v0_norm = v.dot(&v).sqrt();
for basis in &out {
let proj = v.dot(basis);
v.scaled_add(-proj, basis);
}
let norm_after_first = v.dot(&v).sqrt();
for basis in &out {
let proj = v.dot(basis);
v.scaled_add(-proj, basis);
}
let norm = v.dot(&v).sqrt();
let rank_tol = f64::EPSILON.sqrt() * v0_norm;
let collapsed = !(v0_norm.is_finite() && v0_norm > 0.0)
|| !(norm_after_first.is_finite())
|| norm_after_first <= rank_tol
|| !(norm.is_finite())
|| norm <= rank_tol;
if !collapsed {
v.mapv_inplace(|x| x / norm);
out.push(v);
}
}
out
}
fn rademacher_block(master: &mut u64, ncols: usize, dim: usize) -> Vec<Array1<f64>> {
(0..ncols)
.map(|_| {
let mut v = Array1::<f64>::zeros(dim);
let mut bits: u64 = 0;
let mut remaining: u32 = 0;
for value in v.iter_mut() {
if remaining == 0 {
bits = splitmix64(master);
remaining = 64;
}
*value = if bits & 1 == 1 { 1.0 } else { -1.0 };
bits >>= 1;
remaining -= 1;
}
v
})
.collect()
}
fn build_deflation_basis(
matvec: &(impl Fn(ArrayView1<f64>) -> Array1<f64> + Sync),
dim: usize,
rank: usize,
iters: usize,
seed: u64,
) -> Vec<Array1<f64>> {
let r = rank.min(dim);
if r == 0 {
return Vec::new();
}
let mut master = splitmix64_hash(seed.wrapping_add(0xD1B5_4A32_D192_ED03));
let mut cols = orthonormalize(&rademacher_block(&mut master, r, dim));
for _ in 0..iters {
if cols.is_empty() {
break;
}
let applied: Vec<Array1<f64>> = cols.iter().map(|c| matvec(c.view())).collect();
cols = orthonormalize(&applied);
}
cols
}
fn build_inverse_deflation_basis(
matvec: &(impl Fn(ArrayView1<f64>) -> Array1<f64> + Sync),
dim: usize,
rank: usize,
iters: usize,
seed: u64,
cg_rel_tol: f64,
cg_max_iters: usize,
) -> Option<Vec<Array1<f64>>> {
let r = rank.min(dim);
if r == 0 {
return Some(Vec::new());
}
let mut master = splitmix64_hash(seed.wrapping_add(0x2545_F491_4F6C_DD1D));
let mut cols = orthonormalize(&rademacher_block(&mut master, r, dim));
let zero = Array1::<f64>::zeros(dim);
for _ in 0..iters {
if cols.is_empty() {
break;
}
let applied: Option<Vec<Array1<f64>>> = cols
.iter()
.map(|c| shifted_cg(matvec, 0.0, c, &zero, cg_rel_tol, cg_max_iters).map(|(y, _)| y))
.collect();
cols = orthonormalize(&applied?);
}
Some(cols)
}
fn solve_shift_ladder(
matvec: &(impl Fn(ArrayView1<f64>) -> Array1<f64> + Sync),
nodes: &[(f64, f64)],
order: &[usize],
vectors: &[Array1<f64>],
cg_rel_tol: f64,
cg_max_iters: usize,
) -> Option<(Vec<Vec<Array1<f64>>>, usize)> {
let m = vectors.len();
let dim = vectors.first().map(|v| v.len()).unwrap_or(0);
let mut solves: Vec<Vec<Array1<f64>>> = vec![Vec::with_capacity(m); nodes.len()];
let mut warm: Vec<Array1<f64>> = vec![Array1::zeros(dim); m];
let mut total = 0usize;
for &ell in order {
let (t, _) = nodes[ell];
let mut per = Vec::with_capacity(m);
for (j, v) in vectors.iter().enumerate() {
let (y, iters) = shifted_cg(matvec, t, v, &warm[j], cg_rel_tol, cg_max_iters)?;
total += iters;
warm[j] = y.clone();
per.push(y);
}
solves[ell] = per;
}
Some((solves, total))
}
#[cfg(test)]
mod tests {
use super::*;
use ndarray::array;
fn next_uniform(state: &mut u64, lo: f64, hi: f64) -> f64 {
let bits = splitmix64(state) >> 11;
let unit = (bits as f64) / ((1u64 << 53) as f64);
lo + (hi - lo) * unit
}
fn spd_with_spectrum(dim: usize, lambdas: &[f64], seed: u64) -> (Array2<f64>, f64) {
let mut state = seed;
let mut g = Array2::<f64>::zeros((dim, dim));
for v in g.iter_mut() {
let u1 = next_uniform(&mut state, 1e-12, 1.0);
let u2 = next_uniform(&mut state, 0.0, 1.0);
*v = (-2.0 * u1.ln()).sqrt() * (2.0 * std::f64::consts::PI * u2).cos();
}
let mut q = Array2::<f64>::zeros((dim, dim));
for c in 0..dim {
let mut col = g.column(c).to_owned();
for prev in 0..c {
let proj = q.column(prev).dot(&col);
let prev_col = q.column(prev).to_owned();
col.scaled_add(-proj, &prev_col);
}
let norm = col.dot(&col).sqrt();
let col = col / norm;
q.column_mut(c).assign(&col);
}
let mut a = Array2::<f64>::zeros((dim, dim));
for (i, &l) in lambdas.iter().enumerate() {
let qi = q.column(i);
for r in 0..dim {
for c in 0..dim {
a[[r, c]] += l * qi[r] * qi[c];
}
}
}
let logdet: f64 = lambdas.iter().map(|l| l.ln()).sum();
(a, logdet)
}
#[test]
fn quadrature_is_exact_on_scalar_spectrum() {
for &x in &[1e-6, 1e-3, 0.5, 1.0, 7.3, 1e4, 1e8] {
let plan = RationalLogdetPlan::build(1, 1, 7, x, x, 1e-10).expect("plan");
let a = Array2::from_elem((1, 1), x);
let eval = plan
.evaluate(&|v: ArrayView1<f64>| a.dot(&v), 1e-14, 10_000)
.expect("eval");
let err = (eval.estimate - x.ln()).abs() / x.ln().abs().max(1.0);
assert!(
err < 1e-8,
"quadrature error {err:.3e} at x={x:e} (est {} vs {})",
eval.estimate,
x.ln()
);
}
}
#[test]
fn matches_dense_logdet_within_probe_error_at_wide_kappa() {
let dim = 96;
let mut state = 42u64;
let lambdas: Vec<f64> = (0..dim)
.map(|_| 10f64.powf(next_uniform(&mut state, -4.0, 4.0)))
.collect();
let (a, logdet) = spd_with_spectrum(dim, &lambdas, 1234);
let lmin = lambdas.iter().cloned().fold(f64::INFINITY, f64::min);
let lmax = lambdas.iter().cloned().fold(0.0f64, f64::max);
let plan = RationalLogdetPlan::build(dim, 64, 11, lmin, lmax, 1e-9).expect("plan");
let eval = plan
.evaluate(&|v: ArrayView1<f64>| a.dot(&v), 1e-12, 50_000)
.expect("eval");
let err = (eval.estimate - logdet).abs();
let budget = 5.0 * eval.std_err + 1e-3 * logdet.abs().max(1.0);
assert!(
err < budget,
"estimate {} vs exact {} — |err| {err:.3e} exceeds 5σ+quad budget {budget:.3e} \
(std_err {:.3e})",
eval.estimate,
logdet,
eval.std_err
);
assert!(
eval.std_err.is_finite() && eval.std_err > 0.0,
"multi-probe eval must report a positive error bar"
);
}
#[test]
fn directional_derivative_matches_fd_of_the_surrogate_itself() {
let dim = 40;
let mut state = 9u64;
let lambdas: Vec<f64> = (0..dim)
.map(|_| 10f64.powf(next_uniform(&mut state, -2.0, 2.0)))
.collect();
let (a, _) = spd_with_spectrum(dim, &lambdas, 77);
let d_lambdas: Vec<f64> = (0..dim)
.map(|_| next_uniform(&mut state, 0.1, 1.0))
.collect();
let (da, _) = spd_with_spectrum(dim, &d_lambdas, 78);
let plan = RationalLogdetPlan::build(dim, 8, 5, 1e-2, 1e2, 1e-9).expect("plan");
let eval = plan
.evaluate(&|v: ArrayView1<f64>| a.dot(&v), 1e-13, 20_000)
.expect("eval");
let grad = plan
.directional_derivative(&eval, &|v: ArrayView1<f64>| da.dot(&v))
.expect("grad");
let h = 1e-5;
let a_plus = &a + &(&da * h);
let a_minus = &a - &(&da * h);
let f_plus = plan
.evaluate(&|v: ArrayView1<f64>| a_plus.dot(&v), 1e-13, 20_000)
.expect("eval+")
.estimate;
let f_minus = plan
.evaluate(&|v: ArrayView1<f64>| a_minus.dot(&v), 1e-13, 20_000)
.expect("eval-")
.estimate;
let fd = (f_plus - f_minus) / (2.0 * h);
let rel = (grad - fd).abs() / fd.abs().max(1e-12);
assert!(
rel < 1e-5,
"surrogate gradient {grad:.9e} vs its own FD {fd:.9e} (rel {rel:.3e})"
);
assert!(
grad > 0.0,
"SPD direction must increase log det, got {grad}"
);
}
#[test]
fn fixed_probe_derivative_bundle_matches_rational_directional_not_raw_inverse() {
let diagonal = array![0.2, 3.0, 17.0];
let direction = array![1.0, 2.0, 4.0];
let matvec = |v: ArrayView1<f64>| &diagonal * &v;
let dmatvec = |v: ArrayView1<f64>| &direction * &v;
let plan = RationalLogdetPlan::build(3, 3, 71, 0.2, 17.0, 0.25)
.expect("fixed rational plan");
let eval = plan
.evaluate(&matvec, 1.0e-13, 64)
.expect("fixed rational evaluation");
let authority = plan
.directional_derivative(&eval, &dmatvec)
.expect("rational directional derivative");
let bundle = plan
.into_directional_derivative_bundle(eval)
.expect("lossless rational derivative bundle");
let represented = bundle
.directional_derivative(&dmatvec)
.expect("represented directional derivative");
let scale = authority.abs().max(1.0);
assert!(
(represented - authority).abs() <= 64.0 * f64::EPSILON * scale,
"lossless bundle derivative {represented:.16e} != rational authority \
{authority:.16e}"
);
let raw_shift_zero = direction
.iter()
.zip(diagonal.iter())
.map(|(&d, &s)| d / s)
.sum::<f64>();
assert!(
(raw_shift_zero - authority).abs() > 1.0e-4,
"fixture must separate the rational derivative ({authority:.9e}) from the \
raw shift-zero inverse trace ({raw_shift_zero:.9e})"
);
}
#[test]
fn evaluate_is_deterministic_across_calls() {
let dim = 24;
let lambdas: Vec<f64> = (1..=dim).map(|i| i as f64).collect();
let (a, _) = spd_with_spectrum(dim, &lambdas, 3);
let plan = RationalLogdetPlan::build(dim, 4, 99, 1.0, dim as f64, 1e-8).expect("plan");
let e1 = plan
.evaluate(&|v: ArrayView1<f64>| a.dot(&v), 1e-12, 10_000)
.expect("eval1")
.estimate;
let e2 = plan
.evaluate(&|v: ArrayView1<f64>| a.dot(&v), 1e-12, 10_000)
.expect("eval2")
.estimate;
assert_eq!(e1, e2, "fixed plan must be bit-deterministic");
}
#[test]
fn shifted_cg_refuses_an_unconverged_iteration_cap() {
let a = array![[1.0, 0.0], [0.0, 4.0]];
let b = array![1.0, 1.0];
let zero = Array1::<f64>::zeros(2);
let matvec = |v: ArrayView1<f64>| a.dot(&v);
assert!(
shifted_cg(&matvec, 0.0, &b, &zero, 1.0e-12, 1).is_none(),
"one CG step cannot solve a two-eigenvalue system to 1e-12; the \
iteration-capped last iterate must be refused"
);
let (solved, iterations) = shifted_cg(&matvec, 0.0, &b, &zero, 1.0e-12, 2)
.expect("two-dimensional SPD CG must converge in at most two steps");
let residual = &b - &matvec(solved.view());
assert!(
residual.dot(&residual).sqrt() <= 1.0e-12 * b.dot(&b).sqrt(),
"returned shifted solve must satisfy its true-residual contract"
);
assert_eq!(iterations, 2);
}
#[test]
fn two_sided_deflation_propagates_bottom_solve_nonconvergence() {
let a = array![[1.0, 0.0], [0.0, 4.0]];
let matvec = |v: ArrayView1<f64>| a.dot(&v);
let plan =
RationalLogdetPlan::build(2, 2, 17, 1.0, 4.0, 1.0e-9).expect("valid rational plan");
assert!(
plan.with_two_sided_deflation(&matvec, 0, 1, 1, 91, (1.0e-12, 0))
.is_none(),
"a requested bottom-tail inverse solve may not silently fall back to \
the unamplified start column"
);
}
#[test]
fn full_rank_deflation_is_exact_no_hutchinson() {
let dim = 28;
let lambdas: Vec<f64> = (1..=dim).map(|i| 0.3 + 0.7 * i as f64).collect();
let (a, logdet) = spd_with_spectrum(dim, &lambdas, 31);
let lmin = lambdas.iter().cloned().fold(f64::INFINITY, f64::min);
let lmax = lambdas.iter().cloned().fold(0.0f64, f64::max);
let matvec = |v: ArrayView1<f64>| a.dot(&v);
let plan = RationalLogdetPlan::build(dim, 4, 5, lmin, lmax, 1e-11)
.expect("plan")
.with_deflation(&matvec, dim, 2, 123);
let eval = plan.evaluate(&matvec, 1e-14, 20_000).expect("eval");
assert_eq!(
eval.deflation_basis.len(),
dim,
"full-rank block must realise dim orthonormal columns"
);
assert!(
eval.std_err < 1e-8,
"full deflation leaves ~no Hutchinson variance (P ≈ 0), got std_err={:.3e}",
eval.std_err
);
let rel = (eval.estimate - logdet).abs() / logdet.abs().max(1.0);
assert!(
rel < 1e-6,
"full-rank deflation must be exact to quadrature: rel {rel:.3e} \
(est {} vs {logdet})",
eval.estimate
);
}
#[test]
fn full_rank_deflation_is_exact_at_wide_kappa_deterministic_bias_localizer() {
let dim = 96;
let mut state = 2026u64;
let lambdas: Vec<f64> = (0..dim)
.map(|_| 10f64.powf(next_uniform(&mut state, -4.0, 4.0)))
.collect();
let (a, logdet) = spd_with_spectrum(dim, &lambdas, 4321);
let lmin = lambdas.iter().cloned().fold(f64::INFINITY, f64::min);
let lmax = lambdas.iter().cloned().fold(0.0f64, f64::max);
let matvec = |v: ArrayView1<f64>| a.dot(&v);
let plan = RationalLogdetPlan::build(dim, 8, 5, lmin, lmax, 1e-9)
.expect("plan")
.with_deflation(&matvec, dim, 2, 555);
let eval = evaluate_exact(&plan, &a);
assert_eq!(
eval.deflation_basis.len(),
dim,
"the rank=dim block power must realise a full orthonormal basis even at \
κ≈1e8 (got {}); if it collapses, term2 is nonzero and this stops being a \
zero-variance deterministic check",
eval.deflation_basis.len()
);
assert!(
eval.std_err < 1e-8,
"full deflation must leave ~no Hutchinson variance (P ≈ 0) at wide κ, got \
std_err={:.3e}",
eval.std_err
);
let rel = (eval.estimate - logdet).abs() / logdet.abs().max(1.0);
assert!(
rel < 1e-6,
"wide-κ full-rank deflation must be exact to quadrature — a nonzero value \
is the ONLY signature of a genuine deterministic quadrature/split bias: \
rel {rel:.3e} (est {} vs exact {logdet})",
eval.estimate
);
}
#[test]
fn deflation_cuts_error_bar_and_stays_accurate_at_wide_kappa() {
let dim = 96;
let mut state = 2026u64;
let lambdas: Vec<f64> = (0..dim)
.map(|_| 10f64.powf(next_uniform(&mut state, -4.0, 4.0)))
.collect();
let (a, logdet) = spd_with_spectrum(dim, &lambdas, 4321);
let lmin = lambdas.iter().cloned().fold(f64::INFINITY, f64::min);
let lmax = lambdas.iter().cloned().fold(0.0f64, f64::max);
let matvec = |v: ArrayView1<f64>| a.dot(&v);
let plain = RationalLogdetPlan::build(dim, 32, 17, lmin, lmax, 1e-9).expect("plan");
let defl = plain.clone().with_deflation(&matvec, 16, 3, 555);
let e_plain = plain.evaluate(&matvec, 1e-12, 50_000).expect("plain");
let e_defl = defl.evaluate(&matvec, 1e-12, 50_000).expect("defl");
let rel = (e_defl.estimate - logdet).abs() / logdet.abs().max(1.0);
eprintln!(
"wide-κ: plain std_err={:.3e} defl std_err={:.3e} defl rel={:.3e}",
e_plain.std_err, e_defl.std_err, rel
);
assert!(
rel < 0.05,
"deflated estimate rel err {rel:.3e} (est {} vs exact {logdet})",
e_defl.estimate
);
assert!(
e_defl.std_err < e_plain.std_err,
"deflation must shrink the Hutchinson error bar (plain {:.3e} vs defl {:.3e})",
e_plain.std_err,
e_defl.std_err
);
}
fn evaluate_exact(plan: &RationalLogdetPlan, a: &Array2<f64>) -> RationalLogdetEval {
use gam_linalg::triangular::{
CholeskyGuard, cholesky_factor_in_place, cholesky_solve_vector,
};
let basis: &[Array1<f64>] = plan
.deflation
.as_ref()
.map(|d| d.basis.as_slice())
.unwrap_or(&[]);
let probes_proj = plan.projected_probes(basis);
let exact_ladder = |vectors: &[Array1<f64>]| -> Vec<Vec<Array1<f64>>> {
plan.nodes
.iter()
.map(|&(t, _)| {
let mut at = a.clone();
for i in 0..a.nrows() {
at[[i, i]] += t;
}
let l = cholesky_factor_in_place(at.view(), CholeskyGuard::FiniteStrict)
.expect("shifted SPD system must factor");
vectors
.iter()
.map(|v| cholesky_solve_vector(l.view(), v.view()))
.collect()
})
.collect()
};
let shifted = exact_ladder(&probes_proj);
let deflation_solves = if basis.is_empty() {
Vec::new()
} else {
exact_ladder(basis)
};
plan.assemble_eval(probes_proj, basis, shifted, deflation_solves, 0)
.expect("exact assemble")
}
#[test]
fn deflation_wide_kappa_bias_cg_convergence_discriminator() {
let dim = 96;
let mut state = 2026u64;
let lambdas: Vec<f64> = (0..dim)
.map(|_| 10f64.powf(next_uniform(&mut state, -4.0, 4.0)))
.collect();
let (a, logdet) = spd_with_spectrum(dim, &lambdas, 4321);
let lmin = lambdas.iter().cloned().fold(f64::INFINITY, f64::min);
let lmax = lambdas.iter().cloned().fold(0.0f64, f64::max);
let matvec = |v: ArrayView1<f64>| a.dot(&v);
let plain = RationalLogdetPlan::build(dim, 32, 17, lmin, lmax, 1e-9).expect("plan");
let defl = plain.clone().with_deflation(&matvec, 16, 3, 555);
let e_loose = defl.evaluate(&matvec, 1e-12, 50_000).expect("loose");
let e_tight = defl.evaluate(&matvec, 1e-15, 500_000).expect("tight");
let e_exact = evaluate_exact(&defl, &a);
let rel_loose = (e_loose.estimate - logdet).abs() / logdet.abs().max(1.0);
let rel_tight = (e_tight.estimate - logdet).abs() / logdet.abs().max(1.0);
let rel_exact = (e_exact.estimate - logdet).abs() / logdet.abs().max(1.0);
let gap = logdet - e_loose.estimate;
let sigma_ratio = gap.abs() / e_loose.std_err.max(1e-300);
eprintln!(
"wide-κ 3-arm discriminator: exact_logdet={logdet:.6}\n \
loose(1e-12,50k) est={:.6} rel={rel_loose:.3e} std_err={:.3e}\n \
tight(1e-15,500k) est={:.6} rel={rel_tight:.3e} Δvs_loose={:.3e}\n \
EXACT(cholesky) est={:.6} rel={rel_exact:.3e} Δvs_loose={:.3e}\n \
gap={gap:.4} = {sigma_ratio:.1}σ (loose std_err)",
e_loose.estimate,
e_loose.std_err,
e_tight.estimate,
(e_tight.estimate - e_loose.estimate).abs(),
e_exact.estimate,
(e_exact.estimate - e_loose.estimate).abs(),
);
let frozen: &[Array1<f64>] = defl
.deflation
.as_ref()
.map(|d| d.basis.as_slice())
.unwrap_or(&[]);
assert_eq!(
e_exact.deflation_basis.len(),
frozen.len(),
"exact arm must realise the same frozen Q rank as the plan"
);
for (qe, qf) in e_exact.deflation_basis.iter().zip(frozen) {
assert!(
(qe - qf).mapv(f64::abs).sum() < 1e-12,
"term1's Q must be the plan's frozen Q (no drift)"
);
}
for (i, qi) in frozen.iter().enumerate() {
for (j, qj) in frozen.iter().enumerate() {
let expect = if i == j { 1.0 } else { 0.0 };
assert!(
(qi.dot(qj) - expect).abs() < 1e-9,
"frozen Q must be orthonormal: QᵀQ[{i},{j}] = {}",
qi.dot(qj)
);
}
}
let proj = defl.projected_probes(frozen);
for u in &proj {
for q in frozen {
assert!(
u.dot(q).abs() < 1e-9,
"projected probe must be Q-orthogonal (same P as term1)"
);
}
}
assert!(
rel_exact < 0.05,
"EXACT-solve deflated estimate rel err {rel_exact:.3e} (est {} vs exact {logdet}); \
CG loose rel {rel_loose:.3e}, tight rel {rel_tight:.3e}. With no CG error possible, \
rel_exact ≈ rel_loose ⇒ the wide-κ bias is STRUCTURAL in the deflated split, not \
solve convergence",
e_exact.estimate
);
}
#[test]
fn deflation_wide_kappa_variance_vs_bias_multiseed() {
let dim = 96;
let mut state = 2026u64;
let lambdas: Vec<f64> = (0..dim)
.map(|_| 10f64.powf(next_uniform(&mut state, -4.0, 4.0)))
.collect();
let (a, logdet) = spd_with_spectrum(dim, &lambdas, 4321);
let lmin = lambdas.iter().cloned().fold(f64::INFINITY, f64::min);
let lmax = lambdas.iter().cloned().fold(0.0f64, f64::max);
let matvec = |v: ArrayView1<f64>| a.dot(&v);
let k_seeds = 96usize;
let mut ests = Vec::with_capacity(k_seeds);
let mut internal_bars = Vec::with_capacity(k_seeds);
let mut probe_fingerprints: std::collections::HashSet<u128> =
std::collections::HashSet::new();
for s in 0..k_seeds {
let plan = RationalLogdetPlan::build(dim, 32, 9000 + s as u64, lmin, lmax, 1e-9)
.expect("plan")
.with_deflation(&matvec, 16, 3, 555);
for probe in &plan.probes {
let mut fp = 0u128;
for (i, &x) in probe.iter().enumerate() {
if x > 0.0 {
fp |= 1u128 << i;
}
}
probe_fingerprints.insert(fp);
}
let e = evaluate_exact(&plan, &a);
ests.push(e.estimate);
internal_bars.push(e.std_err);
}
let total_pairs = k_seeds * 32;
let distinct = probe_fingerprints.len();
assert!(
distinct as f64 > 0.95 * total_pairs as f64,
"probe vectors must be jointly independent across seeds for this \
variance-vs-bias split to be valid: only {distinct} distinct of \
{total_pairs} (seed, probe) pairs — the RNG has re-aliased unit-spaced \
seeds (expected ~{total_pairs}), so any reported σ is meaningless"
);
let n = ests.len() as f64;
let mean = ests.iter().sum::<f64>() / n;
let var = ests.iter().map(|e| (e - mean).powi(2)).sum::<f64>() / (n - 1.0);
let sd = var.sqrt();
let se_mean = sd / n.sqrt();
let mean_internal_bar = internal_bars.iter().sum::<f64>() / n;
let bias = mean - logdet;
let bias_frac = bias.abs() / logdet.abs().max(1.0);
let bias_sigma = bias.abs() / se_mean.max(1e-300);
eprintln!(
"wide-κ variance-vs-bias ({k_seeds} seeds, fixed Q, EXACT solves): exact={logdet:.6}\n \
mean={mean:.6} bias={bias:+.6} ({bias_frac:.3e} rel, {bias_sigma:.2}σ of the mean)\n \
seed-to-seed sd={sd:.4} se_mean={se_mean:.4} ⟨internal std_err⟩={mean_internal_bar:.4}\n \
VERDICT: {}",
if bias_sigma < 3.0 {
"VARIANCE-dominated — split+quadrature UNBIASED at κ=1e8; fix = variance reduction (probes/rank/control-variate), NOT re-derivation"
} else {
"genuine DETERMINISTIC bias survives probe-averaging — quadrature/split derivation work needed"
}
);
assert!(
(mean_internal_bar / sd).ln().abs() < 1.0,
"internal std_err ({mean_internal_bar:.3}) must track the empirical seed spread ({sd:.3}) \
within a factor e; a mismatch means the surrogate's error bar is miscalibrated"
);
assert!(
bias_sigma < 3.0 || bias_frac < 0.02,
"probe-averaged estimate is biased by {bias:+.4} ({bias_frac:.3e} rel, {bias_sigma:.2}σ): \
deterministic split/quadrature bias survives — genuine derivation work, not variance"
);
}
#[test]
fn two_sided_deflation_drops_wide_kappa_std_err_below_two_percent() {
let dim = 96;
let mut state = 2026u64;
let lambdas: Vec<f64> = (0..dim)
.map(|_| 10f64.powf(next_uniform(&mut state, -4.0, 4.0)))
.collect();
let (a, logdet) = spd_with_spectrum(dim, &lambdas, 4321);
let lmin = lambdas.iter().cloned().fold(f64::INFINITY, f64::min);
let lmax = lambdas.iter().cloned().fold(0.0f64, f64::max);
let matvec = |v: ArrayView1<f64>| a.dot(&v);
let base = RationalLogdetPlan::build(dim, 256, 17, lmin, lmax, 1e-9).expect("plan");
let top16 = base.clone().with_deflation(&matvec, 16, 3, 555); let top64 = base.clone().with_deflation(&matvec, 64, 3, 555); let two = base
.clone()
.with_two_sided_deflation(&matvec, 32, 32, 3, 555, (1e-3, 5000))
.expect("bottom-tail inverse iteration must converge");
let e16 = evaluate_exact(&top16, &a);
let e64 = evaluate_exact(&top64, &a);
let e2 = evaluate_exact(&two, &a);
let f = |se: f64| se / logdet.abs().max(1.0);
eprintln!(
"wide-κ variance reduction (256 probes, EXACT estimator): |logdet|={:.4}\n \
top-only r16 (current): std_err={:.4} ({:.4} of |ld|)\n \
top-only r64 (eq-rank): std_err={:.4} ({:.4} of |ld|)\n \
two-sided 32+32: std_err={:.4} ({:.4} of |ld|) rel={:.4}\n \
=> vs top-r16 {:.2}×, vs eq-rank top-r64 {:.2}×",
logdet.abs(),
e16.std_err,
f(e16.std_err),
e64.std_err,
f(e64.std_err),
e2.std_err,
f(e2.std_err),
(e2.estimate - logdet).abs() / logdet.abs().max(1.0),
e16.std_err / e2.std_err.max(1e-300),
e64.std_err / e2.std_err.max(1e-300),
);
assert_eq!(
e2.deflation_basis.len(),
64,
"two-sided block must realise 32 top + 32 bottom orthonormal columns (got {})",
e2.deflation_basis.len()
);
assert!(
e2.std_err < 0.02 * logdet.abs(),
"two-sided wide-κ std_err {:.4} must fall below 2% of |logdet| ({:.4})",
e2.std_err,
0.02 * logdet.abs()
);
assert!(
e2.std_err < 0.5 * e64.std_err,
"two-sided ({:.4}) must beat equal-rank one-sided ({:.4}) by ≥2× — the win is \
peeling BOTH tails, not merely deflating more columns",
e2.std_err,
e64.std_err
);
assert!(
(e2.estimate - logdet).abs() < 5.0 * e2.std_err,
"two-sided estimate {:.4} must stay within 5σ ({:.4}) of exact {:.4} — variance \
reduction must not bias the value",
e2.estimate,
5.0 * e2.std_err,
logdet
);
}
#[test]
fn deflated_directional_derivative_matches_fd_of_surrogate() {
let dim = 40;
let mut state = 9u64;
let lambdas: Vec<f64> = (0..dim)
.map(|_| 10f64.powf(next_uniform(&mut state, -2.0, 2.0)))
.collect();
let (a, _) = spd_with_spectrum(dim, &lambdas, 77);
let d_lambdas: Vec<f64> = (0..dim)
.map(|_| next_uniform(&mut state, 0.1, 1.0))
.collect();
let (da, _) = spd_with_spectrum(dim, &d_lambdas, 78);
let matvec = |v: ArrayView1<f64>| a.dot(&v);
let plan = RationalLogdetPlan::build(dim, 8, 5, 1e-2, 1e2, 1e-9)
.expect("plan")
.with_deflation(&matvec, 6, 3, 4242);
assert!(
plan.deflation.as_ref().is_some_and(|d| !d.basis.is_empty()),
"deflation basis must have been frozen"
);
let eval = plan.evaluate(&matvec, 1e-13, 20_000).expect("eval");
let grad = plan
.directional_derivative(&eval, &|v: ArrayView1<f64>| da.dot(&v))
.expect("grad");
let h = 1e-5;
let a_plus = &a + &(&da * h);
let a_minus = &a - &(&da * h);
let f_plus = plan
.evaluate(&|v: ArrayView1<f64>| a_plus.dot(&v), 1e-13, 20_000)
.expect("eval+")
.estimate;
let f_minus = plan
.evaluate(&|v: ArrayView1<f64>| a_minus.dot(&v), 1e-13, 20_000)
.expect("eval-")
.estimate;
let fd = (f_plus - f_minus) / (2.0 * h);
let rel = (grad - fd).abs() / fd.abs().max(1e-12);
assert!(
rel < 1e-5,
"deflated surrogate gradient {grad:.9e} vs its own FD {fd:.9e} (rel {rel:.3e})"
);
assert!(
grad > 0.0,
"SPD direction must increase log det, got {grad}"
);
}
}