use crate::error::{CsError, CsResult};
use crate::linalg::jacobi_svd::jacobi_svd_thin;
use crate::matrix_completion::CompletionResult;
pub fn svt(
m: &[f64],
mask: &[bool],
h: usize,
w: usize,
tau: f64,
delta: f64,
max_iter: usize,
tol: f64,
) -> CsResult<CompletionResult> {
if m.len() != h * w {
return Err(CsError::ShapeMismatch {
expected: vec![h, w],
got: vec![m.len()],
});
}
if mask.len() != h * w {
return Err(CsError::DimensionMismatch {
a: mask.len(),
b: h * w,
});
}
if tau <= 0.0 {
return Err(CsError::InvalidParameter("tau must be > 0".into()));
}
if delta <= 0.0 {
return Err(CsError::InvalidParameter("delta must be > 0".into()));
}
let mut y = vec![0.0_f64; h * w];
let mut x = vec![0.0_f64; h * w];
let mut iter = 0usize;
let mut last_residual = f64::INFINITY;
let r_dim = h.min(w);
for _ in 0..max_iter {
let (u, s, v) = if h >= w {
jacobi_svd_thin(&y, h, w)?
} else {
let mut yt = vec![0.0_f64; w * h];
for i in 0..h {
for j in 0..w {
yt[j * h + i] = y[i * w + j];
}
}
let (u2, s2, v2) = jacobi_svd_thin(&yt, w, h)?;
(v2, s2, u2)
};
let mut s_new = vec![0.0_f64; s.len()];
for (i, &si) in s.iter().enumerate() {
s_new[i] = (si - tau).max(0.0);
}
x.fill(0.0);
if h >= w {
for i in 0..h {
for j in 0..w {
let mut acc = 0.0_f64;
for k in 0..w {
acc += u[i * w + k] * s_new[k] * v[j * w + k];
}
x[i * w + j] = acc;
}
}
} else {
for i in 0..h {
for j in 0..w {
let mut acc = 0.0_f64;
for k in 0..r_dim {
acc += u[i * r_dim + k] * s_new[k] * v[j * r_dim + k];
}
x[i * w + j] = acc;
}
}
}
let mut residual_sq = 0.0_f64;
let mut obs_norm_sq = 0.0_f64;
for k in 0..(h * w) {
if mask[k] {
let r = m[k] - x[k];
y[k] += delta * r;
residual_sq += r * r;
obs_norm_sq += m[k] * m[k];
}
}
iter += 1;
let denom = obs_norm_sq.sqrt().max(1.0e-300);
let cur_res = residual_sq.sqrt() / denom;
if (last_residual - cur_res).abs() < tol && cur_res < tol {
break;
}
last_residual = cur_res;
}
Ok(CompletionResult {
x,
residual: last_residual,
iterations: iter,
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn svt_recovers_low_rank_2x2() {
let m = vec![1.0_f64, 2.0, 2.0, 4.0];
let mask = vec![true, true, true, false];
let r = svt(&m, &mask, 2, 2, 0.5, 1.5, 1000, 1.0e-9).expect("ok");
assert!(r.iterations > 0);
assert!((r.x[0] - 1.0).abs() < 1.0);
assert!((r.x[1] - 2.0).abs() < 1.0);
assert!((r.x[2] - 2.0).abs() < 1.0);
}
}