Skip to main content

sparse_ir_minipole/minipole/
mini_pole.rs

1//! MPM: minimal poles from Matsubara data on a uniform grid (port of
2//! `mini_pole/mini_pole.py`).
3//!
4//! Ported from Green-Phys/MiniPole (commit 15e4a54, MIT License,
5//! Copyright (c) 2024 lzphy); see `LICENSE-THIRD-PARTY`.
6
7use 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/// Choice of `n0`, the number of low frequencies left out of the contour.
20#[derive(Debug, Clone, Copy, PartialEq, Eq)]
21pub enum N0 {
22    /// Chosen from the data, plus `shift`.
23    Auto {
24        /// Shift added to the automatic choice.
25        shift: usize,
26    },
27    /// Fixed.
28    Fixed(usize),
29}
30
31/// Plane in which the pole weights are computed.
32#[derive(Debug, Clone, Copy, PartialEq, Eq)]
33pub enum Plane {
34    /// Least squares on the Matsubara data.
35    Z,
36    /// From the ESPRIT weights in the mapped plane.
37    W,
38}
39
40/// Parameters of [`mini_pole`], with the defaults of the reference.
41#[derive(Debug, Clone, PartialEq)]
42pub struct MiniPoleParams {
43    /// Choice of `n0` (default automatic, shift 0).
44    pub n0: N0,
45    /// Error tolerance, at least the noise level. Required: the reference's
46    /// default (knee detection) is not ported.
47    pub err: f64,
48    /// Interpretation of `err` (default absolute).
49    pub err_type: ErrType,
50    /// Number of poles; `None` uses the precision of the first ESPRIT.
51    pub m: Option<usize>,
52    /// Preserve up-down symmetry.
53    pub symmetry: bool,
54    /// Symmetrize the data as `G_ij(z) = G_ji(z)`.
55    pub g_symmetric: bool,
56    /// Fit a constant term of `G`.
57    pub compute_const: bool,
58    /// Plane for the pole weights; `None` uses `Z` without and `W` with
59    /// symmetry.
60    pub plane: Option<Plane>,
61    /// Include the first `n0` points in the weight fit in the z plane.
62    pub include_n0: bool,
63    /// Maximum number of contour integrals (default 999).
64    pub k_max: usize,
65    /// Maximum ratio of oscillation when choosing `n0` (default 10).
66    pub ratio_max: f64,
67}
68
69impl MiniPoleParams {
70    /// Parameters with tolerance `err` and defaults otherwise.
71    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
88/// Minimal pole representation of Matsubara data.
89///
90/// `g_w` has shape `[n_w]` or `[n_w, n_orb, n_orb]`, `w` the corresponding
91/// finite, increasing, uniformly spaced non-negative Matsubara frequencies
92/// `ω_n` (real).
93///
94/// # Errors
95/// [`Error::ShapeMismatch`] for inconsistent shapes,
96/// [`Error::InvalidParameter`] for a grid that is not finite, increasing,
97/// uniform and non-negative or invalid parameters, and ESPRIT errors.
98pub 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    // First ESPRIT on each component.
176    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        // ConMapGapless needs ω_min > 0 (e.g. n0 = 0 on a bosonic grid).
239        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        // Complex poles for the data in [iω_max, i∞).
270        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        // cal_hk_gapless_symmetric_indiv on one function
300        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                        // Channel of G_ij in the column-major layout.
324                        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    // find_poles: second ESPRIT on the contour integrals.
351    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(); // (frequency index, mirrored)
400        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    // Discard poles with negligible weights.
429    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
453/// Rows `h_k` for `k = 0, 1, …` until two successive rows are below
454/// `cutoff` in every component, at most `k_max` rows; column-major `K x d`.
455fn 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}