use crate::error::{CsError, CsResult};
#[derive(Debug, Clone)]
pub struct ListaConfig {
pub n_measurements: usize,
pub n_atoms: usize,
pub threshold: f64,
pub n_layers: usize,
}
#[derive(Debug, Clone)]
pub struct Lista {
w: Vec<f64>,
s: Vec<f64>,
config: ListaConfig,
}
impl Lista {
pub fn from_dict(dict: &[f64], config: ListaConfig) -> CsResult<Self> {
let m = config.n_atoms;
let n = config.n_measurements;
if dict.len() != m * n {
return Err(CsError::ShapeMismatch {
expected: vec![m, n],
got: vec![dict.len()],
});
}
if m == 0 || n == 0 {
return Err(CsError::InvalidParameter(
"n_atoms and n_measurements must be > 0".into(),
));
}
let l: f64 = dict.iter().map(|&v| v * v).sum();
let inv_l = if l > 0.0 { 1.0 / l } else { 1.0 };
let w: Vec<f64> = dict.iter().map(|&v| inv_l * v).collect();
let mut s = vec![0.0_f64; m * m];
for i in 0..m {
for k in 0..m {
let mut dot = 0.0_f64;
for j in 0..n {
dot += dict[i * n + j] * dict[k * n + j];
}
let identity_val = if i == k { 1.0 } else { 0.0 };
s[i * m + k] = identity_val - inv_l * dot;
}
}
Ok(Self { w, s, config })
}
pub fn encode(&self, y: &[f64]) -> CsResult<Vec<f64>> {
let m = self.config.n_atoms;
let n = self.config.n_measurements;
if y.len() != n {
return Err(CsError::DimensionMismatch { a: y.len(), b: n });
}
let lam = self.config.threshold;
let t = self.config.n_layers;
let mut x = vec![0.0_f64; m];
let wy = matvec(&self.w, y, m, n);
for _ in 0..t {
let sx = matvec(&self.s, &x, m, m);
for i in 0..m {
x[i] = Self::soft_threshold(wy[i] + sx[i], lam);
}
}
Ok(x)
}
pub fn encode_batch(&self, ys: &[f64], n_batch: usize) -> CsResult<Vec<f64>> {
let n = self.config.n_measurements;
let m = self.config.n_atoms;
if ys.len() != n_batch * n {
return Err(CsError::ShapeMismatch {
expected: vec![n_batch, n],
got: vec![ys.len()],
});
}
let mut out = Vec::with_capacity(n_batch * m);
for b in 0..n_batch {
let y_b = &ys[b * n..(b + 1) * n];
let code = self.encode(y_b)?;
out.extend_from_slice(&code);
}
Ok(out)
}
#[inline]
#[must_use]
pub fn soft_threshold(v: f64, lambda: f64) -> f64 {
if v > lambda {
v - lambda
} else if v < -lambda {
v + lambda
} else {
0.0
}
}
pub fn reconstruct(dict: &[f64], code: &[f64], n: usize, m: usize) -> CsResult<Vec<f64>> {
if dict.len() != m * n {
return Err(CsError::ShapeMismatch {
expected: vec![m, n],
got: vec![dict.len()],
});
}
if code.len() != m {
return Err(CsError::DimensionMismatch {
a: code.len(),
b: m,
});
}
let mut y_hat = vec![0.0_f64; n];
for i in 0..m {
for j in 0..n {
y_hat[j] += dict[i * n + j] * code[i];
}
}
Ok(y_hat)
}
#[must_use]
pub fn n_atoms(&self) -> usize {
self.config.n_atoms
}
#[must_use]
pub fn n_measurements(&self) -> usize {
self.config.n_measurements
}
#[must_use]
pub fn threshold(&self) -> f64 {
self.config.threshold
}
}
fn matvec(a: &[f64], x: &[f64], out_dim: usize, in_dim: usize) -> Vec<f64> {
let mut y = vec![0.0_f64; out_dim];
for i in 0..out_dim {
let mut s = 0.0_f64;
for j in 0..in_dim {
s += a[i * in_dim + j] * x[j];
}
y[i] = s;
}
y
}
#[cfg(test)]
mod tests {
use super::*;
fn make_dict(m: usize, n: usize) -> Vec<f64> {
(0..m * n)
.map(|k| {
let i = k / n;
let j = k % n;
if (i + j) % 3 == 0 { 0.5 } else { -0.1 }
})
.collect()
}
fn make_config(n: usize, m: usize, lambda: f64, layers: usize) -> ListaConfig {
ListaConfig {
n_measurements: n,
n_atoms: m,
threshold: lambda,
n_layers: layers,
}
}
#[test]
fn encode_output_shape() {
let n = 8;
let m = 12;
let dict = make_dict(m, n);
let cfg = make_config(n, m, 0.1, 10);
let lista = Lista::from_dict(&dict, cfg).expect("ok");
let y = vec![0.5_f64; n];
let code = lista.encode(&y).expect("ok");
assert_eq!(code.len(), m);
}
#[test]
fn encode_sparse() {
let n = 8;
let m = 12;
let dict = make_dict(m, n);
let cfg = make_config(n, m, 0.5, 20);
let lista = Lista::from_dict(&dict, cfg).expect("ok");
let y = vec![0.3_f64; n];
let code = lista.encode(&y).expect("ok");
let nnz = code.iter().filter(|&&v| v.abs() > 1e-10).count();
assert!(nnz < m, "expected sparse code, got {nnz}/{m} nonzeros");
}
#[test]
fn soft_threshold_positive() {
let v = Lista::soft_threshold(1.5, 0.5);
assert!((v - 1.0).abs() < 1.0e-12);
}
#[test]
fn soft_threshold_negative() {
let v = Lista::soft_threshold(-1.5, 0.5);
assert!((v + 1.0).abs() < 1.0e-12);
}
#[test]
fn soft_threshold_within_band() {
let v = Lista::soft_threshold(0.3, 0.5);
assert!(v.abs() < 1.0e-12);
}
#[test]
fn encode_all_finite() {
let n = 10;
let m = 15;
let dict = make_dict(m, n);
let cfg = make_config(n, m, 0.05, 15);
let lista = Lista::from_dict(&dict, cfg).expect("ok");
let y: Vec<f64> = (0..n).map(|i| (i as f64) * 0.1 - 0.5).collect();
let code = lista.encode(&y).expect("ok");
assert!(
code.iter().all(|v| v.is_finite()),
"code has non-finite value"
);
}
#[test]
fn reconstruct_shape() {
let n = 8;
let m = 12;
let dict = make_dict(m, n);
let code = vec![0.1_f64; m];
let y_hat = Lista::reconstruct(&dict, &code, n, m).expect("ok");
assert_eq!(y_hat.len(), n);
}
#[test]
fn encode_batch_shape() {
let n = 8;
let m = 12;
let n_batch = 5;
let dict = make_dict(m, n);
let cfg = make_config(n, m, 0.1, 10);
let lista = Lista::from_dict(&dict, cfg).expect("ok");
let ys = vec![0.4_f64; n_batch * n];
let codes = lista.encode_batch(&ys, n_batch).expect("ok");
assert_eq!(codes.len(), n_batch * m);
}
#[test]
fn larger_lambda_sparser() {
let n = 8;
let m = 12;
let dict = make_dict(m, n);
let y = vec![0.5_f64; n];
let cfg_small = make_config(n, m, 0.05, 10);
let lista_small = Lista::from_dict(&dict, cfg_small).expect("ok");
let code_small = lista_small.encode(&y).expect("ok");
let cfg_large = make_config(n, m, 0.5, 10);
let lista_large = Lista::from_dict(&dict, cfg_large).expect("ok");
let code_large = lista_large.encode(&y).expect("ok");
let nnz_small = code_small.iter().filter(|&&v| v.abs() > 1e-10).count();
let nnz_large = code_large.iter().filter(|&&v| v.abs() > 1e-10).count();
assert!(
nnz_large <= nnz_small,
"larger lambda should produce sparser code: nnz_small={nnz_small}, nnz_large={nnz_large}"
);
}
#[test]
fn from_dict_dim_mismatch() {
let n = 8;
let m = 12;
let bad_dict = vec![0.0_f64; m * n + 1];
let cfg = make_config(n, m, 0.1, 10);
let result = Lista::from_dict(&bad_dict, cfg);
assert!(result.is_err(), "expected Err for dim mismatch");
}
#[test]
fn zero_layers_returns_zeros() {
let n = 8;
let m = 12;
let dict = make_dict(m, n);
let cfg = make_config(n, m, 0.1, 0);
let lista = Lista::from_dict(&dict, cfg).expect("ok");
let y = vec![1.0_f64; n];
let code = lista.encode(&y).expect("ok");
assert!(
code.iter().all(|&v| v == 0.0),
"zero layers must return all zeros"
);
}
}