use oxiblas_core::scalar::{Field, Real, Scalar};
use oxiblas_matrix::Mat;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub(crate) enum BidiagDcError {
NotConverged,
SecularEquationFailed,
}
const DIRECT_THRESHOLD: usize = 25;
const CUPPEN_BASE: usize = 8;
const MAX_QR_SWEEPS: usize = 40;
const MAX_SECULAR_ITER: usize = 80;
const MAX_JACOBI_SWEEPS: usize = 100;
#[inline]
fn fabs<R: Real>(x: R) -> R {
if x < R::zero() { -x } else { x }
}
#[inline]
fn from_f64<R: Real>(v: f64) -> R {
R::from_f64(v).unwrap_or_else(R::zero)
}
#[inline]
fn rsqrt<R: Real>(x: R) -> R {
<R as Real>::sqrt(x)
}
pub(crate) fn bidiagonal_svd_dc<R>(
d: &[R],
e: &[R],
) -> Result<(Mat<R>, Vec<R>, Mat<R>), BidiagDcError>
where
R: Field + Real + bytemuck::Zeroable,
{
let n = d.len();
if n == 0 {
return Ok((Mat::zeros(0, 0), Vec::new(), Mat::zeros(0, 0)));
}
if n == 1 {
let sigma = vec![fabs(d[0])];
let mut u = Mat::zeros(1, 1);
let mut vt = Mat::zeros(1, 1);
u[(0, 0)] = if d[0] >= R::zero() {
R::one()
} else {
-R::one()
};
vt[(0, 0)] = R::one();
return Ok((u, sigma, vt));
}
if n <= DIRECT_THRESHOLD {
return bidiagonal_svd_qr(d, e);
}
let mut t_diag = vec![R::zero(); n];
let mut t_off = vec![R::zero(); n - 1];
for i in 0..n {
let mut diag = d[i] * d[i];
if i > 0 {
diag = diag + e[i - 1] * e[i - 1];
}
t_diag[i] = diag;
}
for i in 0..n - 1 {
t_off[i] = d[i] * e[i];
}
let (evals, vmat) = cuppen_tridiag(&t_diag, &t_off)?;
assemble_svd(d, e, &evals, vmat)
}
fn assemble_svd<R>(
d: &[R],
e: &[R],
evals: &[R],
mut vmat: Mat<R>,
) -> Result<(Mat<R>, Vec<R>, Mat<R>), BidiagDcError>
where
R: Field + Real + bytemuck::Zeroable,
{
let n = d.len();
if evals.len() != n || vmat.nrows() != n || vmat.ncols() != n {
return Err(BidiagDcError::SecularEquationFailed);
}
mgs_orthonormalize(&mut vmat);
let mut w = Mat::zeros(n, n);
let mut sigma_col = vec![R::zero(); n];
for j in 0..n {
let mut norm_sq = R::zero();
for i in 0..n {
let mut bv = d[i] * vmat[(i, j)];
if i + 1 < n {
bv = bv + e[i] * vmat[(i + 1, j)];
}
w[(i, j)] = bv;
norm_sq = norm_sq + bv * bv;
}
sigma_col[j] = rsqrt(norm_sq);
}
let mut order: Vec<usize> = (0..n).collect();
order.sort_by(|&a, &b| {
sigma_col[b]
.partial_cmp(&sigma_col[a])
.unwrap_or(core::cmp::Ordering::Equal)
});
let scale = sigma_col.iter().copied().fold(R::zero(), |m, s| m.max(s));
let eps = <R as Scalar>::epsilon();
let tiny = eps * scale.max(R::one()) * from_f64::<R>(8.0);
let mut sigma = vec![R::zero(); n];
let mut u = Mat::zeros(n, n);
let mut vt = Mat::zeros(n, n);
let mut needs_fill = Vec::new();
for (new_j, &old_j) in order.iter().enumerate() {
sigma[new_j] = sigma_col[old_j];
for i in 0..n {
vt[(new_j, i)] = vmat[(i, old_j)];
}
if sigma_col[old_j] > tiny {
let inv = R::one() / sigma_col[old_j];
for i in 0..n {
u[(i, new_j)] = w[(i, old_j)] * inv;
}
} else {
needs_fill.push(new_j);
}
}
if !needs_fill.is_empty() {
fill_orthonormal_columns(&mut u, &needs_fill);
}
Ok((u, sigma, vt))
}
fn bidiagonal_svd_qr<R>(d: &[R], e: &[R]) -> Result<(Mat<R>, Vec<R>, Mat<R>), BidiagDcError>
where
R: Field + Real + bytemuck::Zeroable,
{
let n = d.len();
let mut d_work: Vec<R> = d.to_vec();
let mut e_work: Vec<R> = e.to_vec();
let mut u = Mat::zeros(n, n);
let mut vt = Mat::zeros(n, n);
for i in 0..n {
u[(i, i)] = R::one();
vt[(i, i)] = R::one();
}
let eps = <R as Scalar>::epsilon();
let tol = eps * from_f64::<R>(4.0);
let max_iter = MAX_QR_SWEEPS * n + 20;
let mut converged = false;
for _iter in 0..max_iter {
for i in 0..e_work.len() {
let thresh = tol * (fabs(d_work[i]) + fabs(d_work[i + 1]));
if fabs(e_work[i]) <= thresh {
e_work[i] = R::zero();
}
}
let mut p = e_work.len();
while p > 0 && e_work[p - 1] == R::zero() {
p -= 1;
}
if p == 0 {
converged = true;
break;
}
let mut q = p - 1;
while q > 0 && e_work[q - 1] != R::zero() {
q -= 1;
}
golub_kahan_step(&mut d_work, &mut e_work, &mut u, &mut vt, q, p + 1);
}
if !converged {
return Err(BidiagDcError::NotConverged);
}
for i in 0..n {
if d_work[i] < R::zero() {
d_work[i] = -d_work[i];
for r in 0..n {
u[(r, i)] = -u[(r, i)];
}
}
}
let mut order: Vec<usize> = (0..n).collect();
order.sort_by(|&a, &b| {
d_work[b]
.partial_cmp(&d_work[a])
.unwrap_or(core::cmp::Ordering::Equal)
});
let mut sigma = vec![R::zero(); n];
let mut u_sorted = Mat::zeros(n, n);
let mut vt_sorted = Mat::zeros(n, n);
for (new_idx, &old_idx) in order.iter().enumerate() {
sigma[new_idx] = d_work[old_idx];
for r in 0..n {
u_sorted[(r, new_idx)] = u[(r, old_idx)];
vt_sorted[(new_idx, r)] = vt[(old_idx, r)];
}
}
Ok((u_sorted, sigma, vt_sorted))
}
fn golub_kahan_step<R>(
d: &mut [R],
e: &mut [R],
u: &mut Mat<R>,
vt: &mut Mat<R>,
start: usize,
end: usize,
) where
R: Field + Real + bytemuck::Zeroable,
{
let n = u.nrows();
let last = end - 1;
let d_last = d[last];
let d_prev = d[last - 1];
let e_prev = e[last - 1];
let e_prev2 = if last >= start + 2 {
e[last - 2]
} else {
R::zero()
};
let t11 = d_prev * d_prev + e_prev2 * e_prev2;
let t22 = d_last * d_last + e_prev * e_prev;
let t12 = d_prev * e_prev;
let shift = wilkinson_shift(t11, t12, t22);
let mut f = d[start] * d[start] - shift;
let mut g = d[start] * e[start];
for k in start..last {
let (c, s, r) = givens(f, g);
if k > start {
e[k - 1] = r;
}
f = c * d[k] + s * e[k];
e[k] = c * e[k] - s * d[k];
g = s * d[k + 1];
d[k + 1] = c * d[k + 1];
for j in 0..n {
let vk = vt[(k, j)];
let vk1 = vt[(k + 1, j)];
vt[(k, j)] = c * vk + s * vk1;
vt[(k + 1, j)] = c * vk1 - s * vk;
}
let (c, s, r) = givens(f, g);
d[k] = r;
f = c * e[k] + s * d[k + 1];
d[k + 1] = c * d[k + 1] - s * e[k];
if k < last - 1 {
g = s * e[k + 1];
e[k + 1] = c * e[k + 1];
}
for i in 0..n {
let uk = u[(i, k)];
let uk1 = u[(i, k + 1)];
u[(i, k)] = c * uk + s * uk1;
u[(i, k + 1)] = c * uk1 - s * uk;
}
}
e[last - 1] = f;
}
#[inline]
fn wilkinson_shift<R: Real>(t11: R, t12: R, t22: R) -> R {
if t12 == R::zero() {
return t22;
}
let two = from_f64::<R>(2.0);
let delta = (t11 - t22) / two;
let sign = if delta < R::zero() {
-R::one()
} else {
R::one()
};
let denom = delta + sign * rsqrt(delta * delta + t12 * t12);
if denom == R::zero() {
t22
} else {
t22 - (t12 * t12) / denom
}
}
#[inline]
fn givens<R: Field + Real>(f: R, g: R) -> (R, R, R) {
if g == R::zero() {
(R::one(), R::zero(), f)
} else if f == R::zero() {
let sign = if g < R::zero() { -R::one() } else { R::one() };
(R::zero(), sign, fabs(g))
} else {
let r = <R as Real>::hypot(f, g);
(f / r, g / r, r)
}
}
fn cuppen_tridiag<R>(diag: &[R], off: &[R]) -> Result<(Vec<R>, Mat<R>), BidiagDcError>
where
R: Field + Real + bytemuck::Zeroable,
{
let (evals, evecs) = cuppen_impl(diag, off)?;
let n = evals.len();
let mut v = Mat::zeros(n, n);
for (k, col) in evecs.iter().enumerate() {
for (i, &val) in col.iter().enumerate() {
v[(i, k)] = val;
}
}
Ok((evals, v))
}
fn cuppen_impl<R>(diag: &[R], off: &[R]) -> Result<(Vec<R>, Vec<Vec<R>>), BidiagDcError>
where
R: Field + Real + bytemuck::Zeroable,
{
let n = diag.len();
if n == 1 {
return Ok((vec![diag[0]], vec![vec![R::one()]]));
}
if n <= CUPPEN_BASE {
return jacobi_tridiag(diag, off);
}
let mid = n / 2;
let beta = off[mid - 1];
let abs_beta = fabs(beta);
let scale = tridiag_scale(diag, off);
let eps = <R as Scalar>::epsilon();
if abs_beta <= eps * scale {
let (ev1, evec1) = cuppen_impl(&diag[..mid], &off[..mid - 1])?;
let (ev2, evec2) = cuppen_impl(&diag[mid..], &off[mid..])?;
return Ok(combine_block_diag(ev1, evec1, ev2, evec2, mid, n));
}
let mut diag1 = diag[..mid].to_vec();
let mut diag2 = diag[mid..].to_vec();
diag1[mid - 1] = diag1[mid - 1] - abs_beta;
diag2[0] = diag2[0] - abs_beta;
let off1 = &off[..mid - 1];
let off2 = &off[mid..];
let (evals1, evecs1) = cuppen_impl(&diag1, off1)?;
let (evals2, evecs2) = cuppen_impl(&diag2, off2)?;
let d: Vec<R> = evals1.iter().chain(evals2.iter()).copied().collect();
let sign_beta = if beta < R::zero() {
-R::one()
} else {
R::one()
};
let mut u = Vec::with_capacity(n);
for ev1 in evecs1.iter().take(mid) {
u.push(sign_beta * ev1[mid - 1]);
}
for ev2 in evecs2.iter().take(n - mid) {
u.push(ev2[0]);
}
let deflation_tol = eps * from_f64::<R>(n as f64) * scale.max(R::one());
let (d_defl, u_defl, active_idx, trivial, rot) = deflate(&d, &u, deflation_tol);
let (secular_evals, secular_evecs) = if d_defl.is_empty() {
(Vec::new(), Vec::new())
} else {
secular_eigenpairs(&d_defl, &u_defl, abs_beta)?
};
let total = trivial.len() + secular_evals.len();
if total != n {
return jacobi_tridiag(diag, off);
}
let mut pairs: Vec<(R, Vec<R>)> = Vec::with_capacity(n);
for &(eval, col) in trivial.iter() {
let mut ev = vec![R::zero(); n];
for (p, ev_p) in ev.iter_mut().enumerate() {
*ev_p = rot[(p, col)];
}
pairs.push((eval, ev));
}
for (lam, evec_active) in secular_evals.into_iter().zip(secular_evecs) {
let mut ev = vec![R::zero(); n];
for (k, &ai) in active_idx.iter().enumerate() {
let x = evec_active[k];
if x != R::zero() {
for (p, ev_p) in ev.iter_mut().enumerate() {
*ev_p = *ev_p + x * rot[(p, ai)];
}
}
}
pairs.push((lam, ev));
}
pairs.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap_or(core::cmp::Ordering::Equal));
let n2 = n - mid;
let mut evals_final = Vec::with_capacity(n);
let mut evecs_final = Vec::with_capacity(n);
for (lam, v_d) in pairs {
evals_final.push(lam);
let mut out = vec![R::zero(); n];
for (k, evec1) in evecs1.iter().enumerate().take(mid) {
let x = v_d[k];
if x != R::zero() {
for (i, &c) in evec1.iter().enumerate().take(mid) {
out[i] = out[i] + x * c;
}
}
}
for (k, evec2) in evecs2.iter().enumerate().take(n2) {
let x = v_d[mid + k];
if x != R::zero() {
for (i, &c) in evec2.iter().enumerate().take(n2) {
out[mid + i] = out[mid + i] + x * c;
}
}
}
evecs_final.push(out);
}
Ok((evals_final, evecs_final))
}
#[allow(clippy::type_complexity)]
fn combine_block_diag<R: Real>(
ev1: Vec<R>,
evec1: Vec<Vec<R>>,
ev2: Vec<R>,
evec2: Vec<Vec<R>>,
mid: usize,
n: usize,
) -> (Vec<R>, Vec<Vec<R>>) {
let n2 = n - mid;
let mut pairs: Vec<(R, Vec<R>)> = Vec::with_capacity(n);
for (k, lam) in ev1.into_iter().enumerate() {
let mut ev = vec![R::zero(); n];
for i in 0..mid {
ev[i] = evec1[k][i];
}
pairs.push((lam, ev));
}
for (k, lam) in ev2.into_iter().enumerate() {
let mut ev = vec![R::zero(); n];
for i in 0..n2 {
ev[mid + i] = evec2[k][i];
}
pairs.push((lam, ev));
}
pairs.sort_by(|a, b| a.0.partial_cmp(&b.0).unwrap_or(core::cmp::Ordering::Equal));
let evals = pairs.iter().map(|p| p.0).collect();
let evecs = pairs.into_iter().map(|p| p.1).collect();
(evals, evecs)
}
fn tridiag_scale<R: Real>(diag: &[R], off: &[R]) -> R {
let mut s = R::zero();
for &x in diag {
s = s.max(fabs(x));
}
for &x in off {
s = s.max(fabs(x));
}
if s > R::zero() { s } else { R::one() }
}
#[allow(clippy::type_complexity)]
fn deflate<R>(d: &[R], u: &[R], tol: R) -> (Vec<R>, Vec<R>, Vec<usize>, Vec<(R, usize)>, Mat<R>)
where
R: Field + Real + bytemuck::Zeroable,
{
let n = d.len();
let mut u_norm_sq = R::zero();
for &x in u {
u_norm_sq = u_norm_sq + x * x;
}
let u_norm = rsqrt(u_norm_sq).max(R::min_positive());
let mut rot = Mat::zeros(n, n);
for i in 0..n {
rot[(i, i)] = R::one();
}
let mut defl_d = Vec::new();
let mut defl_u = Vec::new();
let mut active_idx = Vec::new();
let mut trivial: Vec<(R, usize)> = Vec::new();
for i in 0..n {
if fabs(u[i]) < tol * u_norm {
trivial.push((d[i], i));
} else {
defl_d.push(d[i]);
defl_u.push(u[i]);
active_idx.push(i);
}
}
let mut i = 0;
while i < defl_d.len() {
let mut j = i + 1;
while j < defl_d.len() {
if fabs(defl_d[j] - defl_d[i]) < tol * (fabs(defl_d[i]) + R::one()) {
let ui = defl_u[i];
let uj = defl_u[j];
let r = rsqrt(ui * ui + uj * uj);
let pos_i = active_idx[i];
let pos_j = active_idx[j];
if r > R::min_positive() {
let c = ui / r;
let s = uj / r;
for p in 0..n {
let a = rot[(p, pos_i)];
let b = rot[(p, pos_j)];
rot[(p, pos_i)] = c * a + s * b;
rot[(p, pos_j)] = c * b - s * a;
}
}
defl_u[i] = r;
trivial.push((defl_d[j], pos_j));
defl_d.remove(j);
defl_u.remove(j);
active_idx.remove(j);
} else {
j += 1;
}
}
i += 1;
}
(defl_d, defl_u, active_idx, trivial, rot)
}
#[allow(clippy::type_complexity)]
fn secular_eigenpairs<R>(d: &[R], u: &[R], beta: R) -> Result<(Vec<R>, Vec<Vec<R>>), BidiagDcError>
where
R: Field + Real,
{
let m = d.len();
if m == 0 {
return Ok((Vec::new(), Vec::new()));
}
let mut idx: Vec<usize> = (0..m).collect();
idx.sort_by(|&a, &b| {
d[a].partial_cmp(&d[b])
.unwrap_or(core::cmp::Ordering::Equal)
});
let d_s: Vec<R> = idx.iter().map(|&i| d[i]).collect();
let u_s: Vec<R> = idx.iter().map(|&i| u[i]).collect();
let mut weight_sum = R::zero();
for &ui in &u_s {
weight_sum = weight_sum + ui * ui;
}
let mut lam = vec![R::zero(); m];
for i in 0..m {
let (lo, hi) = if i < m - 1 {
(d_s[i], d_s[i + 1])
} else {
(d_s[m - 1], d_s[m - 1] + beta * weight_sum + R::one())
};
lam[i] = find_secular_root(&d_s, &u_s, beta, lo, hi)?;
}
lam.sort_by(|a, b| a.partial_cmp(b).unwrap_or(core::cmp::Ordering::Equal));
let mut w_hat = vec![R::zero(); m];
for i in 0..m {
let mut prod = R::one();
for k in 0..m {
prod = prod * (lam[k] - d_s[i]);
if k != i {
prod = prod / (d_s[k] - d_s[i]);
}
}
let mag = rsqrt(prod.max(R::zero()));
let sign = if u_s[i] < R::zero() {
-R::one()
} else {
R::one()
};
w_hat[i] = sign * mag;
}
let mut evecs = vec![vec![R::zero(); m]; m];
for k in 0..m {
let mut col = vec![R::zero(); m];
let mut norm_sq = R::zero();
for i in 0..m {
let denom = d_s[i] - lam[k];
let val = if fabs(denom) < R::min_positive() {
let s = if w_hat[i] < R::zero() {
-R::one()
} else {
R::one()
};
s / R::min_positive()
} else {
w_hat[i] / denom
};
col[i] = val;
norm_sq = norm_sq + val * val;
}
let norm = rsqrt(norm_sq);
let inv = if norm > R::min_positive() {
R::one() / norm
} else {
R::one()
};
for i in 0..m {
evecs[k][idx[i]] = col[i] * inv;
}
}
Ok((lam, evecs))
}
#[inline]
fn secular_f_df<R: Field + Real>(d: &[R], u: &[R], beta: R, lam: R) -> (R, R) {
let mut f = R::one();
let mut df = R::zero();
for (&di, &ui) in d.iter().zip(u.iter()) {
let t = di - lam;
let w = ui * ui / t;
f = f + beta * w;
df = df + beta * w / t;
}
(f, df)
}
fn find_secular_root<R>(d: &[R], u: &[R], beta: R, lo: R, hi: R) -> Result<R, BidiagDcError>
where
R: Field + Real,
{
let eps = <R as Scalar>::epsilon();
let two = from_f64::<R>(2.0);
let margin = eps * (fabs(lo) + fabs(hi) + R::one()) * from_f64::<R>(4.0);
let mut lo_m = lo + margin;
let mut hi_m = hi - margin;
if hi_m <= lo_m {
return Ok((lo + hi) / two);
}
let conv = eps * from_f64::<R>(8.0);
let mut x = (lo_m + hi_m) / two;
for _ in 0..MAX_SECULAR_ITER {
let (fx, dfx) = secular_f_df(d, u, beta, x);
if !fx.is_finite() {
x = (lo_m + hi_m) / two;
continue;
}
if fx < R::zero() {
lo_m = x;
} else if fx > R::zero() {
hi_m = x;
} else {
return Ok(x);
}
let newton_ok = dfx > R::min_positive();
let x_new = if newton_ok { x - fx / dfx } else { x };
let next = if newton_ok && x_new > lo_m && x_new < hi_m {
x_new
} else {
(lo_m + hi_m) / two
};
let step = fabs(next - x);
x = next;
if step < conv * (fabs(x) + R::one()) || hi_m - lo_m < conv * (fabs(x) + R::one()) {
return Ok(x);
}
}
Ok(x)
}
fn jacobi_tridiag<R>(diag: &[R], off: &[R]) -> Result<(Vec<R>, Vec<Vec<R>>), BidiagDcError>
where
R: Field + Real + bytemuck::Zeroable,
{
let n = diag.len();
if n == 1 {
return Ok((vec![diag[0]], vec![vec![R::one()]]));
}
let mut a = Mat::zeros(n, n);
for i in 0..n {
a[(i, i)] = diag[i];
}
for i in 0..n - 1 {
a[(i, i + 1)] = off[i];
a[(i + 1, i)] = off[i];
}
let mut q = Mat::zeros(n, n);
for i in 0..n {
q[(i, i)] = R::one();
}
let eps = <R as Scalar>::epsilon();
let scale = tridiag_scale(diag, off);
let tol = eps * scale * from_f64::<R>(n as f64);
let two = from_f64::<R>(2.0);
let mut converged = false;
for _sweep in 0..MAX_JACOBI_SWEEPS {
let mut off_norm = R::zero();
for i in 0..n {
for j in (i + 1)..n {
off_norm = off_norm + a[(i, j)] * a[(i, j)];
}
}
if rsqrt(off_norm) <= tol {
converged = true;
break;
}
for p in 0..n {
for qcol in (p + 1)..n {
let apq = a[(p, qcol)];
if apq == R::zero() {
continue;
}
let app = a[(p, p)];
let aqq = a[(qcol, qcol)];
let theta = (aqq - app) / (two * apq);
let sign_t = if theta < R::zero() {
-R::one()
} else {
R::one()
};
let t = sign_t / (fabs(theta) + rsqrt(theta * theta + R::one()));
let c = R::one() / rsqrt(t * t + R::one());
let s = t * c;
for k in 0..n {
let akp = a[(k, p)];
let akq = a[(k, qcol)];
a[(k, p)] = c * akp - s * akq;
a[(k, qcol)] = s * akp + c * akq;
}
for k in 0..n {
let apk = a[(p, k)];
let aqk = a[(qcol, k)];
a[(p, k)] = c * apk - s * aqk;
a[(qcol, k)] = s * apk + c * aqk;
}
a[(p, qcol)] = R::zero();
a[(qcol, p)] = R::zero();
for k in 0..n {
let qkp = q[(k, p)];
let qkq = q[(k, qcol)];
q[(k, p)] = c * qkp - s * qkq;
q[(k, qcol)] = s * qkp + c * qkq;
}
}
}
}
if !converged {
return Err(BidiagDcError::NotConverged);
}
let mut pairs: Vec<(R, Vec<R>)> = (0..n)
.map(|k| {
let col: Vec<R> = (0..n).map(|i| q[(i, k)]).collect();
(a[(k, k)], col)
})
.collect();
pairs.sort_by(|x, y| x.0.partial_cmp(&y.0).unwrap_or(core::cmp::Ordering::Equal));
let evals = pairs.iter().map(|p| p.0).collect();
let evecs = pairs.into_iter().map(|p| p.1).collect();
Ok((evals, evecs))
}
fn mgs_orthonormalize<R>(mat: &mut Mat<R>)
where
R: Field + Real + bytemuck::Zeroable,
{
let m = mat.nrows();
let n = mat.ncols();
let eps = <R as Scalar>::epsilon();
let tol = eps * from_f64::<R>(16.0);
for j in 0..n {
for k in 0..j {
let mut dot = R::zero();
for i in 0..m {
dot = dot + mat[(i, j)] * mat[(i, k)];
}
for i in 0..m {
mat[(i, j)] = mat[(i, j)] - dot * mat[(i, k)];
}
}
let mut norm_sq = R::zero();
for i in 0..m {
norm_sq = norm_sq + mat[(i, j)] * mat[(i, j)];
}
let norm = rsqrt(norm_sq);
if norm > tol {
let inv = R::one() / norm;
for i in 0..m {
mat[(i, j)] = mat[(i, j)] * inv;
}
} else {
fill_orthonormal_columns(mat, &[j]);
}
}
}
fn fill_orthonormal_columns<R>(mat: &mut Mat<R>, cols: &[usize])
where
R: Field + Real + bytemuck::Zeroable,
{
let m = mat.nrows();
let n = mat.ncols();
let eps = <R as Scalar>::epsilon();
let tol = eps * from_f64::<R>(16.0);
for &j in cols {
for i in 0..m {
mat[(i, j)] = R::zero();
}
let mut placed = false;
for basis in 0..m {
for i in 0..m {
mat[(i, j)] = if i == basis { R::one() } else { R::zero() };
}
for k in 0..n {
if k == j {
continue;
}
let mut dot = R::zero();
for i in 0..m {
dot = dot + mat[(i, j)] * mat[(i, k)];
}
for i in 0..m {
mat[(i, j)] = mat[(i, j)] - dot * mat[(i, k)];
}
}
let mut norm_sq = R::zero();
for i in 0..m {
norm_sq = norm_sq + mat[(i, j)] * mat[(i, j)];
}
let norm = rsqrt(norm_sq);
if norm > tol {
let inv = R::one() / norm;
for i in 0..m {
mat[(i, j)] = mat[(i, j)] * inv;
}
placed = true;
break;
}
}
if !placed {
for i in 0..m {
mat[(i, j)] = R::zero();
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
fn build_bidiag(d: &[f64], e: &[f64]) -> Vec<Vec<f64>> {
let n = d.len();
let mut b = vec![vec![0.0; n]; n];
for i in 0..n {
b[i][i] = d[i];
if i + 1 < n {
b[i][i + 1] = e[i];
}
}
b
}
fn svd_errors(d: &[f64], e: &[f64]) -> (f64, f64, f64, Vec<f64>) {
let n = d.len();
let (u, sigma, vt) = bidiagonal_svd_dc(d, e).expect("svd");
let mut uu_err = 0.0f64;
let mut vv_err = 0.0f64;
for i in 0..n {
for j in 0..n {
let mut uu = 0.0;
let mut vv = 0.0;
for k in 0..n {
uu += u[(k, i)] * u[(k, j)];
vv += vt[(i, k)] * vt[(j, k)];
}
let expect = if i == j { 1.0 } else { 0.0 };
uu_err = uu_err.max((uu - expect).abs());
vv_err = vv_err.max((vv - expect).abs());
}
}
let b = build_bidiag(d, e);
let mut rec_err = 0.0f64;
for i in 0..n {
for j in 0..n {
let mut acc = 0.0;
for k in 0..n {
acc += u[(i, k)] * sigma[k] * vt[(k, j)];
}
rec_err = rec_err.max((acc - b[i][j]).abs());
}
}
for i in 1..n {
assert!(sigma[i] <= sigma[i - 1] + 1e-12, "not descending at {i}");
}
(uu_err, vv_err, rec_err, sigma)
}
#[test]
fn svd_small_qr_path() {
let d = vec![2.0, 3.0, 1.5, 4.0, 0.7];
let e = vec![0.5, -0.3, 0.8, 0.2];
let (uu, vv, rec, _) = svd_errors(&d, &e);
assert!(uu < 1e-12 && vv < 1e-12 && rec < 1e-12, "{uu} {vv} {rec}");
}
#[test]
fn svd_accuracy_across_sizes() {
for &n in &[26usize, 50, 100, 200] {
let d: Vec<f64> = (0..n)
.map(|i| 2.0 + (i as f64 * 0.37).sin() + 0.5 * (i as f64 * 1.7).cos())
.collect();
let e: Vec<f64> = (0..n - 1)
.map(|i| 0.6 + 0.4 * (i as f64 * 0.9).cos() + 0.2 * (i as f64 * 0.3).sin())
.collect();
let (uu, vv, rec, _) = svd_errors(&d, &e);
println!("n={n}: UtU={uu:.2e} VtV={vv:.2e} recon={rec:.2e}");
assert!(uu < 1e-6, "n={n} UtU={uu}");
assert!(vv < 1e-9, "n={n} VtV={vv}");
assert!(rec < 1e-8, "n={n} recon={rec}");
}
}
#[test]
fn svd_dc_path_with_zero_coupling() {
let mut d = vec![0.0; 30];
let mut e = vec![0.6; 29];
for (i, di) in d.iter_mut().enumerate() {
*di = 3.0 + (i as f64 * 0.5).sin();
}
e[15] = 0.0;
let (uu, vv, rec, _) = svd_errors(&d, &e);
assert!(uu < 1e-6 && vv < 1e-9 && rec < 1e-8, "{uu} {vv} {rec}");
}
#[test]
fn deflate_coincident_merge_applies_givens() {
let d = vec![1.0_f64, 5.0, 5.0, 9.0];
let u = vec![0.5_f64, 0.8, 0.6, 0.4];
let beta = 0.7_f64; let tol = 1e-12_f64;
let n = d.len();
let (d_defl, u_defl, active_idx, trivial, rot) = deflate(&d, &u, tol);
assert_eq!(d_defl.len(), 3, "one active component should have deflated");
assert_eq!(u_defl.len(), 3);
assert_eq!(active_idx.len(), 3);
assert_eq!(
trivial.len(),
1,
"exactly one coincident pair should deflate"
);
let survivor = active_idx
.iter()
.position(|&p| p == 1)
.expect("position 1 must remain active");
assert!(
(u_defl[survivor] - 1.0).abs() < 1e-14,
"combined coupling r, got {}",
u_defl[survivor]
);
let (defl_lambda, defl_col) = trivial[0];
assert_eq!(defl_col, 2, "deflated column");
assert!((defl_lambda - 5.0).abs() < 1e-14, "deflated eigenvalue");
let mut max_orth = 0.0_f64;
for a in 0..n {
for b in 0..n {
let mut dot = 0.0;
for p in 0..n {
dot += rot[(p, a)] * rot[(p, b)];
}
let expect = if a == b { 1.0 } else { 0.0 };
max_orth = max_orth.max((dot - expect).abs());
}
}
assert!(max_orth < 1e-14, "rot not orthonormal: {max_orth}");
let u_hat: Vec<f64> = (0..n)
.map(|k| (0..n).map(|p| rot[(p, k)] * u[p]).sum::<f64>())
.collect();
assert!(u_hat[2].abs() < 1e-14, "coupling not zeroed: {}", u_hat[2]);
assert!(
(u_hat[1] - 1.0).abs() < 1e-14,
"survivor coupling: {}",
u_hat[1]
);
let q: Vec<f64> = (0..n).map(|p| rot[(p, 2)]).collect();
let u_dot_q: f64 = (0..n).map(|p| u[p] * q[p]).sum();
let mut resid = 0.0_f64;
for p in 0..n {
let mq = d[p] * q[p] + beta * u[p] * u_dot_q;
resid = resid.max((mq - defl_lambda * q[p]).abs());
}
assert!(
resid < 1e-13,
"deflated column not an eigenvector: residual {resid}"
);
}
#[test]
fn svd_repeated_singular_values_reconstructs() {
let k = 16usize;
let d0: Vec<f64> = (0..k)
.map(|i| 2.5 + (i as f64 * 0.53).sin() + 0.4 * (i as f64 * 1.3).cos())
.collect();
let e0: Vec<f64> = (0..k - 1)
.map(|i| 0.7 + 0.3 * (i as f64 * 0.8).cos())
.collect();
let mut d = Vec::with_capacity(2 * k);
d.extend_from_slice(&d0);
d.extend_from_slice(&d0);
let mut e = Vec::with_capacity(2 * k - 1);
e.extend_from_slice(&e0);
e.push(0.0); e.extend_from_slice(&e0);
let (uu, vv, rec, sigma) = svd_errors(&d, &e);
for (a, &sa) in sigma.iter().enumerate() {
let has_partner = sigma
.iter()
.enumerate()
.any(|(b, &sb)| b != a && (sa - sb).abs() < 1e-8);
assert!(has_partner, "no coincident partner for sigma[{a}] = {sa}");
}
assert!(uu < 1e-6, "UtU orthonormality: {uu}");
assert!(vv < 1e-9, "VtV orthonormality: {vv}");
assert!(rec < 1e-8, "reconstruction: {rec}");
}
}