use crate::error::EmlError;
#[derive(Debug, Clone)]
pub struct QrFactors {
pub data: Vec<f64>,
pub betas: Vec<f64>,
pub m: usize,
pub n: usize,
}
pub fn qr(a: &[f64], m: usize, n: usize) -> Result<QrFactors, EmlError> {
if a.len() != m * n {
return Err(EmlError::DimensionMismatch(m * n, a.len()));
}
if m < n {
return Err(EmlError::DimensionMismatch(m, n));
}
let k = m.min(n);
let mut data = a.to_vec();
let mut betas = vec![0.0f64; k];
for j in 0..k {
let col_len = m - j;
let mut v: Vec<f64> = (0..col_len).map(|i| data[(j + i) * n + j]).collect();
let norm_x = v.iter().map(|x| x * x).sum::<f64>().sqrt();
if norm_x < 1e-15 {
continue;
}
let sign = if v[0] >= 0.0 { 1.0 } else { -1.0 };
v[0] += sign * norm_x;
let v_sq_sum: f64 = v.iter().map(|x| x * x).sum();
if v_sq_sum < 1e-30 {
continue;
}
let beta = 2.0 / v_sq_sum;
for jj in j..n {
let w: f64 = v
.iter()
.enumerate()
.map(|(i, vi)| vi * data[(j + i) * n + jj])
.sum::<f64>();
for i in 0..col_len {
data[(j + i) * n + jj] -= beta * w * v[i];
}
}
let v0 = v[0];
betas[j] = beta * v0 * v0;
for i in 1..col_len {
data[(j + i) * n + j] = v[i] / v0;
}
}
Ok(QrFactors { data, betas, m, n })
}
pub fn q_from_qr(factors: &QrFactors) -> Vec<f64> {
let m = factors.m;
let k = m.min(factors.n);
let mut q = vec![0.0f64; m * m];
for i in 0..m {
q[i * m + i] = 1.0;
}
for j in (0..k).rev() {
let col_len = m - j;
let mut v = vec![1.0f64];
for i in 1..col_len {
v.push(factors.data[(j + i) * factors.n + j]);
}
let beta = factors.betas[j];
for jj in 0..m {
let w: f64 = v
.iter()
.enumerate()
.map(|(i, vi)| vi * q[(j + i) * m + jj])
.sum::<f64>();
for i in 0..col_len {
q[(j + i) * m + jj] -= beta * w * v[i];
}
}
}
q
}
pub(crate) fn apply_qt_to_vec(factors: &QrFactors, b: &mut [f64]) {
let k = factors.m.min(factors.n);
for j in 0..k {
let col_len = factors.m - j;
let mut v = vec![1.0f64];
for i in 1..col_len {
v.push(factors.data[(j + i) * factors.n + j]);
}
let beta = factors.betas[j];
let w: f64 = v
.iter()
.enumerate()
.map(|(i, vi)| vi * b[j + i])
.sum::<f64>();
for i in 0..col_len {
b[j + i] -= beta * w * v[i];
}
}
}
#[derive(Debug, Clone)]
pub struct SvdResult {
pub u: Vec<f64>,
pub s: Vec<f64>,
pub v: Vec<f64>,
pub m: usize,
pub n: usize,
}
fn jacobi_svd_square(b_cols: &[f64], n: usize) -> (Vec<f64>, Vec<f64>, Vec<f64>) {
let mut c = vec![0.0f64; n * n];
for i in 0..n {
for j in i..n {
let dot: f64 = (0..n).map(|k| b_cols[k + i * n] * b_cols[k + j * n]).sum();
c[i * n + j] = dot;
c[j * n + i] = dot;
}
}
let mut v = vec![0.0f64; n * n];
for i in 0..n {
v[i * n + i] = 1.0;
}
const MAX_SWEEPS: usize = 100;
const TOL: f64 = 1e-14;
for _sweep in 0..MAX_SWEEPS {
let mut converged = true;
for p in 0..n {
for q in (p + 1)..n {
let cpq = c[p * n + q];
if cpq.abs() <= TOL * (c[p * n + p] * c[q * n + q]).abs().sqrt() {
continue;
}
converged = false;
let theta = (c[q * n + q] - c[p * n + p]) / (2.0 * cpq);
let t = if theta.abs() < 1e15 {
let sign = if theta >= 0.0 { 1.0 } else { -1.0 };
sign / (theta.abs() + (1.0 + theta * theta).sqrt())
} else {
1.0 / (2.0 * theta)
};
let cos_a = 1.0 / (1.0 + t * t).sqrt();
let sin_a = t * cos_a;
let cpp = c[p * n + p];
let cqq = c[q * n + q];
c[p * n + p] = cpp - t * cpq;
c[q * n + q] = cqq + t * cpq;
c[p * n + q] = 0.0;
c[q * n + p] = 0.0;
for r in 0..n {
if r != p && r != q {
let crp = c[r * n + p];
let crq = c[r * n + q];
let new_crp = cos_a * crp - sin_a * crq;
let new_crq = sin_a * crp + cos_a * crq;
c[r * n + p] = new_crp;
c[p * n + r] = new_crp;
c[r * n + q] = new_crq;
c[q * n + r] = new_crq;
}
}
for i in 0..n {
let vip = v[i * n + p];
let viq = v[i * n + q];
v[i * n + p] = cos_a * vip - sin_a * viq;
v[i * n + q] = sin_a * vip + cos_a * viq;
}
}
}
if converged {
break;
}
}
let singular_vals: Vec<f64> = (0..n).map(|j| c[j * n + j].max(0.0).sqrt()).collect();
let mut u_cols = vec![0.0f64; n * n];
for j in 0..n {
let sigma = singular_vals[j];
if sigma > 1e-15 {
for i in 0..n {
let mut s = 0.0f64;
for p in 0..n {
s += b_cols[i + p * n] * v[p * n + j];
}
u_cols[i + j * n] = s / sigma;
}
}
}
(u_cols, singular_vals, v)
}
pub fn svd(a: &[f64], m: usize, n: usize) -> Result<SvdResult, EmlError> {
if a.len() != m * n {
return Err(EmlError::DimensionMismatch(m * n, a.len()));
}
if m < n {
return Err(EmlError::DimensionMismatch(m, n));
}
let qr_f = qr(a, m, n)?;
let mut r_cols = vec![0.0f64; n * n]; for i in 0..n {
for j in i..n {
r_cols[i + j * n] = qr_f.data[i * n + j];
}
}
let (u_r_cols, singular_vals, v_r_row) = jacobi_svd_square(&r_cols, n);
let q_full = q_from_qr(&qr_f);
let mut order: Vec<usize> = (0..n).collect();
order.sort_by(|&a_idx, &b_idx| singular_vals[b_idx].total_cmp(&singular_vals[a_idx]));
let s_sorted: Vec<f64> = order.iter().map(|&i| singular_vals[i]).collect();
let mut u_sorted = vec![0.0f64; m * m];
for (new_j, &old_j) in order.iter().enumerate() {
for i in 0..m {
let mut s = 0.0f64;
for k in 0..n {
s += q_full[i * m + k] * u_r_cols[k + old_j * n];
}
u_sorted[i * m + new_j] = s;
}
}
for new_j in n..m {
let q_col = new_j; for i in 0..m {
u_sorted[i * m + new_j] = q_full[i * m + q_col];
}
}
let mut v_sorted = vec![0.0f64; n * n];
for (new_j, &old_j) in order.iter().enumerate() {
for i in 0..n {
v_sorted[i * n + new_j] = v_r_row[i * n + old_j];
}
}
if s_sorted.iter().any(|s| !s.is_finite()) {
return Err(EmlError::SingularMatrix);
}
Ok(SvdResult {
u: u_sorted,
s: s_sorted,
v: v_sorted,
m,
n,
})
}
#[cfg(test)]
mod tests {
use super::*;
fn matmul_rm(a: &[f64], p: usize, q: usize, b: &[f64], r: usize) -> Vec<f64> {
let mut c = vec![0.0f64; p * r];
for i in 0..p {
for k in 0..q {
for j in 0..r {
c[i * r + j] += a[i * q + k] * b[k * r + j];
}
}
}
c
}
fn transpose_rm(a: &[f64], p: usize, q: usize) -> Vec<f64> {
let mut t = vec![0.0f64; q * p];
for i in 0..p {
for j in 0..q {
t[j * p + i] = a[i * q + j];
}
}
t
}
#[test]
fn test_qr_qt_q_is_identity() {
let a = vec![1.0_f64, 4.0, 2.0, 5.0, 3.0, 6.0];
let m = 3;
let n = 2;
let factors = qr(&a, m, n).unwrap();
let q = q_from_qr(&factors);
let qt = transpose_rm(&q, m, m);
let qtq = matmul_rm(&qt, m, m, &q, m);
for i in 0..m {
for j in 0..m {
let expected = if i == j { 1.0 } else { 0.0 };
assert!(
(qtq[i * m + j] - expected).abs() < 1e-10,
"Q^T Q [{i},{j}] = {} (expected {})",
qtq[i * m + j],
expected
);
}
}
}
#[test]
fn test_qr_q_r_equals_a() {
let a = vec![1.0_f64, 4.0, 2.0, 5.0, 3.0, 6.0];
let m = 3;
let n = 2;
let factors = qr(&a, m, n).unwrap();
let q = q_from_qr(&factors);
let mut r_full = vec![0.0f64; m * n];
for i in 0..n {
for j in i..n {
r_full[i * n + j] = factors.data[i * n + j];
}
}
let qr_prod = matmul_rm(&q, m, m, &r_full, n);
for (idx, (v1, v2)) in qr_prod.iter().zip(a.iter()).enumerate() {
assert!(
(v1 - v2).abs() < 1e-9,
"QR[{idx}] = {} != A[{idx}] = {}",
v1,
v2
);
}
}
#[test]
fn test_svd_reconstruction() {
let a = vec![
1.0_f64, 5.0, 9.0, 2.0, 6.0, 10.0, 3.0, 7.0, 11.0, 4.0, 8.0, 12.0,
];
let m = 4;
let n = 3;
let svd_r = svd(&a, m, n).unwrap();
let mut diag_s = vec![0.0f64; m * n];
for i in 0..n {
diag_s[i * n + i] = svd_r.s[i];
}
let u_diag = matmul_rm(&svd_r.u, m, m, &diag_s, n);
let vt = transpose_rm(&svd_r.v, n, n);
let recon = matmul_rm(&u_diag, m, n, &vt, n);
for (idx, (v1, v2)) in recon.iter().zip(a.iter()).enumerate() {
assert!(
(v1 - v2).abs() < 1e-6,
"SVD recon[{idx}] = {} != A[{idx}] = {}",
v1,
v2
);
}
}
#[test]
fn test_svd_rank_deficient_small_sigma() {
let a = vec![1.0_f64, 2.0, 3.0, 2.0, 4.0, 6.0, 3.0, 6.0, 9.0];
let m = 3;
let n = 3;
let svd_r = svd(&a, m, n).unwrap();
assert!(svd_r.s[1].abs() < 1e-8, "s[1] = {}", svd_r.s[1]);
assert!(svd_r.s[2].abs() < 1e-8, "s[2] = {}", svd_r.s[2]);
}
#[test]
fn test_qr_dimension_mismatch() {
let a = vec![1.0_f64; 5];
assert!(qr(&a, 3, 2).is_err());
}
#[test]
fn test_svd_dimension_mismatch() {
let a = vec![1.0_f64; 5];
assert!(svd(&a, 3, 2).is_err());
}
}