1use crate::error::{Error, Result};
11use crate::linalg::{cpow, eigvals, lstsq, svd_s_vh};
12use num_complex::Complex;
13
14type C64 = Complex<f64>;
15
16#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
18pub enum ErrType {
19 #[default]
21 Abs,
22 Rel,
24}
25
26#[derive(Debug, Clone, PartialEq)]
28pub struct EspritParams {
29 pub x_min: f64,
31 pub x_max: f64,
33 pub err: Option<f64>,
37 pub err_type: ErrType,
39 pub m: Option<usize>,
41 pub lfactor: f64,
43 pub tol: f64,
46 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#[derive(Debug, Clone)]
75pub struct Esprit {
76 pub n: usize,
78 pub dim: usize,
80 pub l: usize,
82 pub x_min: f64,
84 pub x_max: f64,
86 pub s: Vec<f64>,
88 pub m: usize,
90 pub sigma: f64,
92 pub gamma: Vec<C64>,
94 pub omega: Vec<C64>,
96 pub err_max: f64,
98 pub err_ave: f64,
100 data_type: DataType,
101}
102
103impl Esprit {
104 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 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 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 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 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 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 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
296fn 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
305pub(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;