use crate::error::{CsError, CsResult};
use crate::linalg::cholesky::{cholesky_factor, cholesky_solve};
use crate::linalg::{mat_t_vec, mat_vec, norm2};
use crate::thresholding::iht::soft_threshold;
#[derive(Debug, Clone)]
pub struct BasisPursuitResult {
pub x: Vec<f64>,
pub primal_residual: f64,
pub dual_residual: f64,
pub iterations: usize,
}
pub fn basis_pursuit(
phi: &[f64],
m: usize,
n: usize,
y: &[f64],
rho: f64,
max_iter: usize,
tol: f64,
) -> CsResult<BasisPursuitResult> {
if phi.len() != m * n {
return Err(CsError::ShapeMismatch {
expected: vec![m, n],
got: vec![phi.len()],
});
}
if y.len() != m {
return Err(CsError::DimensionMismatch { a: y.len(), b: m });
}
if rho <= 0.0 {
return Err(CsError::InvalidParameter("rho must be > 0".into()));
}
let mut pp = vec![0.0_f64; m * m];
for i in 0..m {
for k in 0..n {
let pik = phi[i * n + k];
for j in 0..m {
pp[i * m + j] += pik * phi[j * n + k];
}
}
}
for i in 0..m {
pp[i * m + i] += 1.0e-10;
}
let l = cholesky_factor(&pp, m)?;
let mut x = vec![0.0_f64; n];
let mut z = vec![0.0_f64; n];
let mut u = vec![0.0_f64; n];
let mut iter = 0usize;
let mut primal_r = f64::INFINITY;
let mut dual_r = f64::INFINITY;
for _ in 0..max_iter {
let mut zmu = vec![0.0_f64; n];
for j in 0..n {
zmu[j] = z[j] - u[j];
}
let phi_zmu = mat_vec(phi, m, n, &zmu)?;
let mut rhs = vec![0.0_f64; m];
for i in 0..m {
rhs[i] = y[i] - phi_zmu[i];
}
let lam = cholesky_solve(&l, m, &rhs)?;
let phi_t_lam = mat_t_vec(phi, m, n, &lam)?;
for j in 0..n {
x[j] = zmu[j] + phi_t_lam[j];
}
let mut x_plus_u = vec![0.0_f64; n];
for j in 0..n {
x_plus_u[j] = x[j] + u[j];
}
let z_new = soft_threshold(&x_plus_u, 1.0 / rho);
let mut dr_sq = 0.0_f64;
for j in 0..n {
let d = z_new[j] - z[j];
dr_sq += d * d;
}
dual_r = rho * dr_sq.sqrt();
z = z_new;
for j in 0..n {
u[j] += x[j] - z[j];
}
let mut pr_sq = 0.0_f64;
for j in 0..n {
let d = x[j] - z[j];
pr_sq += d * d;
}
primal_r = pr_sq.sqrt();
iter += 1;
if primal_r < tol && dual_r < tol {
break;
}
}
Ok(BasisPursuitResult {
x: z,
primal_residual: primal_r,
dual_residual: dual_r,
iterations: iter,
})
}
pub fn basis_pursuit_denoise(
phi: &[f64],
m: usize,
n: usize,
y: &[f64],
eps: f64,
rho: f64,
max_iter: usize,
tol: f64,
) -> CsResult<BasisPursuitResult> {
if phi.len() != m * n {
return Err(CsError::ShapeMismatch {
expected: vec![m, n],
got: vec![phi.len()],
});
}
if y.len() != m {
return Err(CsError::DimensionMismatch { a: y.len(), b: m });
}
if eps < 0.0 {
return Err(CsError::InvalidParameter("eps must be ≥ 0".into()));
}
if rho <= 0.0 {
return Err(CsError::InvalidParameter("rho must be > 0".into()));
}
let mut g = vec![0.0_f64; n * n];
for k in 0..m {
for i in 0..n {
let pki = phi[k * n + i];
for j in 0..n {
g[i * n + j] += pki * phi[k * n + j];
}
}
}
for i in 0..n {
g[i * n + i] += rho;
}
let l = cholesky_factor(&g, n)?;
let phi_t_y = mat_t_vec(phi, m, n, y)?;
let mut x = vec![0.0_f64; n];
let mut z = vec![0.0_f64; n];
let mut u = vec![0.0_f64; n];
let mut iter = 0usize;
let lambda = 1.0; let mut primal_r = f64::INFINITY;
let mut dual_r = f64::INFINITY;
for _ in 0..max_iter {
let mut rhs = vec![0.0_f64; n];
for j in 0..n {
rhs[j] = phi_t_y[j] + rho * (z[j] - u[j]);
}
x = cholesky_solve(&l, n, &rhs)?;
let mut x_plus_u = vec![0.0_f64; n];
for j in 0..n {
x_plus_u[j] = x[j] + u[j];
}
let z_new = soft_threshold(&x_plus_u, lambda / rho);
let mut dr_sq = 0.0_f64;
for j in 0..n {
let d = z_new[j] - z[j];
dr_sq += d * d;
}
dual_r = rho * dr_sq.sqrt();
z = z_new;
for j in 0..n {
u[j] += x[j] - z[j];
}
let mut pr_sq = 0.0_f64;
for j in 0..n {
let d = x[j] - z[j];
pr_sq += d * d;
}
primal_r = pr_sq.sqrt();
iter += 1;
let ax = mat_vec(phi, m, n, &z)?;
let mut res = vec![0.0_f64; m];
for i in 0..m {
res[i] = ax[i] - y[i];
}
let res_norm = norm2(&res);
if primal_r < tol && dual_r < tol && res_norm <= eps + tol {
break;
}
}
Ok(BasisPursuitResult {
x: z,
primal_residual: primal_r,
dual_residual: dual_r,
iterations: iter,
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn basis_pursuit_canonical() {
let phi = vec![1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0, 0.0, 0.0, 0.0, 1.0, 0.0];
let y = vec![1.0, 0.0, 0.5];
let r = basis_pursuit(&phi, 3, 4, &y, 1.0, 200, 1.0e-6).expect("ok");
assert!((r.x[0] - 1.0).abs() < 1.0e-3);
assert!((r.x[2] - 0.5).abs() < 1.0e-3);
}
#[test]
fn bpdn_noisy() {
let phi = vec![1.0, 0.0, 0.0, 1.0];
let y = vec![1.0, 0.5];
let r = basis_pursuit_denoise(&phi, 2, 2, &y, 0.1, 1.0, 100, 1.0e-6).expect("ok");
assert!(r.iterations > 0);
}
}