1use super::con_map::{ConMap, ConMapGapless, ConMapGeneric};
8use super::quad::oscillatory;
9use super::{MiniPoleResult, assemble};
10use crate::error::{ArrayRole, Error, Result};
11use crate::esprit::{ErrType, Esprit, EspritParams, linspace};
12use crate::linalg::lstsq;
13use num_complex::Complex;
14use std::f64::consts::PI;
15use tenferro_tensor::TypedTensor;
16
17type C64 = Complex<f64>;
18
19#[derive(Debug, Clone, Copy, PartialEq, Eq)]
21pub enum N0 {
22 Auto {
24 shift: usize,
26 },
27 Fixed(usize),
29}
30
31#[derive(Debug, Clone, Copy, PartialEq, Eq)]
33pub enum Plane {
34 Z,
36 W,
38}
39
40#[derive(Debug, Clone, PartialEq)]
42pub struct MiniPoleParams {
43 pub n0: N0,
45 pub err: f64,
48 pub err_type: ErrType,
50 pub m: Option<usize>,
52 pub symmetry: bool,
54 pub g_symmetric: bool,
56 pub compute_const: bool,
58 pub plane: Option<Plane>,
61 pub include_n0: bool,
63 pub k_max: usize,
65 pub ratio_max: f64,
67}
68
69impl MiniPoleParams {
70 pub fn new(err: f64) -> Self {
72 Self {
73 n0: N0::Auto { shift: 0 },
74 err,
75 err_type: ErrType::Abs,
76 m: None,
77 symmetry: false,
78 g_symmetric: false,
79 compute_const: false,
80 plane: None,
81 include_n0: false,
82 k_max: 999,
83 ratio_max: 10.0,
84 }
85 }
86}
87
88pub fn mini_pole(
99 g_w: &TypedTensor<C64>,
100 w: &[f64],
101 params: &MiniPoleParams,
102) -> Result<MiniPoleResult> {
103 let shape = g_w.shape().to_vec();
104 let nw = w.len();
105 let n_orb = match shape.len() {
106 1 => 1,
107 3 if shape[1] == shape[2] => shape[1],
108 _ => {
109 return Err(Error::ShapeMismatch {
110 which: ArrayRole::Input,
111 expected: vec![
112 nw,
113 shape.get(1).copied().unwrap_or(1),
114 shape.get(1).copied().unwrap_or(1),
115 ],
116 actual: shape,
117 });
118 }
119 };
120 if shape[0] != nw {
121 let mut expected = shape.clone();
122 expected[0] = nw;
123 return Err(Error::ShapeMismatch {
124 which: ArrayRole::Input,
125 expected,
126 actual: shape,
127 });
128 }
129 if nw < 3
130 || w.iter().any(|x| !x.is_finite() || *x < 0.0)
131 || w.windows(2).any(|pair| pair[1] <= pair[0])
132 {
133 return Err(Error::InvalidParameter {
134 name: "w",
135 value: format!("of length {nw} starting at {:?}", w.first()),
136 reason: "must have at least 3 finite, non-negative, strictly increasing frequencies"
137 .to_string(),
138 });
139 }
140 let wabs = w.iter().map(|x| x.abs()).fold(0.0, f64::max);
141 let dd = (0..nw - 2)
142 .map(|i| ((w[i + 2] - 2.0 * w[i + 1] + w[i]) / wabs).abs())
143 .fold(0.0, f64::max);
144 if dd.is_nan() || dd >= 1e-6 {
145 return Err(Error::InvalidParameter {
146 name: "w",
147 value: format!("with second differences up to {dd:e}"),
148 reason: "must be uniformly spaced".to_string(),
149 });
150 }
151 if params.symmetry && params.compute_const {
152 return Err(Error::InvalidParameter {
153 name: "compute_const",
154 value: "true".to_string(),
155 reason: "set symmetry to false to calculate the overall constant".to_string(),
156 });
157 }
158 let d = n_orb * n_orb;
159 let tr = |c: usize| (c % n_orb) * n_orb + c / n_orb;
160 let raw = g_w.host_data()?;
161 let g: Vec<C64> = if params.g_symmetric {
162 (0..nw * d)
163 .map(|idx| {
164 let (i, c) = (idx % nw, idx / nw);
165 (raw[idx] + raw[i + nw * tr(c)]) * 0.5
166 })
167 .collect()
168 } else {
169 raw.to_vec()
170 };
171 let plane = params
172 .plane
173 .unwrap_or(if params.symmetry { Plane::W } else { Plane::Z });
174
175 let first = |lfactor: f64| -> Result<Vec<Esprit>> {
177 (0..d)
178 .map(|c| {
179 Esprit::new(
180 &g[nw * c..nw * (c + 1)],
181 nw,
182 1,
183 &EspritParams {
184 x_min: w[0],
185 x_max: w[nw - 1],
186 err: Some(params.err),
187 err_type: params.err_type,
188 lfactor,
189 ..EspritParams::default()
190 },
191 )
192 })
193 .collect()
194 };
195 let p_o = first(0.4)?;
196 let (n0, err_max) = match params.n0 {
197 N0::Auto { shift } => {
198 let p_o2 = first(0.5)?;
199 let w_cont = linspace(w[0], w[nw - 1], 10 * nw - 9);
200 let err_max = p_o
201 .iter()
202 .chain(&p_o2)
203 .map(|p| p.err_max)
204 .fold(f64::NEG_INFINITY, f64::max);
205 let mut n0 = 0;
206 for c in 0..d {
207 let l1 = p_o[c].get_value(&w_cont);
208 let l2 = p_o2[c].get_value(&w_cont);
209 let diff: Vec<f64> = (0..nw - 1)
210 .map(|j| {
211 (0..10)
212 .map(|t| (l2[10 * j + t] - l1[10 * j + t]).norm())
213 .fold(f64::NEG_INFINITY, f64::max)
214 })
215 .collect();
216 let first_ok = (0..nw - 2)
217 .position(|j| diff[j] <= err_max && diff[j] / diff[j + 1] < params.ratio_max)
218 .unwrap_or(0);
219 n0 = n0.max(first_ok);
220 }
221 (n0 + shift, err_max)
222 }
223 N0::Fixed(n0) => (
224 n0,
225 p_o.iter()
226 .map(|p| p.err_max)
227 .fold(f64::NEG_INFINITY, f64::max),
228 ),
229 };
230 if n0 + 1 >= nw {
231 return Err(Error::InvalidParameter {
232 name: "n0",
233 value: n0.to_string(),
234 reason: format!("must be less than len(w) - 1 = {}", nw - 1),
235 });
236 }
237 if params.symmetry && w[n0] <= 0.0 {
238 return Err(Error::InvalidParameter {
240 name: "n0",
241 value: n0.to_string(),
242 reason: format!("must give w[n0] > 0 with symmetry, got {:?}", w[n0]),
243 });
244 }
245 let head = |c: usize, x: f64| p_o[c].get_value_indiv(x, 0);
246 let cutoff = err_max;
247 let qerr = 0.01 * cutoff;
248
249 let generic;
250 let gapless;
251 let mut constant = vec![C64::new(0.0, 0.0); d];
252 let (map, h, nk): (&dyn ConMap, Vec<C64>, usize) = if !params.symmetry {
253 generic = ConMapGeneric {
254 w_m: 0.5 * (w[n0] + w[nw - 1]),
255 dw_h: 0.5 * (w[nw - 1] - w[n0]),
256 };
257 let (w_m, dw_h) = (generic.w_m, generic.dw_h);
258 let (h, nk) = moments(params.k_max, d, cutoff, |k, c| {
259 let f = |x: f64| head(c, w_m + dw_h * x.sin());
260 let v = oscillatory(&f, -0.5 * PI, 0.5 * PI, (k + 1) as f64, k & 1 == 0, qerr);
261 if k & 1 == 0 {
262 v * C64::new(0.0, 1.0 / PI)
263 } else {
264 v / PI
265 }
266 });
267 (&generic, h, nk)
268 } else {
269 let sub = mini_pole(
271 g_w,
272 w,
273 &MiniPoleParams {
274 m: None,
275 symmetry: false,
276 plane: None,
277 include_n0: false,
278 ..params.clone()
279 },
280 )?;
281 constant = sub.constant.clone();
282 let sub_loc = sub.pole_location.clone();
283 let sub_w = sub.pole_weight.host_data()?.to_vec();
284 let r_sub = sub_loc.len();
285 let tail = move |c: usize, x: f64| -> C64 {
286 let z = C64::new(0.0, x);
287 (0..r_sub)
288 .map(|j| sub_w[j + r_sub * c] / (z - sub_loc[j]))
289 .sum()
290 };
291 gapless = ConMapGapless { w_min: w[n0] };
292 let w_min = gapless.w_min;
293 let theta0 = (w_min / w[nw - 1]).asin();
294 let (ha, hb) = (theta0 + 1e-12, 0.5 * PI);
295 let (ta, tb) = (1e-6, theta0 - 1e-12);
296 let integral = |f: &dyn Fn(f64) -> C64, a: f64, b: f64, k: usize| {
297 oscillatory(f, a, b, (k + 1) as f64, k & 1 == 0, qerr)
298 };
299 let sym = |gf: &dyn Fn(f64) -> C64, k: usize, a: f64, b: f64| -> C64 {
301 if k & 1 == 0 {
302 let f = |x: f64| C64::new(gf(w_min / x.sin()).im, 0.0);
303 integral(&f, a, b, k) * (-2.0 / PI)
304 } else {
305 let f = |x: f64| C64::new(gf(w_min / x.sin()).re, 0.0);
306 integral(&f, a, b, k) * (2.0 / PI)
307 }
308 };
309 let raw_int = |gf: &dyn Fn(f64) -> C64, k: usize, a: f64, b: f64| -> C64 {
310 let f = |x: f64| gf(w_min / x.sin());
311 integral(&f, a, b, k)
312 };
313 let sym_hk = |c: usize, k: usize| -> C64 {
314 sym(&|x| head(c, x), k, ha, hb) + sym(&|x| tail(c, x), k, ta, tb)
315 };
316 let (h, nk) = if params.g_symmetric {
317 moments(params.k_max, d, cutoff, |k, c| sym_hk(c, k))
318 } else {
319 moments_rows(params.k_max, d, cutoff, |k| {
320 let mut row = vec![C64::new(0.0, 0.0); d];
321 for i in 0..n_orb {
322 for j in i..n_orb {
323 let c1 = i + n_orb * j;
325 let c2 = j + n_orb * i;
326 if i == j {
327 row[c1] = sym_hk(c1, k);
328 } else {
329 let h1 = raw_int(&|x| head(c1, x), k, ha, hb)
330 + raw_int(&|x| tail(c1, x), k, ta, tb);
331 let h2 = raw_int(&|x| head(c2, x), k, ha, hb)
332 + raw_int(&|x| tail(c2, x), k, ta, tb);
333 if k & 1 == 0 {
334 let f = C64::new(0.0, 1.0 / PI);
335 row[c1] = f * (h1 - h2.conj());
336 row[c2] = f * (h2 - h1.conj());
337 } else {
338 row[c1] = (h1 + h2.conj()) / PI;
339 row[c2] = (h2 + h1.conj()) / PI;
340 }
341 }
342 }
343 }
344 row
345 })
346 };
347 (&gapless, h, nk)
348 };
349
350 let esprit = Esprit::new(
352 &h,
353 nk,
354 d,
355 &EspritParams {
356 err: if params.m.is_none() {
357 Some(0.5 * err_max)
358 } else {
359 None
360 },
361 m: params.m,
362 lfactor: 0.5,
363 ..EspritParams::default()
364 },
365 )?;
366 let keep: Vec<usize> = (0..esprit.gamma.len())
367 .filter(|&j| esprit.gamma[j].norm() < 1.0)
368 .collect();
369 let r = keep.len();
370 let location: Vec<C64> = keep.iter().map(|&j| map.z(esprit.gamma[j])).collect();
371 let mut weight = vec![C64::new(0.0, 0.0); r * d];
372 for (jn, &j) in keep.iter().enumerate() {
373 let f = map.dz(esprit.gamma[j]);
374 for c in 0..d {
375 weight[jn + r * c] = esprit.omega[j + esprit.m * c] * f;
376 }
377 }
378 if params.compute_const {
379 let nf = nw - n0;
380 let mut cst = vec![C64::new(0.0, 0.0); d];
381 for (c, v) in cst.iter_mut().enumerate() {
382 for i in n0..nw {
383 let z = C64::new(0.0, w[i]);
384 let approx: C64 = (0..r).map(|j| weight[j + r * c] / (z - location[j])).sum();
385 *v += g[i + nw * c] - approx;
386 }
387 *v /= nf as f64;
388 }
389 let big = cst.iter().map(|v| v.norm()).fold(0.0, f64::max) > 100.0 * err_max;
390 constant = if big {
391 cst
392 } else {
393 vec![C64::new(0.0, 0.0); d]
394 };
395 }
396 if plane == Plane::Z {
397 let start = if params.include_n0 { 0 } else { n0 };
398 let mut ws: Vec<f64> = Vec::new();
399 let mut rows: Vec<(usize, bool)> = Vec::new(); if params.symmetry {
401 for i in (start..nw).rev() {
402 ws.push(-w[i]);
403 rows.push((i, true));
404 }
405 }
406 ws.extend_from_slice(&w[start..]);
407 rows.extend((start..nw).map(|i| (i, false)));
408 let nr = ws.len();
409 let mut a = vec![C64::new(0.0, 0.0); nr * r];
410 for (j, &loc) in location.iter().enumerate() {
411 for (i, &x) in ws.iter().enumerate() {
412 a[i + nr * j] = (C64::new(0.0, x) - loc).inv();
413 }
414 }
415 let mut b = vec![C64::new(0.0, 0.0); nr * d];
416 for c in 0..d {
417 for (i, &(fi, mirrored)) in rows.iter().enumerate() {
418 let v = if mirrored {
419 g[fi + nw * tr(c)].conj()
420 } else {
421 g[fi + nw * c]
422 };
423 b[i + nr * c] = v - constant[c];
424 }
425 }
426 weight = lstsq(&a, nr, r, &b, d)?;
427 }
428 let keep: Vec<usize> = (0..r)
430 .filter(|&j| (0..d).map(|c| weight[j + r * c].norm()).fold(0.0, f64::max) > err_max)
431 .collect();
432 let location_k: Vec<C64> = keep.iter().map(|&j| location[j]).collect();
433 let mut weight_k = vec![C64::new(0.0, 0.0); keep.len() * d];
434 for (jn, &j) in keep.iter().enumerate() {
435 for c in 0..d {
436 weight_k[jn + keep.len() * c] = weight[j + r * c];
437 }
438 }
439 let mut hshape = shape;
440 hshape[0] = nk;
441 assemble(
442 location_k,
443 weight_k,
444 constant,
445 h,
446 hshape,
447 esprit,
448 n0,
449 Some(err_max),
450 )
451}
452
453fn moments(
456 k_max: usize,
457 d: usize,
458 cutoff: f64,
459 f: impl Fn(usize, usize) -> C64,
460) -> (Vec<C64>, usize) {
461 moments_rows(k_max, d, cutoff, |k| (0..d).map(|c| f(k, c)).collect())
462}
463
464fn moments_rows(
465 k_max: usize,
466 d: usize,
467 cutoff: f64,
468 f: impl Fn(usize) -> Vec<C64>,
469) -> (Vec<C64>, usize) {
470 let mut rows: Vec<Vec<C64>> = Vec::new();
471 for k in 0..k_max {
472 rows.push(f(k));
473 if k >= 1 && (0..d).all(|c| rows[k][c].norm() < cutoff && rows[k - 1][c].norm() < cutoff) {
474 break;
475 }
476 }
477 let nk = rows.len();
478 let mut h = vec![C64::new(0.0, 0.0); nk * d];
479 for (k, row) in rows.iter().enumerate() {
480 for c in 0..d {
481 h[k + nk * c] = row[c];
482 }
483 }
484 (h, nk)
485}