Skip to main content

sparse_ir_basis/
poly.rs

1//! Piecewise Legendre polynomial implementations for SparseIR
2//!
3//! This module provides high-performance piecewise Legendre polynomial
4//! functionality compatible with the C++ implementation.
5
6use crate::error::Error;
7use crate::matrix::{Mat, Mat3};
8
9/// A single piecewise Legendre polynomial
10#[derive(Debug, Clone)]
11pub struct PiecewiseLegendrePoly {
12    /// Polynomial order (degree of Legendre polynomials in each segment)
13    pub(crate) polyorder: usize,
14    /// Minimum x value of the domain
15    pub(crate) xmin: f64,
16    /// Maximum x value of the domain
17    pub(crate) xmax: f64,
18    /// Knot points defining the segments
19    pub(crate) knots: Vec<f64>,
20    /// Segment widths (for numerical stability)
21    pub(crate) delta_x: Vec<f64>,
22    /// Coefficient matrix: [degree][segment_index]
23    pub(crate) data: Mat<f64>,
24    /// Symmetry parameter
25    pub(crate) symm: i32,
26    /// Index of this function in the sequence of singular functions it belongs
27    /// to (0-based, in order of non-increasing singular value)
28    ///
29    /// SVE results set it to the position of the function in the result, and
30    /// `PiecewiseLegendrePolyVector::from_3d_data` to the position in the
31    /// vector. `PiecewiseLegendreFT` takes the parity `(-1)^l` of the function
32    /// from it for the asymptotic expansion used at |n| >= n_asymp. For the
33    /// centrosymmetric kernels of this crate the even and odd singular
34    /// functions interlace, so that parity equals `symm`.
35    pub(crate) l: i32,
36    /// Segment midpoints
37    pub(crate) xm: Vec<f64>,
38    /// Inverse segment widths
39    pub(crate) inv_xs: Vec<f64>,
40    /// Normalization factors
41    pub(crate) norms: Vec<f64>,
42}
43
44/// `Ok` if there are `nsegments + 1` finite knots and every segment length is
45/// a positive normal double (so that `2 / length` is finite)
46fn check_knots(knots: &[f64], nsegments: usize) -> Result<(), Error> {
47    if knots.len() != nsegments + 1 {
48        return Err(Error::InvalidParameter {
49            name: "knots",
50            value: format!("{} knots", knots.len()),
51            reason: format!(
52                "must have {} entries, one more than the segments of data",
53                nsegments + 1
54            ),
55        });
56    }
57    if let Some((i, k)) = knots.iter().enumerate().find(|(_, k)| !k.is_finite()) {
58        return Err(Error::InvalidParameter {
59            name: "knots",
60            value: format!("{k:?} at index {i}"),
61            reason: "must be finite".to_string(),
62        });
63    }
64    for i in 1..knots.len() {
65        let length = knots[i] - knots[i - 1];
66        if !(length > 0.0 && length.is_normal()) {
67            return Err(Error::InvalidParameter {
68                name: "knots",
69                value: format!("{:?} after {:?} at index {i}", knots[i], knots[i - 1]),
70                reason: "must be strictly increasing, with each segment length a normal double"
71                    .to_string(),
72            });
73        }
74    }
75    Ok(())
76}
77
78/// `Ok` if `delta_x` has one entry per segment, each equal to the knot
79/// spacing `e = knots[i + 1] - knots[i]` within
80/// `max(1e-10 * |e|, 8 * f64::EPSILON * max(|knots[i]|, |knots[i + 1]|))`
81/// (NaN is rejected)
82///
83/// The first term is a relative tolerance on the segment length. The second
84/// covers the rounding of the knots, which grows with their magnitude: knots
85/// and widths scaled separately (as `FiniteTempBasis::from_sve_result` does,
86/// by β / 2) differ by a few units in the last place of the knots, which next
87/// to large knots can exceed 1e-10 times a narrow segment.
88fn check_delta_x(delta_x: &[f64], knots: &[f64]) -> Result<(), Error> {
89    let nsegments = knots.len() - 1;
90    if delta_x.len() != nsegments {
91        return Err(Error::InvalidParameter {
92            name: "delta_x",
93            value: format!("{} entries", delta_x.len()),
94            reason: format!("must have one entry per segment ({nsegments})"),
95        });
96    }
97    for (i, &d) in delta_x.iter().enumerate() {
98        let expected = knots[i + 1] - knots[i];
99        let rounding = 8.0 * f64::EPSILON * knots[i].abs().max(knots[i + 1].abs());
100        let tolerance = (1e-10 * expected.abs()).max(rounding);
101        if !((d - expected).abs() <= tolerance) {
102            return Err(Error::InvalidParameter {
103                name: "delta_x",
104                value: format!("{d:?} at index {i}"),
105                reason: format!(
106                    "must equal the knot spacing {expected:?} to a relative 1e-10, or to 8 machine epsilons times the magnitude of the knots"
107                ),
108            });
109        }
110    }
111    Ok(())
112}
113
114impl PiecewiseLegendrePoly {
115    /// Create a new PiecewiseLegendrePoly from data and knots
116    ///
117    /// `data` holds the Legendre coefficients, one column per segment; the
118    /// `nsegments + 1` knots bound the segments. `delta_x` (the segment
119    /// widths) is computed from the knots when `None`. `symm` is the parity of
120    /// the polynomial: 1 (even), -1 (odd) or 0 (no definite parity).
121    ///
122    /// # Errors
123    ///
124    /// * [`Error::EmptyInput`] if `data` has no row or no column
125    /// * [`Error::InvalidParameter`] if `knots` does not have one entry more
126    ///   than `data` has columns, a knot is not finite, or a segment length
127    ///   `knots[i] - knots[i - 1]` is not a positive normal double (this
128    ///   includes decreasing knots, NaN, lengths that overflow and subnormal
129    ///   lengths); or if `delta_x` does not have one entry per segment or
130    ///   differs from the knot spacing `e` by more than
131    ///   `max(1e-10 * |e|, 8 * f64::EPSILON * max(|knots[i]|, |knots[i + 1]|))`
132    /// * [`Error::InvalidParameter`] if `symm` is not -1, 0 or 1
133    pub fn new(
134        data: Mat<f64>,
135        knots: Vec<f64>,
136        l: i32,
137        delta_x: Option<Vec<f64>>,
138        symm: i32,
139    ) -> Result<Self, Error> {
140        let polyorder = data.shape().0;
141        let nsegments = data.shape().1;
142        if polyorder == 0 || nsegments == 0 {
143            return Err(Error::EmptyInput { name: "data" });
144        }
145        check_knots(&knots, nsegments)?;
146
147        // Compute delta_x if not provided
148        let delta_x =
149            delta_x.unwrap_or_else(|| (1..knots.len()).map(|i| knots[i] - knots[i - 1]).collect());
150        check_delta_x(&delta_x, &knots)?;
151        if !matches!(symm, -1..=1) {
152            return Err(Error::InvalidParameter {
153                name: "symm",
154                value: symm.to_string(),
155                reason: "must be -1, 0 or 1".to_string(),
156            });
157        }
158
159        // Compute segment midpoints
160        let xm: Vec<f64> = (0..nsegments)
161            .map(|i| 0.5 * (knots[i] + knots[i + 1]))
162            .collect();
163
164        // Compute inverse segment widths
165        let inv_xs: Vec<f64> = delta_x.iter().map(|&dx| 2.0 / dx).collect();
166
167        // Compute normalization factors
168        let norms: Vec<f64> = inv_xs.iter().map(|&inv_x| inv_x.sqrt()).collect();
169
170        Ok(Self {
171            polyorder,
172            xmin: knots[0],
173            xmax: knots[knots.len() - 1],
174            knots,
175            delta_x,
176            data,
177            symm,
178            l,
179            xm,
180            inv_xs,
181            norms,
182        })
183    }
184
185    /// Create a new PiecewiseLegendrePoly with new data but same structure
186    ///
187    /// Crate-internal: `new_data` is not checked against the knots, so a
188    /// caller could break the invariants of the type.
189    pub(crate) fn with_data(&self, new_data: Mat<f64>) -> Self {
190        Self {
191            data: new_data,
192            ..self.clone()
193        }
194    }
195
196    /// Get the symmetry parameter
197    pub fn symm(&self) -> i32 {
198        self.symm
199    }
200
201    /// The polynomial with every coefficient negated
202    ///
203    /// Knots, widths and normalizations are those of `self`, so this equals
204    /// `new` on the negated data without repeating its checks.
205    pub(crate) fn negated(&self) -> Self {
206        self.with_data(Mat::<f64>::from_fn(self.data.dims(), |idx| -self.data[idx]))
207    }
208
209    /// Rescale domain: create a new polynomial with the same data but different knots
210    ///
211    /// This is useful for transforming from one domain to another, e.g.,
212    /// from x ∈ [-1, 1] to τ ∈ [0, β].
213    ///
214    /// # Arguments
215    ///
216    /// * `new_knots` - New knot points
217    /// * `new_delta_x` - Optional new segment widths (computed from knots if None)
218    /// * `new_symm` - Optional new symmetry parameter (keeps old if None)
219    ///
220    /// # Returns
221    ///
222    /// New polynomial with rescaled domain
223    ///
224    /// # Errors
225    ///
226    /// The errors of [`new`](Self::new) for the new knots and widths
227    pub fn rescale_domain(
228        &self,
229        new_knots: Vec<f64>,
230        new_delta_x: Option<Vec<f64>>,
231        new_symm: Option<i32>,
232    ) -> Result<Self, Error> {
233        Self::new(
234            self.data.clone(),
235            new_knots,
236            self.l,
237            new_delta_x,
238            new_symm.unwrap_or(self.symm),
239        )
240    }
241
242    /// Scale all data values by a constant factor
243    ///
244    /// This is useful for normalizations, e.g., multiplying by √β for
245    /// Fourier transform preparations.
246    ///
247    /// # Arguments
248    ///
249    /// * `factor` - Scaling factor to multiply all data by
250    ///
251    /// # Returns
252    ///
253    /// New polynomial with scaled data
254    pub fn scale_data(&self, factor: f64) -> Self {
255        Self::with_data(
256            self,
257            Mat::<f64>::from_fn(self.data.dims(), |idx| self.data[idx] * factor),
258        )
259    }
260
261    /// `Ok` if `x` lies in [xmin, xmax] (NaN does not)
262    fn check_in_domain(&self, name: &'static str, x: f64) -> Result<(), Error> {
263        if x >= self.xmin && x <= self.xmax {
264            Ok(())
265        } else {
266            Err(Error::OutOfDomain {
267                name,
268                value: x,
269                domain: (self.xmin, self.xmax),
270            })
271        }
272    }
273
274    /// Evaluate the polynomial at a given point
275    ///
276    /// # Panics
277    ///
278    /// Panics if `x` is outside [xmin, xmax] or NaN; see [`Self::try_evaluate`].
279    pub fn evaluate(&self, x: f64) -> f64 {
280        self.try_evaluate(x).unwrap_or_else(|e| panic!("{e}"))
281    }
282
283    /// [`Self::evaluate`] returning an error instead of panicking
284    ///
285    /// # Errors
286    ///
287    /// [`Error::OutOfDomain`] if `x` is outside [xmin, xmax] or NaN
288    pub fn try_evaluate(&self, x: f64) -> Result<f64, Error> {
289        self.check_in_domain("x", x)?;
290        Ok(self.evaluate_in_domain(x))
291    }
292
293    /// [`Self::evaluate`] for an `x` already checked to lie in the domain
294    fn evaluate_in_domain(&self, x: f64) -> f64 {
295        let (i, x_tilde) = self.split_in_domain(x);
296        // Extract column i into a Vec
297        let coeffs: Vec<f64> = (0..self.data.shape().0)
298            .map(|row| self.data[[row, i]])
299            .collect();
300        let value = self.evaluate_legendre_polynomial(x_tilde, &coeffs);
301        value * self.norms[i]
302    }
303
304    /// Evaluate the polynomial at multiple points
305    ///
306    /// # Panics
307    ///
308    /// Panics if a point is outside [xmin, xmax] or NaN; see
309    /// [`Self::try_evaluate_many`].
310    pub fn evaluate_many(&self, xs: &[f64]) -> Vec<f64> {
311        self.try_evaluate_many(xs).unwrap_or_else(|e| panic!("{e}"))
312    }
313
314    /// [`Self::evaluate_many`] returning an error instead of panicking
315    ///
316    /// # Errors
317    ///
318    /// [`Error::OutOfDomain`] for the first point of `xs` that is outside
319    /// [xmin, xmax] or NaN; no point is evaluated then
320    pub fn try_evaluate_many(&self, xs: &[f64]) -> Result<Vec<f64>, Error> {
321        for &x in xs {
322            self.check_in_domain("xs", x)?;
323        }
324        Ok(xs.iter().map(|&x| self.evaluate_in_domain(x)).collect())
325    }
326
327    /// Split x into segment index and normalized x
328    ///
329    /// # Panics
330    ///
331    /// Panics if `x` is outside [xmin, xmax] or NaN; see [`Self::try_split`].
332    pub fn split(&self, x: f64) -> (usize, f64) {
333        self.try_split(x).unwrap_or_else(|e| panic!("{e}"))
334    }
335
336    /// [`Self::split`] returning an error instead of panicking
337    ///
338    /// # Errors
339    ///
340    /// [`Error::OutOfDomain`] if `x` is outside [xmin, xmax] or NaN
341    pub fn try_split(&self, x: f64) -> Result<(usize, f64), Error> {
342        self.check_in_domain("x", x)?;
343        Ok(self.split_in_domain(x))
344    }
345
346    /// [`Self::split`] for an `x` already checked to lie in the domain
347    fn split_in_domain(&self, x: f64) -> (usize, f64) {
348        // Find the segment containing x
349        for i in 0..self.knots.len() - 1 {
350            if x >= self.knots[i] && x <= self.knots[i + 1] {
351                // Transform x to [-1, 1] for Legendre polynomials
352                let x_tilde = 2.0 * (x - self.xm[i]) / self.delta_x[i];
353                return (i, x_tilde);
354            }
355        }
356
357        // Handle edge case: x exactly at the last knot
358        let last_idx = self.knots.len() - 2;
359        let x_tilde = 2.0 * (x - self.xm[last_idx]) / self.delta_x[last_idx];
360        (last_idx, x_tilde)
361    }
362
363    /// Evaluate Legendre polynomial using recurrence relation
364    pub fn evaluate_legendre_polynomial(&self, x: f64, coeffs: &[f64]) -> f64 {
365        if coeffs.is_empty() {
366            return 0.0;
367        }
368
369        let mut result = 0.0;
370        let mut p_prev = 1.0; // P_0(x) = 1
371        let mut p_curr = x; // P_1(x) = x
372
373        // Add first two terms
374        if !coeffs.is_empty() {
375            result += coeffs[0] * p_prev;
376        }
377        if coeffs.len() > 1 {
378            result += coeffs[1] * p_curr;
379        }
380
381        // Use recurrence relation: P_{n+1}(x) = ((2n+1)x*P_n(x) - n*P_{n-1}(x))/(n+1)
382        for n in 1..coeffs.len() - 1 {
383            let p_next =
384                ((2.0 * (n as f64) + 1.0) * x * p_curr - (n as f64) * p_prev) / ((n + 1) as f64);
385            result += coeffs[n + 1] * p_next;
386            p_prev = p_curr;
387            p_curr = p_next;
388        }
389
390        result
391    }
392
393    /// Compute derivative of the polynomial
394    ///
395    /// The result has `polyorder` equal to the number of its coefficient rows
396    /// (at least 1).
397    pub fn deriv(&self, n: usize) -> Self {
398        if n == 0 {
399            return self.clone();
400        }
401
402        // Compute derivative coefficients
403        let mut ddata = self.data.clone();
404        for _ in 0..n {
405            ddata = self.compute_derivative_coefficients(&ddata);
406        }
407
408        // Apply scaling factors (C++: ddata.col(i) *= std::pow(inv_xs[i], n))
409        let ddata_shape = *ddata.shape();
410        for i in 0..ddata_shape.1 {
411            let inv_x_power = self.inv_xs[i].powi(n as i32);
412            for j in 0..ddata_shape.0 {
413                ddata[[j, i]] *= inv_x_power;
414            }
415        }
416
417        // Update symmetry: C++: int new_symm = std::pow(-1, n) * symm;
418        let new_symm = if n % 2 == 0 { self.symm } else { -self.symm };
419
420        Self {
421            polyorder: ddata.shape().0,
422            data: ddata,
423            symm: new_symm,
424            ..self.clone()
425        }
426    }
427
428    /// Compute derivative coefficients using the same algorithm as C++ legder function
429    fn compute_derivative_coefficients(&self, coeffs: &Mat<f64>) -> Mat<f64> {
430        let mut c = coeffs.clone();
431        let c_shape = *c.shape();
432        let mut n = c_shape.0;
433
434        // Single derivative step (equivalent to C++ legder with cnt=1)
435        if n <= 1 {
436            return Mat::<f64>::from_elem([1, c.shape().1], 0.0);
437        }
438
439        n -= 1;
440        let mut der = Mat::<f64>::from_elem([n, c.shape().1], 0.0);
441
442        // C++ implementation: for (int j = n; j >= 2; --j)
443        for j in (2..=n).rev() {
444            // C++: der.row(j - 1) = (2 * j - 1) * c.row(j);
445            for col in 0..c_shape.1 {
446                der[[j - 1, col]] = (2.0 * (j as f64) - 1.0) * c[[j, col]];
447            }
448            // C++: c.row(j - 2) += c.row(j);
449            for col in 0..c_shape.1 {
450                c[[j - 2, col]] += c[[j, col]];
451            }
452        }
453
454        // C++: if (n > 1) der.row(1) = 3 * c.row(2);
455        if n > 1 {
456            for col in 0..c_shape.1 {
457                der[[1, col]] = 3.0 * c[[2, col]];
458            }
459        }
460
461        // C++: der.row(0) = c.row(1);
462        for col in 0..c_shape.1 {
463            der[[0, col]] = c[[1, col]];
464        }
465
466        der
467    }
468
469    /// Compute derivatives at a point x
470    ///
471    /// Returns the values of the derivatives of order 0 to `polyorder - 1`.
472    ///
473    /// # Panics
474    ///
475    /// Panics if `x` is outside [xmin, xmax] or NaN; see [`Self::try_derivs`].
476    pub fn derivs(&self, x: f64) -> Vec<f64> {
477        self.try_derivs(x).unwrap_or_else(|e| panic!("{e}"))
478    }
479
480    /// [`Self::derivs`] returning an error instead of panicking
481    ///
482    /// # Errors
483    ///
484    /// [`Error::OutOfDomain`] if `x` is outside [xmin, xmax] or NaN
485    pub fn try_derivs(&self, x: f64) -> Result<Vec<f64>, Error> {
486        self.check_in_domain("x", x)?;
487        let mut results = Vec::new();
488
489        // Compute up to polyorder derivatives
490        // The derivatives have the knots of self, so x lies in their domain.
491        for n in 0..self.polyorder {
492            let deriv_poly = self.deriv(n);
493            results.push(deriv_poly.evaluate_in_domain(x));
494        }
495
496        Ok(results)
497    }
498
499    /// Compute overlap integral with a function
500    pub fn overlap<F>(&self, f: F) -> f64
501    where
502        F: Fn(f64) -> f64,
503    {
504        let mut integral = 0.0;
505
506        for i in 0..self.knots.len() - 1 {
507            let segment_integral =
508                self.gauss_legendre_quadrature(self.knots[i], self.knots[i + 1], |x| {
509                    self.evaluate(x) * f(x)
510                });
511            integral += segment_integral;
512        }
513
514        integral
515    }
516
517    /// Gauss-Legendre quadrature over [a, b]
518    fn gauss_legendre_quadrature<F>(&self, a: f64, b: f64, f: F) -> f64
519    where
520        F: Fn(f64) -> f64,
521    {
522        // 5-point Gauss-Legendre quadrature
523        const XG: [f64; 5] = [
524            -0.906179845938664,
525            -0.538469310105683,
526            0.0,
527            0.538469310105683,
528            0.906179845938664,
529        ];
530        const WG: [f64; 5] = [
531            0.236926885056189,
532            0.478628670499366,
533            0.568888888888889,
534            0.478628670499366,
535            0.236926885056189,
536        ];
537
538        let c1 = (b - a) / 2.0;
539        let c2 = (b + a) / 2.0;
540
541        let mut integral = 0.0;
542        for j in 0..5 {
543            let x = c1 * XG[j] + c2;
544            integral += WG[j] * f(x);
545        }
546
547        integral * c1
548    }
549
550    /// Find roots of the polynomial using C++ compatible algorithm
551    pub fn roots(&self) -> Vec<f64> {
552        let xmid = (self.xmax + self.xmin) / 2.0;
553
554        // Exploit symmetry: only search the right half, then mirror.
555        // This matches the Python 1.x / Julia v1 algorithm and guarantees
556        // exactly symmetric root positions.
557        let grid = if self.symm != 0 {
558            let nsegments = self.knots.len() - 1;
559            let mid_idx = nsegments / 2;
560            if (self.knots[mid_idx] - xmid).abs() < 1e-15 {
561                self.knots[mid_idx..].to_vec()
562            } else {
563                let mut g = vec![xmid];
564                g.extend(self.knots.iter().filter(|&&x| x > xmid));
565                g
566            }
567        } else {
568            self.knots.clone()
569        };
570
571        let refined_grid = self.refine_grid(&grid, 4);
572        let roots_half = self.find_all_roots(&refined_grid);
573
574        if self.symm == 1 {
575            // Even symmetry: roots on right half, mirror to left
576            let mut all_roots: Vec<f64> = roots_half
577                .iter()
578                .rev()
579                .map(|&r| (self.xmax + self.xmin) - r)
580                .collect();
581            all_roots.extend_from_slice(&roots_half);
582            all_roots
583        } else if self.symm == -1 {
584            // Odd symmetry: there must be a zero at xmid
585            let mut right = roots_half;
586            if !right.is_empty() {
587                // Remove the root at xmid if found (may be slightly off),
588                // or if f(xmid) and f'(xmid) have opposite signs (spurious zero)
589                let f_mid = self.evaluate(xmid);
590                let f_deriv_mid = self.deriv(1).evaluate(xmid);
591                if (right[0] - xmid).abs() < 1e-13 || f_mid * f_deriv_mid < 0.0 {
592                    right.remove(0);
593                }
594            }
595            let mut all_roots: Vec<f64> = right
596                .iter()
597                .rev()
598                .map(|&r| (self.xmax + self.xmin) - r)
599                .collect();
600            all_roots.push(xmid);
601            all_roots.extend_from_slice(&right);
602            all_roots
603        } else {
604            // No symmetry: search the full domain
605            let full_grid = self.refine_grid(&self.knots, 4);
606            self.find_all_roots(&full_grid)
607        }
608    }
609
610    /// Refine grid by factor alpha (C++ compatible)
611    fn refine_grid(&self, grid: &[f64], alpha: usize) -> Vec<f64> {
612        let mut refined = Vec::new();
613
614        for i in 0..grid.len() - 1 {
615            let start = grid[i];
616            let step = (grid[i + 1] - grid[i]) / (alpha as f64);
617            for j in 0..alpha {
618                refined.push(start + (j as f64) * step);
619            }
620        }
621        refined.push(grid[grid.len() - 1]);
622        refined
623    }
624
625    /// Find all roots using refined grid (C++ compatible)
626    fn find_all_roots(&self, xgrid: &[f64]) -> Vec<f64> {
627        if xgrid.is_empty() {
628            return Vec::new();
629        }
630
631        // Evaluate function at all grid points
632        let fx: Vec<f64> = xgrid.iter().map(|&x| self.evaluate(x)).collect();
633
634        // Find exact zeros (direct hits)
635        let mut x_hit = Vec::new();
636        for i in 0..fx.len() {
637            if fx[i] == 0.0 {
638                x_hit.push(xgrid[i]);
639            }
640        }
641
642        // Find sign changes
643        let mut sign_change = Vec::new();
644        for i in 0..fx.len() - 1 {
645            let has_sign_change = fx[i].signum() != fx[i + 1].signum();
646            let not_hit = fx[i] != 0.0 && fx[i + 1] != 0.0;
647            let sc = has_sign_change && not_hit;
648            sign_change.push(sc);
649        }
650
651        // If no sign changes, return only direct hits
652        if sign_change.iter().all(|&sc| !sc) {
653            x_hit.sort_by(|a, b| a.partial_cmp(b).unwrap());
654            return x_hit;
655        }
656
657        // Find intervals with sign changes
658        let mut a_intervals = Vec::new();
659        let mut b_intervals = Vec::new();
660        let mut fa_values = Vec::new();
661
662        for i in 0..sign_change.len() {
663            if sign_change[i] {
664                a_intervals.push(xgrid[i]);
665                b_intervals.push(xgrid[i + 1]);
666                fa_values.push(fx[i]);
667            }
668        }
669
670        // Calculate epsilon for convergence
671        let max_elm = xgrid.iter().map(|&x| x.abs()).fold(0.0, f64::max);
672        let epsilon_x = f64::EPSILON * max_elm;
673
674        // Use bisection for each interval with sign change
675        for i in 0..a_intervals.len() {
676            let root = self.bisect(a_intervals[i], b_intervals[i], fa_values[i], epsilon_x);
677            x_hit.push(root);
678        }
679
680        // Sort and return
681        x_hit.sort_by(|a, b| a.partial_cmp(b).unwrap());
682        x_hit
683    }
684
685    /// Bisection method to find root (C++ compatible)
686    fn bisect(&self, a: f64, b: f64, fa: f64, eps: f64) -> f64 {
687        let mut a = a;
688        let mut b = b;
689        let mut fa = fa;
690
691        loop {
692            let mid = (a + b) / 2.0;
693            if self.close_enough(a, mid, eps) {
694                return mid;
695            }
696
697            let fmid = self.evaluate(mid);
698            if fa.signum() != fmid.signum() {
699                b = mid;
700            } else {
701                a = mid;
702                fa = fmid;
703            }
704        }
705    }
706
707    /// Check if two values are close enough (C++ compatible)
708    fn close_enough(&self, a: f64, b: f64, eps: f64) -> bool {
709        (a - b).abs() <= eps
710    }
711
712    // Accessor methods to match C++ interface
713    pub fn get_xmin(&self) -> f64 {
714        self.xmin
715    }
716    pub fn get_xmax(&self) -> f64 {
717        self.xmax
718    }
719    pub fn get_l(&self) -> i32 {
720        self.l
721    }
722    pub fn get_domain(&self) -> (f64, f64) {
723        (self.xmin, self.xmax)
724    }
725    pub fn get_knots(&self) -> &[f64] {
726        &self.knots
727    }
728    pub fn get_delta_x(&self) -> &[f64] {
729        &self.delta_x
730    }
731    pub fn get_symm(&self) -> i32 {
732        self.symm
733    }
734    pub fn get_data(&self) -> &Mat<f64> {
735        &self.data
736    }
737    pub fn get_norms(&self) -> &[f64] {
738        &self.norms
739    }
740    pub fn get_polyorder(&self) -> usize {
741        self.polyorder
742    }
743}
744
745/// Vector of piecewise Legendre polynomials
746#[derive(Debug, Clone)]
747pub struct PiecewiseLegendrePolyVector {
748    /// Individual polynomials
749    pub(crate) polyvec: Vec<PiecewiseLegendrePoly>,
750}
751
752impl PiecewiseLegendrePolyVector {
753    /// Constructor with a vector of PiecewiseLegendrePoly
754    ///
755    /// # Errors
756    ///
757    /// * [`Error::EmptyInput`] if `polyvec` is empty
758    /// * [`Error::InvalidParameter`] if a polynomial has other knots or
759    ///   another data shape than the first: the accessors of the vector
760    ///   (`get_knots`, `get_data`, ...) describe all of them by the first
761    pub fn new(polyvec: Vec<PiecewiseLegendrePoly>) -> Result<Self, Error> {
762        let Some(first) = polyvec.first() else {
763            return Err(Error::EmptyInput { name: "polyvec" });
764        };
765        if let Some(i) = polyvec
766            .iter()
767            .position(|p| p.knots != first.knots || p.data.shape() != first.data.shape())
768        {
769            return Err(Error::InvalidParameter {
770                name: "polyvec",
771                value: format!("polynomial {i}"),
772                reason: "must have the knots and the data shape of polynomial 0".to_string(),
773            });
774        }
775        Ok(Self { polyvec })
776    }
777
778    /// Constructor that skips the checks of [`Self::new`]
779    ///
780    /// For the SVE, which builds every polynomial of a vector from the same
781    /// knots and the same data shape and may legitimately build an empty
782    /// vector (a symmetrized half of an expansion with no singular value of
783    /// that parity).
784    pub(crate) fn from_polys_unchecked(polyvec: Vec<PiecewiseLegendrePoly>) -> Self {
785        Self { polyvec }
786    }
787
788    /// Get the polynomials
789    pub fn get_polys(&self) -> &[PiecewiseLegendrePoly] {
790        &self.polyvec
791    }
792
793    /// Constructor with a 3D array, knots, and symmetry vector
794    ///
795    /// `data3d` has the shape `(polyorder, nsegments, npolys)`; polynomial `i`
796    /// gets `l = i` and the symmetry `symm[i]` (0 if `symm` is `None`).
797    ///
798    /// # Errors
799    ///
800    /// * [`Error::EmptyInput`] if `data3d` has no polynomial
801    /// * [`Error::InvalidParameter`] if `symm` does not have one entry per
802    ///   polynomial
803    /// * The errors of [`PiecewiseLegendrePoly::new`] for the data and knots
804    pub fn from_3d_data(
805        data3d: Mat3<f64>,
806        knots: Vec<f64>,
807        symm: Option<Vec<i32>>,
808    ) -> Result<Self, Error> {
809        let npolys = data3d.shape().2;
810        if npolys == 0 {
811            return Err(Error::EmptyInput { name: "data3d" });
812        }
813        let mut polyvec = Vec::with_capacity(npolys);
814
815        if let Some(ref symm_vec) = symm {
816            if symm_vec.len() != npolys {
817                return Err(Error::InvalidParameter {
818                    name: "symm",
819                    value: format!("{} entries", symm_vec.len()),
820                    reason: format!("must have one entry per polynomial ({npolys})"),
821                });
822            }
823        }
824
825        // Compute delta_x from knots
826        let delta_x: Vec<f64> = (1..knots.len()).map(|i| knots[i] - knots[i - 1]).collect();
827
828        for i in 0..npolys {
829            // Extract 2D data for this polynomial
830            let data3d_shape = data3d.shape();
831            let mut data = Mat::<f64>::from_elem([data3d_shape.0, data3d_shape.1], 0.0);
832            for j in 0..data3d_shape.0 {
833                for k in 0..data3d_shape.1 {
834                    data[[j, k]] = data3d[[j, k, i]];
835                }
836            }
837
838            let poly = PiecewiseLegendrePoly::new(
839                data,
840                knots.clone(),
841                i as i32,
842                Some(delta_x.clone()),
843                symm.as_ref().map_or(0, |s| s[i]),
844            )?;
845
846            polyvec.push(poly);
847        }
848
849        Ok(Self { polyvec })
850    }
851
852    /// Get the size of the vector
853    pub fn size(&self) -> usize {
854        self.polyvec.len()
855    }
856
857    /// Rescale domain for all polynomials in the vector
858    ///
859    /// Creates a new PiecewiseLegendrePolyVector where each polynomial has
860    /// the same data but new knots and delta_x.
861    ///
862    /// # Arguments
863    ///
864    /// * `new_knots` - New knot points (same for all polynomials)
865    /// * `new_delta_x` - Optional new segment widths
866    /// * `new_symm` - Optional vector of new symmetry parameters (one per polynomial)
867    ///
868    /// # Returns
869    ///
870    /// New vector with rescaled domains
871    ///
872    /// # Errors
873    ///
874    /// * [`Error::InvalidParameter`] if `new_symm` does not have one entry per
875    ///   polynomial
876    /// * The errors of [`PiecewiseLegendrePoly::rescale_domain`]
877    pub fn rescale_domain(
878        &self,
879        new_knots: Vec<f64>,
880        new_delta_x: Option<Vec<f64>>,
881        new_symm: Option<Vec<i32>>,
882    ) -> Result<Self, Error> {
883        if let Some(symm) = &new_symm {
884            if symm.len() != self.polyvec.len() {
885                return Err(Error::InvalidParameter {
886                    name: "new_symm",
887                    value: format!("{} entries", symm.len()),
888                    reason: format!(
889                        "must have one entry per polynomial ({})",
890                        self.polyvec.len()
891                    ),
892                });
893            }
894        }
895        let polyvec = self
896            .polyvec
897            .iter()
898            .enumerate()
899            .map(|(i, poly)| {
900                let symm = new_symm.as_ref().map(|s| s[i]);
901                poly.rescale_domain(new_knots.clone(), new_delta_x.clone(), symm)
902            })
903            .collect::<Result<_, _>>()?;
904        Ok(Self { polyvec })
905    }
906
907    /// Scale all data values by a constant factor
908    ///
909    /// Multiplies the data of all polynomials by the same factor.
910    ///
911    /// # Arguments
912    ///
913    /// * `factor` - Scaling factor to multiply all data by
914    ///
915    /// # Returns
916    ///
917    /// New vector with scaled data
918    pub fn scale_data(&self, factor: f64) -> Self {
919        let polyvec = self
920            .polyvec
921            .iter()
922            .map(|poly| poly.scale_data(factor))
923            .collect();
924
925        Self { polyvec }
926    }
927
928    /// Get polynomial by index (immutable)
929    pub fn get(&self, index: usize) -> Option<&PiecewiseLegendrePoly> {
930        self.polyvec.get(index)
931    }
932
933    /// Get polynomial by index (mutable) - deprecated, use immutable design instead
934    #[deprecated(
935        note = "PiecewiseLegendrePolyVector is designed to be immutable. Use get() and create new instances for modifications."
936    )]
937    pub fn get_mut(&mut self, index: usize) -> Option<&mut PiecewiseLegendrePoly> {
938        self.polyvec.get_mut(index)
939    }
940
941    /// Extract a single polynomial as a vector
942    pub fn slice_single(&self, index: usize) -> Option<Self> {
943        self.polyvec.get(index).map(|poly| Self {
944            polyvec: vec![poly.clone()],
945        })
946    }
947
948    /// Extract multiple polynomials by indices, in the order of `indices`
949    ///
950    /// # Errors
951    ///
952    /// * [`Error::EmptyInput`] if `indices` is empty
953    /// * [`Error::InvalidParameter`] for the first index that is out of range
954    ///   or repeats an earlier one
955    pub fn slice_multi(&self, indices: &[usize]) -> Result<Self, Error> {
956        if indices.is_empty() {
957            return Err(Error::EmptyInput { name: "indices" });
958        }
959        let len = self.polyvec.len();
960        let mut seen = vec![false; len];
961        for &idx in indices {
962            if idx >= len {
963                return Err(Error::InvalidParameter {
964                    name: "indices",
965                    value: format!("{idx}"),
966                    reason: format!("must be less than the size {len}"),
967                });
968            }
969            if seen[idx] {
970                return Err(Error::InvalidParameter {
971                    name: "indices",
972                    value: format!("{idx}"),
973                    reason: "must not repeat".to_string(),
974                });
975            }
976            seen[idx] = true;
977        }
978
979        let new_polyvec: Vec<_> = indices
980            .iter()
981            .map(|&idx| self.polyvec[idx].clone())
982            .collect();
983
984        Ok(Self {
985            polyvec: new_polyvec,
986        })
987    }
988
989    /// Evaluate all polynomials at a single point
990    ///
991    /// # Panics
992    ///
993    /// Panics if `x` is outside the domain of a polynomial or NaN; see
994    /// [`Self::try_evaluate_at`].
995    pub fn evaluate_at(&self, x: f64) -> Vec<f64> {
996        self.try_evaluate_at(x).unwrap_or_else(|e| panic!("{e}"))
997    }
998
999    /// [`Self::evaluate_at`] returning an error instead of panicking
1000    ///
1001    /// # Errors
1002    ///
1003    /// [`Error::OutOfDomain`] if `x` is outside the domain of a polynomial
1004    /// or NaN
1005    pub fn try_evaluate_at(&self, x: f64) -> Result<Vec<f64>, Error> {
1006        self.polyvec
1007            .iter()
1008            .map(|poly| poly.try_evaluate(x))
1009            .collect()
1010    }
1011
1012    /// Evaluate all polynomials at multiple points
1013    ///
1014    /// The result has the shape `(size, xs.len())`.
1015    ///
1016    /// # Panics
1017    ///
1018    /// Panics if a point is outside the domain of a polynomial or NaN; see
1019    /// [`Self::try_evaluate_at_many`].
1020    pub fn evaluate_at_many(&self, xs: &[f64]) -> Mat<f64> {
1021        self.try_evaluate_at_many(xs)
1022            .unwrap_or_else(|e| panic!("{e}"))
1023    }
1024
1025    /// [`Self::evaluate_at_many`] returning an error instead of panicking
1026    ///
1027    /// # Errors
1028    ///
1029    /// [`Error::OutOfDomain`] for the first point of `xs` that is outside
1030    /// the domain of a polynomial or NaN; every point is checked before any
1031    /// is evaluated
1032    pub fn try_evaluate_at_many(&self, xs: &[f64]) -> Result<Mat<f64>, Error> {
1033        // `new` gives every polynomial the knots of the first, so checking the
1034        // points against the first checks them against all of them.
1035        if let Some(first) = self.polyvec.first() {
1036            for &x in xs {
1037                first.check_in_domain("xs", x)?;
1038            }
1039        }
1040        let n_funcs = self.polyvec.len();
1041        let n_points = xs.len();
1042        let mut results = Mat::<f64>::from_elem([n_funcs, n_points], 0.0);
1043
1044        for (i, poly) in self.polyvec.iter().enumerate() {
1045            for (j, &x) in xs.iter().enumerate() {
1046                results[[i, j]] = poly.evaluate_in_domain(x);
1047            }
1048        }
1049
1050        Ok(results)
1051    }
1052
1053    // Accessor methods to match C++ interface
1054    pub fn xmin(&self) -> f64 {
1055        if self.polyvec.is_empty() {
1056            panic!("Cannot get xmin from empty PiecewiseLegendrePolyVector");
1057        }
1058        self.polyvec[0].xmin
1059    }
1060
1061    pub fn xmax(&self) -> f64 {
1062        if self.polyvec.is_empty() {
1063            panic!("Cannot get xmax from empty PiecewiseLegendrePolyVector");
1064        }
1065        self.polyvec[0].xmax
1066    }
1067
1068    pub fn get_knots(&self, tolerance: Option<f64>) -> Vec<f64> {
1069        if self.polyvec.is_empty() {
1070            panic!("Cannot get knots from empty PiecewiseLegendrePolyVector");
1071        }
1072        const DEFAULT_TOLERANCE: f64 = 1e-10;
1073        let tolerance = tolerance.unwrap_or(DEFAULT_TOLERANCE);
1074
1075        // Collect all knots from all polynomials
1076        let mut all_knots = Vec::new();
1077        for poly in &self.polyvec {
1078            for &knot in &poly.knots {
1079                all_knots.push(knot);
1080            }
1081        }
1082
1083        // Sort and remove duplicates
1084        {
1085            all_knots.sort_by(|a, b| a.partial_cmp(b).unwrap());
1086            all_knots.dedup_by(|a, b| (*a - *b).abs() < tolerance);
1087        }
1088        all_knots
1089    }
1090
1091    pub fn get_delta_x(&self) -> Vec<f64> {
1092        if self.polyvec.is_empty() {
1093            panic!("Cannot get delta_x from empty PiecewiseLegendrePolyVector");
1094        }
1095        self.polyvec[0].delta_x.clone()
1096    }
1097
1098    pub fn get_polyorder(&self) -> usize {
1099        if self.polyvec.is_empty() {
1100            panic!("Cannot get polyorder from empty PiecewiseLegendrePolyVector");
1101        }
1102        self.polyvec[0].polyorder
1103    }
1104
1105    pub fn get_norms(&self) -> &[f64] {
1106        if self.polyvec.is_empty() {
1107            panic!("Cannot get norms from empty PiecewiseLegendrePolyVector");
1108        }
1109        &self.polyvec[0].norms
1110    }
1111
1112    pub fn get_symm(&self) -> Vec<i32> {
1113        if self.polyvec.is_empty() {
1114            panic!("Cannot get symm from empty PiecewiseLegendrePolyVector");
1115        }
1116        self.polyvec.iter().map(|poly| poly.symm).collect()
1117    }
1118
1119    /// Get data as 3D tensor: [segment][degree][polynomial]
1120    pub fn get_data(&self) -> Mat3<f64> {
1121        if self.polyvec.is_empty() {
1122            panic!("Cannot get data from empty PiecewiseLegendrePolyVector");
1123        }
1124
1125        let nsegments = self.polyvec[0].data.shape().1;
1126        let polyorder = self.polyvec[0].polyorder;
1127        let npolys = self.polyvec.len();
1128
1129        let mut data = Mat3::<f64>::from_elem([nsegments, polyorder, npolys], 0.0);
1130
1131        for (poly_idx, poly) in self.polyvec.iter().enumerate() {
1132            for segment in 0..nsegments {
1133                for degree in 0..polyorder {
1134                    data[[segment, degree, poly_idx]] = poly.data[[degree, segment]];
1135                }
1136            }
1137        }
1138
1139        data
1140    }
1141
1142    /// Find roots of all polynomials
1143    pub fn roots(&self, tolerance: Option<f64>) -> Vec<f64> {
1144        if self.polyvec.is_empty() {
1145            panic!("Cannot get roots from empty PiecewiseLegendrePolyVector");
1146        }
1147        const DEFAULT_TOLERANCE: f64 = 1e-10;
1148        let tolerance = tolerance.unwrap_or(DEFAULT_TOLERANCE);
1149        let mut all_roots = Vec::new();
1150
1151        for poly in &self.polyvec {
1152            let poly_roots = poly.roots();
1153            for root in poly_roots {
1154                all_roots.push(root);
1155            }
1156        }
1157
1158        // Sort in descending order and remove duplicates (like C++ implementation)
1159        {
1160            all_roots.sort_by(|a, b| b.partial_cmp(a).unwrap());
1161            all_roots.dedup_by(|a, b| (*a - *b).abs() < tolerance);
1162        }
1163        all_roots
1164    }
1165
1166    /// Get reference to last polynomial
1167    ///
1168    /// C++ equivalent: u.polyvec.back()
1169    pub fn last(&self) -> &PiecewiseLegendrePoly {
1170        self.polyvec
1171            .last()
1172            .expect("Cannot get last from empty PiecewiseLegendrePolyVector")
1173    }
1174
1175    /// Get the number of roots
1176    pub fn nroots(&self, tolerance: Option<f64>) -> usize {
1177        if self.polyvec.is_empty() {
1178            panic!("Cannot get nroots from empty PiecewiseLegendrePolyVector");
1179        }
1180        self.roots(tolerance).len()
1181    }
1182}
1183
1184impl std::ops::Index<usize> for PiecewiseLegendrePolyVector {
1185    type Output = PiecewiseLegendrePoly;
1186
1187    fn index(&self, index: usize) -> &Self::Output {
1188        &self.polyvec[index]
1189    }
1190}
1191
1192/// Get default sampling points in [-1, 1]
1193///
1194/// C++ implementation: libsparseir/include/sparseir/basis.hpp:287-310
1195///
1196/// For orthogonal polynomials (the high-T limit of IR), we know that the
1197/// ideal sampling points for a basis of size L are the roots of the L-th
1198/// polynomial. We empirically find that these stay good sampling points
1199/// for our kernels (probably because the kernels are totally positive).
1200///
1201/// If we do not have enough polynomials in the basis, we approximate the
1202/// roots of the L'th polynomial by the extrema of the last basis function,
1203/// which is sensible due to the strong interleaving property of these
1204/// functions' roots.
1205///
1206/// `name` names `u` in the errors (`"u"` or `"v"`).
1207///
1208/// # Errors
1209///
1210/// * [`Error::InvalidParameter`] named `name` if `u` is not on [-1, 1]
1211///   (1e-10), i.e. not the unscaled functions of an SVE
1212/// * [`Error::NotSupported`] if the extrema are needed and the last function
1213///   has none (e.g. an SVE truncated to 2 functions, whose u_1 is monotonic)
1214pub(crate) fn default_sampling_points(
1215    u: &PiecewiseLegendrePolyVector,
1216    name: &'static str,
1217    l: usize,
1218) -> Result<Vec<f64>, Error> {
1219    // C++: if (u.xmin() != -1.0 || u.xmax() != 1.0)
1220    //          throw std::runtime_error("Expecting unscaled functions here.");
1221    if (u.xmin() - (-1.0)).abs() > 1e-10 || (u.xmax() - 1.0).abs() > 1e-10 {
1222        return Err(Error::InvalidParameter {
1223            name,
1224            value: format!("functions on [{:?}, {:?}]", u.xmin(), u.xmax()),
1225            reason: "must be the unscaled functions of an SVE, on [-1, 1]".to_string(),
1226        });
1227    }
1228
1229    let x0 = if l < u.polyvec.len() {
1230        // C++: return u.polyvec[L].roots();
1231        u[l].roots()
1232    } else {
1233        // C++: PiecewiseLegendrePoly poly = u.polyvec.back();
1234        //      Eigen::VectorXd maxima = poly.deriv().roots();
1235        let poly = u.last();
1236        let poly_deriv = poly.deriv(1);
1237        let maxima = poly_deriv.roots();
1238
1239        // C++ reads maxima[0] without a check; SparseIR.jl's first(maxima)
1240        // fails the same way.
1241        let (Some(&first), Some(&last)) = (maxima.first(), maxima.last()) else {
1242            return Err(Error::NotSupported {
1243                what: format!(
1244                    "default sampling points for {l} basis functions: the last singular \
1245                     function (l = {}) has no extrema to stand in for the roots of the \
1246                     missing function l = {l} (an SVE truncated to too few functions)",
1247                    poly.l
1248                ),
1249            });
1250        };
1251
1252        // C++: double left = (maxima[0] + poly.xmin) / 2.0;
1253        let left = (first + poly.xmin) / 2.0;
1254
1255        // C++: double right = (maxima[maxima.size() - 1] + poly.xmax) / 2.0;
1256        let right = (last + poly.xmax) / 2.0;
1257
1258        // C++: Eigen::VectorXd x0(maxima.size() + 2);
1259        //      x0[0] = left;
1260        //      x0.segment(1, maxima.size()) = maxima;
1261        //      x0[x0.size() - 1] = right;
1262        let mut x0_vec = Vec::with_capacity(maxima.len() + 2);
1263        x0_vec.push(left);
1264        x0_vec.extend_from_slice(&maxima);
1265        x0_vec.push(right);
1266        x0_vec
1267    };
1268
1269    // C++: if (x0.size() != L) { warning }
1270    if x0.len() != l {
1271        debug_warn!(
1272            "Expecting to get {} sampling points for corresponding basis function, \
1273             instead got {}. This may happen if not enough precision is left in the polynomial.",
1274            l,
1275            x0.len()
1276        );
1277    }
1278
1279    Ok(x0)
1280}
1281
1282// IndexMut implementation removed - PiecewiseLegendrePolyVector is designed to be immutable
1283// If modification is needed, create a new instance instead
1284
1285// Note: FnOnce implementation removed due to experimental nature
1286// Use evaluate_at() and evaluate_at_many() methods directly
1287
1288#[cfg(test)]
1289#[path = "poly_tests.rs"]
1290mod poly_tests;