use super::faer::FaerLu;
use simplicial::linalg::{CooMatrix, CsrMatrix, Matrix, Vector};
use std::fmt;
#[derive(Debug, Clone, PartialEq)]
pub enum EigenError {
NoFiniteEigenvalue,
SingularPencil { shift: f64 },
NotConverged { shift: f64, residual: f64 },
}
impl fmt::Display for EigenError {
fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
match self {
Self::NoFiniteEigenvalue => write!(f, "B has no positive-seminorm direction"),
Self::SingularPencil { shift } => write!(
f,
"A - {shift}*B is singular under every shift perturbation"
),
Self::NotConverged { shift, residual } => {
write!(
f,
"no convergence near shift {shift}; worst residual {residual:e}"
)
}
}
}
}
impl std::error::Error for EigenError {}
pub fn sparse_shift_invert_eigen(
a: &CsrMatrix,
b: &CsrMatrix,
shift: f64,
k: usize,
) -> Result<(Vector, Matrix), EigenError> {
let n = a.nrows();
assert_eq!(a.nrows(), a.ncols());
assert_eq!(b.nrows(), n);
assert_eq!(b.ncols(), n);
let k = k.min(n);
if k == 0 {
return Ok((Vector::zeros(0), Matrix::zeros(n, 0)));
}
let a_norm = inf_norm(a);
let b_norm = inf_norm(b);
let (lu, used_shift) = factor_with_retry(a, b, shift, a_norm)?;
let target_dim = (4 * k).max(2 * k + 20).min(n);
let (mut basis, mut bbasis) = seed_block(n, k, b)?;
let proj_cap = (target_dim + 2 * k).min(n);
let mut proj = Matrix::zeros(proj_cap, proj_cap);
let mut dim = 0usize;
const RESIDUAL_TOL: f64 = 1e-9;
const MAX_RESTART_CYCLES: usize = 100;
let mut worst = f64::INFINITY;
for _cycle in 0..=MAX_RESTART_CYCLES {
while dim < target_dim && dim < basis.len() {
expand(&mut basis, &mut bbasis, &mut proj, dim, &lu, b);
dim += 1;
}
let exhausted = dim < target_dim;
let block = proj.view_range(0..dim, 0..dim).into_owned();
let (theta, s) = dense_self_adjoint_eigen(&block);
let mut order: Vec<usize> = (0..dim).collect();
order.sort_by(|&i, &j| theta[j].abs().total_cmp(&theta[i].abs()));
let take = k.min(dim);
let pairs: Vec<(f64, Vector, f64)> = {
let ritz = |idx: usize| -> (f64, Vector, f64) {
let lambda = used_shift + 1.0 / theta[idx];
let y = combine(&basis, |l| s[(l, idx)], dim);
let raw_res = residual(a, b, lambda, &y, a_norm, b_norm);
if raw_res <= RESIDUAL_TOL {
return (lambda, y, raw_res);
}
let mut x = lu.solve(&(b * &y));
let bnorm = x.dot(&(b * &x)).sqrt();
if bnorm > 0.0 {
x /= bnorm;
}
let res = residual(a, b, lambda, &x, a_norm, b_norm);
(lambda, x, res)
};
order[..take].iter().map(|&i| ritz(i)).collect()
};
worst = pairs.iter().map(|p| p.2).fold(0.0_f64, f64::max);
if worst <= RESIDUAL_TOL || exhausted {
let mut pairs = pairs;
pairs.sort_by(|p, q| p.0.total_cmp(&q.0));
let eigenvals = Vector::from_iterator(pairs.len(), pairs.iter().map(|p| p.0));
let eigenvecs = Matrix::from_fn(n, pairs.len(), |i, j| pairs[j].1[i]);
return Ok((eigenvals, eigenvecs));
}
let keep = (2 * k).min(target_dim.saturating_sub(1)).max(1);
let mut new_basis = Vec::with_capacity(keep);
let mut new_bbasis = Vec::with_capacity(keep);
for &idx in &order[..keep] {
new_basis.push(combine(&basis, |l| s[(l, idx)], dim));
new_bbasis.push(combine(&bbasis, |l| s[(l, idx)], dim));
}
basis = new_basis;
bbasis = new_bbasis;
proj = Matrix::zeros(proj_cap, proj_cap);
dim = 0;
}
Err(EigenError::NotConverged {
shift,
residual: worst,
})
}
fn expand(
basis: &mut Vec<Vector>,
bbasis: &mut Vec<Vector>,
proj: &mut Matrix,
idx: usize,
lu: &FaerLu,
b: &CsrMatrix,
) {
let mut w = lu.solve(&bbasis[idx]);
let mut h = vec![0.0; basis.len()];
for _pass in 0..2 {
for (j, (vj, bvj)) in basis.iter().zip(bbasis.iter()).enumerate() {
let c = w.dot(bvj);
h[j] += c;
w.axpy(-c, vj, 1.0);
}
}
for (j, &hj) in h.iter().enumerate() {
proj[(j, idx)] = hj;
proj[(idx, j)] = hj;
}
let bw = b * &w;
let beta_sq = w.dot(&bw);
const BREAKDOWN_TOL_SQ: f64 = 1e-20;
if beta_sq <= BREAKDOWN_TOL_SQ {
return;
}
let beta = beta_sq.sqrt();
basis.push(&w / beta);
bbasis.push(&bw / beta);
}
fn seed_block(
n: usize,
bs: usize,
b: &CsrMatrix,
) -> Result<(Vec<Vector>, Vec<Vector>), EigenError> {
let mut basis = Vec::with_capacity(bs);
let mut bbasis = Vec::with_capacity(bs);
let mut seed = 0u64;
while basis.len() < bs && seed < bs as u64 * 32 + 32 {
let mut v = Vector::from_iterator(n, (0..n).map(|i| pseudo_random(seed, i as u64)));
seed += 1;
for _pass in 0..2 {
for (vj, bvj) in basis.iter().zip(bbasis.iter()) {
let c: f64 = v.dot(bvj);
v.axpy(-c, vj, 1.0);
}
}
let bv = b * &v;
let norm_sq: f64 = v.dot(&bv);
const SEED_TOL: f64 = 1e-24;
if norm_sq > SEED_TOL {
let norm = norm_sq.sqrt();
basis.push(&v / norm);
bbasis.push(&bv / norm);
}
}
if basis.is_empty() {
return Err(EigenError::NoFiniteEigenvalue);
}
Ok((basis, bbasis))
}
fn pseudo_random(seed: u64, index: u64) -> f64 {
let mut z = seed
.wrapping_mul(0x9E37_79B9_7F4A_7C15)
.wrapping_add(index.wrapping_mul(0xD1B5_4A32_D192_ED03))
.wrapping_add(0x9E37_79B9_7F4A_7C15);
z = (z ^ (z >> 30)).wrapping_mul(0xBF58_476D_1CE4_E5B9);
z = (z ^ (z >> 27)).wrapping_mul(0x94D0_49BB_1331_11EB);
z ^= z >> 31;
(z >> 11) as f64 / (1u64 << 53) as f64 * 2.0 - 1.0
}
fn combine(vecs: &[Vector], coeff: impl Fn(usize) -> f64, dim: usize) -> Vector {
let mut out = Vector::zeros(vecs[0].nrows());
for (l, v) in vecs.iter().enumerate().take(dim) {
out.axpy(coeff(l), v, 1.0);
}
out
}
fn residual(
a: &CsrMatrix,
b: &CsrMatrix,
lambda: f64,
x: &Vector,
a_norm: f64,
b_norm: f64,
) -> f64 {
let ax = a * x;
let bx = b * x;
let r = &ax - lambda * &bx;
let xnorm = x.norm();
let scale = a_norm * xnorm + lambda.abs() * b_norm * xnorm;
if scale > 0.0 {
r.norm() / scale
} else {
r.norm()
}
}
fn inf_norm(m: &CsrMatrix) -> f64 {
let mut row_sums = vec![0.0; m.nrows()];
for (r, _, &v) in m.triplet_iter() {
row_sums[r] += v.abs();
}
row_sums.into_iter().fold(0.0_f64, f64::max)
}
fn perturbed_shift(shift: f64, eps0: f64, attempt: usize) -> f64 {
if attempt == 0 {
return shift;
}
let step = eps0 * 2f64.powi(((attempt - 1) / 2) as i32);
if attempt % 2 == 1 {
shift + step
} else {
shift - step
}
}
fn factor_with_retry(
a: &CsrMatrix,
b: &CsrMatrix,
shift: f64,
a_norm: f64,
) -> Result<(FaerLu, f64), EigenError> {
const MAX_ATTEMPTS: usize = 16;
let eps0 = f64::EPSILON.sqrt() * a_norm.max(1.0);
for attempt in 0..MAX_ATTEMPTS {
let current_shift = perturbed_shift(shift, eps0, attempt);
let m = shifted_matrix(a, b, current_shift);
if let Some(lu) = FaerLu::try_new(m.clone()) {
let probe = Vector::from_iterator(
m.nrows(),
(0..m.nrows()).map(|i| pseudo_random(!0, i as u64)),
);
let rhs = &m * &probe;
let resolved = lu.solve(&rhs);
if (&resolved - &probe).norm() <= 1e-6 * probe.norm().max(1.0) {
return Ok((lu, current_shift));
}
}
}
Err(EigenError::SingularPencil { shift })
}
fn shifted_matrix(a: &CsrMatrix, b: &CsrMatrix, shift: f64) -> CsrMatrix {
let mut coo = CooMatrix::new(a.nrows(), a.ncols());
for (r, c, &v) in a.triplet_iter() {
coo.push(r, c, v);
}
for (r, c, &v) in b.triplet_iter() {
coo.push(r, c, -shift * v);
}
CsrMatrix::from(&coo)
}
fn dense_self_adjoint_eigen(m: &Matrix) -> (Vector, Matrix) {
let n = m.nrows();
let fm = faer::Mat::from_fn(n, n, |i, j| m[(i, j)]);
let eig = fm.self_adjoint_eigen(faer::Side::Lower).unwrap();
let vals = eig.S().column_vector();
let vecs = eig.U();
let eigenvals = Vector::from_iterator(n, (0..n).map(|i| *vals.get(i)));
let eigenvecs = Matrix::from_fn(n, n, |i, j| *vecs.get(i, j));
(eigenvals, eigenvecs)
}
#[cfg(test)]
mod test {
use super::sparse_shift_invert_eigen;
use simplicial::linalg::{CooMatrix, CsrMatrix, Matrix, Vector};
fn symmetric(n: usize, f: impl Fn(usize, usize) -> f64) -> Matrix {
Matrix::from_fn(n, n, |i, j| f(i.min(j), i.max(j)))
}
fn csr(m: &Matrix) -> CsrMatrix {
m.into()
}
#[test]
fn pairs_solve_the_pencil() {
let n = 6;
let a = symmetric(n, |i, j| ((i * 7 + j * 3) % 11) as f64 - 5.0);
let b = symmetric(n, |i, j| if i == j { n as f64 } else { 0.3 });
for nev in 1..=n {
let (vals, vecs) = sparse_shift_invert_eigen(&csr(&a), &csr(&b), 0.0, nev).unwrap();
for k in 0..vals.len() {
let x = vecs.column(k).into_owned();
let residual = (&a * &x - vals[k] * (&b * &x)).norm();
assert!(residual < 1e-9, "nev={nev} k={k} residual={residual:e}");
}
}
}
#[test]
fn matches_symmetric_evd_oracle() {
let n = 7;
let a = symmetric(n, |i, j| ((i * 5 + j * 2) % 13) as f64 - 6.0);
let id = Matrix::identity(n, n);
let mut oracle: Vec<f64> = a.clone().symmetric_eigenvalues().iter().copied().collect();
oracle.sort_by(|x, y| x.abs().total_cmp(&y.abs()));
for nev in 1..=n {
let (vals, _) = sparse_shift_invert_eigen(&csr(&a), &csr(&id), 0.0, nev).unwrap();
let mut got: Vec<f64> = vals.iter().copied().collect();
let mut want = oracle[..nev].to_vec();
got.sort_by(f64::total_cmp);
want.sort_by(f64::total_cmp);
for (g, w) in got.iter().zip(&want) {
assert!((g - w).abs() < 1e-8, "nev={nev}: got {g} want {w}");
}
}
}
#[test]
fn eigenvectors_are_b_orthonormal() {
let n = 6;
let a = symmetric(n, |i, j| if i == j { 2.0 * n as f64 } else { 0.5 });
let b = symmetric(n, |i, j| if i == j { n as f64 } else { 0.3 });
let (_, vecs) = sparse_shift_invert_eigen(&csr(&a), &csr(&b), 0.0, n).unwrap();
for k in 0..vecs.ncols() {
for l in 0..vecs.ncols() {
let vk = vecs.column(k).into_owned();
let vl = vecs.column(l).into_owned();
let ip = vk.dot(&(&b * &vl));
let want = if k == l { 1.0 } else { 0.0 };
assert!((ip - want).abs() < 1e-8, "k={k} l={l} got {ip} want {want}");
}
}
}
#[test]
fn excludes_the_null_space_of_b() {
let n = 5;
let a = symmetric(n, |i, j| ((i * 3 + j * 7) % 11) as f64 - 5.0);
let b = Matrix::from_fn(n, n, |i, j| if i == j && i + 1 < n { 1.0 } else { 0.0 });
let (vals, vecs) = sparse_shift_invert_eigen(&csr(&a), &csr(&b), 0.1, n - 1).unwrap();
assert_eq!(vals.len(), n - 1);
for k in 0..vals.len() {
assert!(vals[k].is_finite());
let x: Vector = vecs.column(k).into_owned();
let residual = (&a * &x - vals[k] * (&b * &x)).norm();
assert!(residual < 1e-8, "k={k} residual={residual:e}");
}
}
#[test]
fn solves_the_scalar_pencil() {
let a = Matrix::from_element(1, 1, 3.0);
let b = Matrix::from_element(1, 1, 4.0);
let (vals, vecs) = sparse_shift_invert_eigen(&csr(&a), &csr(&b), 0.0, 3).unwrap();
assert_eq!(vals.len(), 1);
assert!((vals[0] - 3.0 / 4.0).abs() < 1e-9);
let x: Vector = vecs.column(0).into_owned();
assert!(
(x.dot(&(&b * &x)) - 1.0).abs() < 1e-9,
"eigenvector is B-normalized"
);
assert!((&a * &x - vals[0] * (&b * &x)).norm() < 1e-9);
let (vals0, vecs0) = sparse_shift_invert_eigen(&csr(&a), &csr(&b), 0.0, 0).unwrap();
assert_eq!(vals0.len(), 0);
assert_eq!(vecs0.ncols(), 0);
}
#[test]
fn reports_a_pencil_without_finite_eigenvalues() {
let a = Matrix::identity(4, 4);
let b = Matrix::zeros(4, 4);
assert_eq!(
sparse_shift_invert_eigen(&csr(&a), &csr(&b), 0.0, 2),
Err(super::EigenError::NoFiniteEigenvalue)
);
}
#[test]
fn resolves_a_degenerate_cluster_at_the_shift() {
let m = 3;
let rest = 4;
let n = m + rest;
let b = Matrix::identity(n, n);
let diag = Matrix::from_fn(n, n, |i, j| {
if i != j || i < m {
0.0
} else {
(i - m + 1) as f64
}
});
let rot = Matrix::from_fn(n, n, |i, j| ((i * 5 + j * 3 + 1) % 7) as f64 - 3.0);
let rot = rot.clone() + rot.transpose() + Matrix::identity(n, n) * (2.0 * n as f64);
let a = &rot * &diag * &rot.transpose();
let a = (&a + a.transpose()) * 0.5;
let (vals, vecs) = sparse_shift_invert_eigen(&csr(&a), &csr(&b), 0.0, m).unwrap();
assert_eq!(vals.len(), m);
for &v in vals.iter() {
assert!(v.abs() < 1e-6, "expected a near-zero eigenvalue, got {v}");
}
for k in 0..m {
let x = vecs.column(k).into_owned();
let residual = (&a * &x - vals[k] * &x).norm();
assert!(residual < 1e-6, "k={k} residual={residual:e}");
}
let gram = vecs.transpose() * &vecs;
let dev = (&gram - Matrix::identity(m, m)).norm();
assert!(
dev < 1e-6,
"recovered directions are not mutually independent: {gram}"
);
}
#[test]
fn handles_a_large_sparse_pencil() {
let n = 3000;
let nev = 5;
let mut coo = CooMatrix::new(n, n);
for i in 0..n {
coo.push(i, i, 2.0);
if i + 1 < n {
coo.push(i, i + 1, -1.0);
coo.push(i + 1, i, -1.0);
}
}
let a = CsrMatrix::from(&coo);
let mut ident = CooMatrix::new(n, n);
for i in 0..n {
ident.push(i, i, 1.0);
}
let b = CsrMatrix::from(&ident);
let (vals, vecs) = sparse_shift_invert_eigen(&a, &b, 0.0, nev).unwrap();
assert_eq!(vals.len(), nev);
for k in 0..nev {
let x = vecs.column(k).into_owned();
let residual = (&a * &x - vals[k] * (&b * &x)).norm();
assert!(residual < 1e-6, "k={k} residual={residual:e}");
let want = 2.0 - 2.0 * (((k + 1) as f64 * std::f64::consts::PI) / (n as f64 + 1.0)).cos();
assert!(
(vals[k] - want).abs() < 1e-6,
"k={k}: got {} want {want}",
vals[k]
);
}
}
}