Skip to main content

sparse_ir_minipole/
esprit.rs

1//! Matrix ESPRIT (port of `mini_pole/esprit.py`).
2//!
3//! Approximates `h_k ∈ C^d`, sampled at `N` uniformly spaced points
4//! `x_k` of `[x_min, x_max]`, by `Σ_j ω_j γ_j^k` with nodes `γ_j` shared by
5//! the `d` columns.
6//!
7//! Ported from Green-Phys/MiniPole (commit 15e4a54, MIT License,
8//! Copyright (c) 2024 lzphy); see `LICENSE-THIRD-PARTY`.
9
10use crate::error::{Error, Result};
11use crate::linalg::{cpow, eigvals, lstsq, svd_s_vh};
12use num_complex::Complex;
13
14type C64 = Complex<f64>;
15
16/// Interpretation of an error tolerance.
17#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
18pub enum ErrType {
19    /// Absolute error.
20    #[default]
21    Abs,
22    /// Error relative to the largest singular value.
23    Rel,
24}
25
26/// Parameters of [`Esprit::new`], with the defaults of the reference.
27#[derive(Debug, Clone, PartialEq)]
28pub struct EspritParams {
29    /// Lower end of the sampling interval (default 0).
30    pub x_min: f64,
31    /// Upper end of the sampling interval (default 1).
32    pub x_max: f64,
33    /// Error tolerance that selects the number of nodes `M`. One of `err`
34    /// and `m` is required: the reference's knee detection for neither is
35    /// not ported.
36    pub err: Option<f64>,
37    /// Interpretation of `err` (default absolute).
38    pub err_type: ErrType,
39    /// Number of nodes `M`; overrides `err`.
40    pub m: Option<usize>,
41    /// Ratio `L / (N - 1)` of the Hankel matrix (default 0.4).
42    pub lfactor: f64,
43    /// Threshold below which the imaginary (real) part of the input counts
44    /// as zero (default 1e-15).
45    pub tol: f64,
46    /// The approximation is accepted when its maximum error is below
47    /// `ctrl_ratio` times the first discarded singular value (default 10).
48    pub ctrl_ratio: f64,
49}
50
51impl Default for EspritParams {
52    fn default() -> Self {
53        Self {
54            x_min: 0.0,
55            x_max: 1.0,
56            err: None,
57            err_type: ErrType::Abs,
58            m: None,
59            lfactor: 0.4,
60            tol: 1e-15,
61            ctrl_ratio: 10.0,
62        }
63    }
64}
65
66#[derive(Debug, Clone, Copy, PartialEq, Eq)]
67enum DataType {
68    Real,
69    Imag,
70    Cplx,
71}
72
73/// Result of the ESPRIT approximation.
74#[derive(Debug, Clone)]
75pub struct Esprit {
76    /// Number of samples `N`.
77    pub n: usize,
78    /// Number of columns `d`.
79    pub dim: usize,
80    /// Hankel parameter `L`.
81    pub l: usize,
82    /// Sampling interval.
83    pub x_min: f64,
84    /// Sampling interval.
85    pub x_max: f64,
86    /// Singular values of the Hankel matrix (empty for zero input).
87    pub s: Vec<f64>,
88    /// Number of nodes `M`.
89    pub m: usize,
90    /// `S[M]`, the first discarded singular value (0 for zero input).
91    pub sigma: f64,
92    /// Nodes `γ_j`.
93    pub gamma: Vec<C64>,
94    /// Weights `ω_j`, column-major `M x d`.
95    pub omega: Vec<C64>,
96    /// Maximum error of the approximation on the samples.
97    pub err_max: f64,
98    /// Largest mean (over samples) error among the columns.
99    pub err_ave: f64,
100    data_type: DataType,
101}
102
103impl Esprit {
104    /// ESPRIT of `h` (column-major `n x dim`).
105    ///
106    /// # Errors
107    /// [`Error::InvalidParameter`] for invalid parameters, and
108    /// [`Error::DecompositionFailed`] if no controlled approximation exists.
109    pub fn new(h: &[C64], n: usize, dim: usize, p: &EspritParams) -> Result<Self> {
110        if h.len() != n * dim || dim == 0 {
111            return Err(Error::InvalidParameter {
112                name: "h_k",
113                value: format!("of length {} for N = {n}, d = {dim}", h.len()),
114                reason: "must have N d entries and d >= 1".to_string(),
115            });
116        }
117        let mut l = (p.lfactor * (n as f64 - 1.0)) as usize;
118        if n < 2 || n - l < l + 1 {
119            return Err(Error::InvalidParameter {
120                name: "lfactor",
121                value: format!("{} for N = {n}", p.lfactor),
122                reason: "must satisfy N - L >= L + 1 with L = int(lfactor (N - 1))".to_string(),
123            });
124        }
125        if p.x_min.is_nan() || p.x_max.is_nan() || p.x_min >= p.x_max {
126            return Err(Error::InvalidParameter {
127                name: "x_min, x_max",
128                value: format!("[{:?}, {:?}]", p.x_min, p.x_max),
129                reason: "must satisfy x_min < x_max".to_string(),
130            });
131        }
132        let max_im = h.iter().map(|v| v.im.abs()).fold(0.0, f64::max);
133        let max_re = h.iter().map(|v| v.re.abs()).fold(0.0, f64::max);
134        let (data_type, h): (DataType, Vec<C64>) = if max_im < p.tol {
135            (
136                DataType::Real,
137                h.iter().map(|v| C64::new(v.re, 0.0)).collect(),
138            )
139        } else if max_re < p.tol {
140            (
141                DataType::Imag,
142                h.iter().map(|v| C64::new(0.0, v.im)).collect(),
143            )
144        } else {
145            (DataType::Cplx, h.to_vec())
146        };
147        let mut out = Self {
148            n,
149            dim,
150            l,
151            x_min: p.x_min,
152            x_max: p.x_max,
153            s: Vec::new(),
154            m: 0,
155            sigma: 0.0,
156            gamma: Vec::new(),
157            omega: Vec::new(),
158            err_max: 0.0,
159            err_ave: 0.0,
160            data_type,
161        };
162        if h.iter().all(|v| *v == C64::new(0.0, 0.0)) {
163            return Ok(out);
164        }
165
166        // H[(d l + c), j] = h[l + j, c], of size d (N - L) x (L + 1); shrink L
167        // if the SVD fails.
168        let (s, vh) = loop {
169            let rows = dim * (n - l);
170            let cols = l + 1;
171            let mut hm = vec![C64::new(0.0, 0.0); rows * cols];
172            for ll in 0..(n - l) {
173                for c in 0..dim {
174                    for j in 0..cols {
175                        hm[(dim * ll + c) + rows * j] = h[(ll + j) + n * c];
176                    }
177                }
178            }
179            match svd_s_vh(&hm, rows, cols) {
180                Ok(f) => break f,
181                Err(e) if l <= 1 => return Err(e),
182                Err(_) => l -= 1,
183            }
184        };
185        out.l = l;
186        let k = s.len();
187
188        let mut m = match p.m {
189            Some(m) => m.min(k - 1),
190            None => {
191                let err = p.err.ok_or_else(|| Error::InvalidParameter {
192                    name: "err, M",
193                    value: "None".to_string(),
194                    reason:
195                        "one of them is required (the knee detection for neither is not ported)"
196                            .to_string(),
197                })?;
198                let mut m = find_m_with_err(&s, err, p.err_type);
199                if s[m] / s[0] < 1e-14 {
200                    m = find_m_with_err(&s, 1e-14, ErrType::Rel);
201                }
202                m
203            }
204        };
205        let x_k: Vec<f64> = linspace(p.x_min, p.x_max, n);
206        loop {
207            out.sigma = s[m];
208            // F = lstsq(W_0ᵀ, W_1ᵀ), W_0 = Vh[:M, :-1], W_1 = Vh[:M, 1:].
209            let mut w0t = vec![C64::new(0.0, 0.0); l * m];
210            let mut w1t = vec![C64::new(0.0, 0.0); l * m];
211            for i in 0..m {
212                for j in 0..l {
213                    w0t[j + l * i] = vh[i + k * j];
214                    w1t[j + l * i] = vh[i + k * (j + 1)];
215                }
216            }
217            let f = lstsq(&w0t, l, m, &w1t, m)?;
218            out.gamma = eigvals(&f, m)?;
219            out.m = m;
220            // omega = lstsq(V, h), V[i, j] = γ_j^i.
221            let mut v = vec![C64::new(0.0, 0.0); n * m];
222            for (j, &g) in out.gamma.iter().enumerate() {
223                for i in 0..n {
224                    v[i + n * j] = cpow(g, i as f64);
225                }
226            }
227            out.omega = lstsq(&v, n, m, &h, dim)?;
228            // cal_err
229            let approx = out.get_value(&x_k);
230            let (mut err_max, mut err_ave) = (0.0_f64, 0.0_f64);
231            for c in 0..dim {
232                let mut sum = 0.0;
233                for i in 0..n {
234                    let e = (approx[i + n * c] - h[i + n * c]).norm();
235                    err_max = err_max.max(e);
236                    sum += e;
237                }
238                err_ave = err_ave.max(sum / n as f64);
239            }
240            out.err_max = err_max;
241            out.err_ave = err_ave;
242            if err_max < (p.ctrl_ratio * out.sigma).max(1e-14 * s[0]) {
243                break;
244            }
245            m = m.wrapping_sub(1);
246            if m == 0 || m == usize::MAX {
247                return Err(Error::DecompositionFailed {
248                    reason: "ESPRIT could not find a controlled approximation".to_string(),
249                });
250            }
251        }
252        out.s = s;
253        Ok(out)
254    }
255
256    /// The approximation at points `x` (column-major `len(x) x d`).
257    pub fn get_value(&self, x: &[f64]) -> Vec<C64> {
258        let nx = x.len();
259        let mut out = vec![C64::new(0.0, 0.0); nx * self.dim];
260        for (i, &xi) in x.iter().enumerate() {
261            let e = (self.n as f64 - 1.0) * ((xi - self.x_min) / (self.x_max - self.x_min));
262            for (j, &g) in self.gamma.iter().enumerate() {
263                let vj = cpow(g, e);
264                for c in 0..self.dim {
265                    out[i + nx * c] += vj * self.omega[j + self.m * c];
266                }
267            }
268        }
269        for v in &mut out {
270            *v = self.project(*v);
271        }
272        out
273    }
274
275    /// The approximation of column `col` at `x`.
276    pub fn get_value_indiv(&self, x: f64, col: usize) -> C64 {
277        let e = (self.n as f64 - 1.0) * ((x - self.x_min) / (self.x_max - self.x_min));
278        let v: C64 = self
279            .gamma
280            .iter()
281            .enumerate()
282            .map(|(j, &g)| cpow(g, e) * self.omega[j + self.m * col])
283            .sum();
284        self.project(v)
285    }
286
287    fn project(&self, v: C64) -> C64 {
288        match self.data_type {
289            DataType::Cplx => v,
290            DataType::Real => C64::new(v.re, 0.0),
291            DataType::Imag => C64::new(0.0, v.im),
292        }
293    }
294}
295
296/// First index with `S[idx] < cutoff`, or the last index.
297fn find_m_with_err(s: &[f64], err: f64, err_type: ErrType) -> usize {
298    let cutoff = match err_type {
299        ErrType::Abs => err,
300        ErrType::Rel => s[0] * err,
301    };
302    s.iter().position(|&v| v < cutoff).unwrap_or(s.len() - 1)
303}
304
305/// `np.linspace(a, b, n)`.
306pub(crate) fn linspace(a: f64, b: f64, n: usize) -> Vec<f64> {
307    if n == 1 {
308        return vec![a];
309    }
310    let step = (b - a) / (n as f64 - 1.0);
311    let mut x: Vec<f64> = (0..n).map(|i| a + i as f64 * step).collect();
312    x[n - 1] = b;
313    x
314}
315
316#[cfg(test)]
317#[path = "esprit_tests.rs"]
318mod tests;