Skip to main content

sparse_ir_basis/sve/
strategy.rs

1//! SVE computation strategies
2
3use crate::error::Error;
4use crate::gauss::{Rule, legendre_generic};
5use crate::kernel::{AbstractKernel, CentrosymmKernel, KernelProperties, SVEHints, SymmetryType};
6use crate::kernelmatrix::{matrix_from_gauss_noncentrosymmetric, matrix_from_gauss_with_segments};
7use crate::matrix::Mat;
8use crate::numeric::CustomNumeric;
9use crate::poly::PiecewiseLegendrePolyVector;
10use std::fmt::Debug;
11
12use super::result::SVEResult;
13use super::utils::{
14    SvdBlock, canonicalize_signs, extend_to_full_domain, merge_blocks,
15    mirror_segments_to_full_domain, remove_weights, svd_to_polynomials,
16};
17
18/// Trait for SVE computation strategies
19pub trait SVEStrategy<T: CustomNumeric> {
20    /// Compute the discretized matrices for SVD
21    fn matrices(&self) -> Vec<Mat<T>>;
22
23    /// Post-process SVD results to create SVEResult
24    ///
25    /// # Errors
26    ///
27    /// The errors of [`SVEResult::new`]
28    fn postprocess(
29        &self,
30        u_list: Vec<Mat<T>>,
31        s_list: Vec<Vec<T>>,
32        v_list: Vec<Mat<T>>,
33    ) -> Result<SVEResult, Error>;
34}
35
36/// Sampling-based SVE computation
37///
38/// This is the general SVE computation strategy that works with discretized kernels.
39/// It does NOT know about symmetry - it just processes a given discretized kernel matrix.
40///
41/// # Responsibility
42///
43/// - Remove weights from SVD results
44/// - Convert to polynomials on the domain specified by segments
45/// - Domain extension is caller's responsibility
46pub struct SamplingSVE<T>
47where
48    T: CustomNumeric + Send + Sync + 'static,
49{
50    segments_x: Vec<T>,
51    segments_y: Vec<T>,
52    gauss_x: Rule<T>,
53    gauss_y: Rule<T>,
54    #[allow(dead_code)]
55    epsilon: f64,
56    n_gauss: usize,
57}
58
59impl<T> SamplingSVE<T>
60where
61    T: CustomNumeric + Send + Sync + 'static,
62{
63    /// Create a new SamplingSVE
64    ///
65    /// This takes only the geometric information needed for polynomial conversion,
66    /// not the kernel itself.
67    pub fn new(
68        segments_x: Vec<T>,
69        segments_y: Vec<T>,
70        gauss_x: Rule<T>,
71        gauss_y: Rule<T>,
72        epsilon: f64,
73        n_gauss: usize,
74    ) -> Self {
75        Self {
76            segments_x,
77            segments_y,
78            gauss_x,
79            gauss_y,
80            epsilon,
81            n_gauss,
82        }
83    }
84
85    /// Post-process a single SVD result to create polynomials
86    ///
87    /// This converts SVD results to piecewise Legendre polynomials
88    /// on the domain specified by segments (e.g., [0, xmax] for reduced kernels).
89    ///
90    /// # Errors
91    ///
92    /// The errors of the crate-internal `svd_to_polynomials`, and
93    /// [`Error::EmptyInput`] if the
94    /// SVD result has no singular values
95    pub fn postprocess_single(
96        &self,
97        u: &Mat<T>,
98        s: &[T],
99        v: &Mat<T>,
100    ) -> Result<
101        (
102            PiecewiseLegendrePolyVector,
103            Vec<f64>,
104            PiecewiseLegendrePolyVector,
105        ),
106        Error,
107    > {
108        let (u_polys, s, v_polys) = self.postprocess_block(u, s, v)?;
109        Ok((
110            PiecewiseLegendrePolyVector::new(u_polys)?,
111            s,
112            PiecewiseLegendrePolyVector::new(v_polys)?,
113        ))
114    }
115
116    /// [`Self::postprocess_single`] returning plain vectors, which may be
117    /// empty
118    fn postprocess_block(&self, u: &Mat<T>, s: &[T], v: &Mat<T>) -> Result<SvdBlock, Error> {
119        // 1. Remove weights
120        // Both U and V have rows corresponding to Gauss points, so is_row=true for both
121        let u_unweighted = remove_weights(u, self.gauss_x.w.as_slice(), true);
122        let v_unweighted = remove_weights(v, self.gauss_y.w.as_slice(), true);
123
124        // 2. Convert to polynomials
125        let gauss_rule_f64 = legendre_generic::<f64>(self.n_gauss);
126        let u_polys = svd_to_polynomials(
127            &u_unweighted,
128            &self.segments_x,
129            &gauss_rule_f64,
130            self.n_gauss,
131        )?;
132        let v_polys = svd_to_polynomials(
133            &v_unweighted,
134            &self.segments_y,
135            &gauss_rule_f64,
136            self.n_gauss,
137        )?;
138
139        // Note: No domain extension here - that's the caller's responsibility
140        Ok((u_polys, s.iter().map(|&x| x.to_f64()).collect(), v_polys))
141    }
142}
143
144/// Centrosymmetric SVE computation
145///
146/// Exploits even/odd symmetry for efficient computation.
147/// This manages the symmetry: creating reduced kernels, extending to full domain, and merging.
148pub struct CentrosymmSVE<T, K>
149where
150    T: CustomNumeric + Send + Sync + 'static,
151    K: CentrosymmKernel + KernelProperties,
152{
153    kernel: K,
154    epsilon: f64,
155    hints: K::SVEHintsType<T>,
156    #[allow(dead_code)]
157    n_gauss: usize,
158
159    // Geometric information (positive domain [0, xmax])
160    #[allow(dead_code)]
161    segments_x: Vec<T>,
162    #[allow(dead_code)]
163    segments_y: Vec<T>,
164    gauss_x: Rule<T>,
165    gauss_y: Rule<T>,
166
167    // The general SVE processor (no symmetry knowledge)
168    sampling_sve: SamplingSVE<T>,
169}
170
171impl<T, K> CentrosymmSVE<T, K>
172where
173    T: CustomNumeric + Send + Sync + Clone + 'static,
174    K: CentrosymmKernel + KernelProperties + Clone,
175    K::SVEHintsType<T>: SVEHints<T> + Clone,
176{
177    /// Create a new CentrosymmSVE
178    ///
179    /// # Errors
180    ///
181    /// [`Error::InvalidParameter`] if the SVE hints of the kernel give
182    /// segments that are not finite and strictly increasing
183    pub fn new(kernel: K, epsilon: f64) -> Result<Self, Error> {
184        let hints = kernel.sve_hints::<T>(epsilon);
185
186        // Get segments for positive domain [0, xmax]
187        let segments_x = hints.segments_x();
188        let segments_y = hints.segments_y();
189        let n_gauss = hints.ngauss();
190
191        // Create composite Gauss rules
192        let rule = legendre_generic::<T>(n_gauss);
193        let gauss_x = rule.piecewise(&segments_x)?;
194        let gauss_y = rule.piecewise(&segments_y)?;
195
196        // Create the general SVE processor
197        let sampling_sve = SamplingSVE::new(
198            segments_x.clone(),
199            segments_y.clone(),
200            gauss_x.clone(),
201            gauss_y.clone(),
202            epsilon,
203            n_gauss,
204        );
205
206        Ok(Self {
207            kernel,
208            epsilon,
209            hints,
210            n_gauss,
211            segments_x,
212            segments_y,
213            gauss_x,
214            gauss_y,
215            sampling_sve,
216        })
217    }
218
219    /// Compute reduced kernel matrix for given symmetry
220    fn compute_reduced_matrix(&self, symmetry: SymmetryType) -> Mat<T> {
221        // Compute K_red(x, y) = K(x, y) + sign * K(x, -y)
222        // where x, y are in [0, xmax] and [0, ymax]
223        let discretized = matrix_from_gauss_with_segments(
224            &self.kernel,
225            &self.gauss_x,
226            &self.gauss_y,
227            symmetry,
228            &self.hints,
229        );
230
231        // Apply weights for SVE
232
233        discretized.apply_weights_for_sve()
234    }
235
236    /// Extend polynomials from [0, xmax] to [-xmax, xmax]
237    fn extend_result_to_full_domain(
238        &self,
239        result: SvdBlock,
240        symmetry: SymmetryType,
241    ) -> Result<SvdBlock, Error> {
242        let (u, s, v) = result;
243
244        // Extend u and v from [0, xmax] to [-xmax, xmax]
245        let u_full = extend_to_full_domain(u, symmetry, self.kernel.xmax())?;
246        let v_full = extend_to_full_domain(v, symmetry, self.kernel.ymax())?;
247
248        Ok((u_full, s, v_full))
249    }
250}
251
252impl<T, K> SVEStrategy<T> for CentrosymmSVE<T, K>
253where
254    T: CustomNumeric + Send + Sync + Clone + 'static,
255    K: CentrosymmKernel + KernelProperties + Clone,
256    K::SVEHintsType<T>: SVEHints<T> + Clone,
257{
258    fn matrices(&self) -> Vec<Mat<T>> {
259        // Compute reduced kernels for even and odd symmetries
260        let even_matrix = self.compute_reduced_matrix(SymmetryType::Even);
261        let odd_matrix = self.compute_reduced_matrix(SymmetryType::Odd);
262
263        vec![even_matrix, odd_matrix]
264    }
265
266    fn postprocess(
267        &self,
268        u_list: Vec<Mat<T>>,
269        s_list: Vec<Vec<T>>,
270        v_list: Vec<Mat<T>>,
271    ) -> Result<SVEResult, Error> {
272        // Process even and odd results using SamplingSVE (which doesn't know
273        // about symmetry). Keep plain vectors until the merge: truncation can
274        // empty a block (keeping only the largest singular value empties the
275        // odd one), and a PiecewiseLegendrePolyVector cannot be empty.
276        let result_even = self
277            .sampling_sve
278            .postprocess_block(&u_list[0], &s_list[0], &v_list[0])?;
279        let result_odd = self
280            .sampling_sve
281            .postprocess_block(&u_list[1], &s_list[1], &v_list[1])?;
282
283        // Now extend to full domain (this is where symmetry comes in)
284        let result_even_full =
285            self.extend_result_to_full_domain(result_even, SymmetryType::Even)?;
286        let result_odd_full = self.extend_result_to_full_domain(result_odd, SymmetryType::Odd)?;
287
288        // Merge the results (at least one singular value is kept)
289        merge_blocks(result_even_full, result_odd_full, self.epsilon)
290    }
291}
292
293/// SVE hints of a kernel expressed on its full domain `[-xmax, xmax] × [-ymax, ymax]`
294///
295/// [`SVEHints`] reports the segments of a centrosymmetric kernel on the
296/// half-domain `[0, xmax] × [0, ymax]` only, which is what [`CentrosymmSVE`]
297/// discretizes. [`NonCentrosymmSVE`] discretizes the full domain, so for a
298/// centrosymmetric kernel this wrapper mirrors the half-domain segments onto
299/// the full domain; the segments of other kernels already cover the full
300/// domain and are passed through unchanged.
301#[derive(Debug, Clone)]
302struct FullDomainHints<H> {
303    inner: H,
304    /// Whether `inner` reports half-domain segments (centrosymmetric kernel)
305    half_domain: bool,
306}
307
308impl<H> FullDomainHints<H> {
309    fn full_domain<T: CustomNumeric>(&self, segments: Vec<T>) -> Vec<T> {
310        if self.half_domain {
311            mirror_segments_to_full_domain(&segments)
312        } else {
313            segments
314        }
315    }
316}
317
318impl<T, H> SVEHints<T> for FullDomainHints<H>
319where
320    T: Copy + Debug + Send + Sync + CustomNumeric,
321    H: SVEHints<T>,
322{
323    fn segments_x(&self) -> Vec<T> {
324        self.full_domain(self.inner.segments_x())
325    }
326
327    fn segments_y(&self) -> Vec<T> {
328        self.full_domain(self.inner.segments_y())
329    }
330
331    fn nsvals(&self) -> usize {
332        self.inner.nsvals()
333    }
334
335    fn ngauss(&self) -> usize {
336        self.inner.ngauss()
337    }
338}
339
340/// Non-centrosymmetric SVE computation
341///
342/// This strategy computes the kernel matrix directly over the full domain
343/// [-xmax, xmax] × [-ymax, ymax]. No symmetry exploitation is performed, so it
344/// works for any kernel.
345///
346/// Centrosymmetric kernels are also expanded correctly: their [`SVEHints`]
347/// segments cover only the half-domain `[0, xmax]` and are mirrored onto the
348/// full domain before discretization. The singular values then agree with
349/// [`CentrosymmSVE`] up to rounding, but the SVD is taken of a single matrix
350/// with twice as many rows and columns as each even/odd block, and the
351/// singular functions carry no parity tag.
352///
353/// The SVD sign gauge is fixed by demanding `u_l(xmax) >= 0`, as
354/// [`CentrosymmSVE`] does and as libsparseir's `SamplingSVE::postprocess`
355/// does via `canonicalize` (convention-matched with
356/// `backend/cxx/include/sparseir/impl/sve_impl.ipp`, libsparseir commit
357/// 4bc58ea).
358#[allow(dead_code)]
359pub struct NonCentrosymmSVE<T, K>
360where
361    T: CustomNumeric + Send + Sync + 'static,
362    K: AbstractKernel + KernelProperties,
363{
364    kernel: K,
365    epsilon: f64,
366    hints: FullDomainHints<K::SVEHintsType<T>>,
367    n_gauss: usize,
368
369    // Geometric information (full domain [-xmax, xmax])
370    segments_x: Vec<T>,
371    segments_y: Vec<T>,
372    gauss_x: Rule<T>,
373    gauss_y: Rule<T>,
374
375    // The general SVE processor
376    sampling_sve: SamplingSVE<T>,
377}
378
379impl<T, K> NonCentrosymmSVE<T, K>
380where
381    T: CustomNumeric + Send + Sync + Clone + 'static,
382    K: AbstractKernel + KernelProperties + Clone,
383    K::SVEHintsType<T>: SVEHints<T> + Clone,
384{
385    /// Create a new NonCentrosymmSVE
386    ///
387    /// # Errors
388    ///
389    /// [`Error::InvalidParameter`] if the SVE hints of the kernel give
390    /// segments that are not finite and strictly increasing
391    pub fn new(kernel: K, epsilon: f64) -> Result<Self, Error> {
392        // SVEHints are half-domain for centrosymmetric kernels; this strategy
393        // needs the full domain (issue #246).
394        let hints = FullDomainHints {
395            inner: kernel.sve_hints::<T>(epsilon),
396            half_domain: kernel.is_centrosymmetric(),
397        };
398
399        // Get segments for full domain [-xmax, xmax]
400        let segments_x = hints.segments_x();
401        let segments_y = hints.segments_y();
402        let n_gauss = hints.ngauss();
403
404        // Create composite Gauss rules for full domain
405        let rule = legendre_generic::<T>(n_gauss);
406        let gauss_x = rule.piecewise(&segments_x)?;
407        let gauss_y = rule.piecewise(&segments_y)?;
408
409        // Create the general SVE processor
410        let sampling_sve = SamplingSVE::new(
411            segments_x.clone(),
412            segments_y.clone(),
413            gauss_x.clone(),
414            gauss_y.clone(),
415            epsilon,
416            n_gauss,
417        );
418
419        Ok(Self {
420            kernel,
421            epsilon,
422            hints,
423            n_gauss,
424            segments_x,
425            segments_y,
426            gauss_x,
427            gauss_y,
428            sampling_sve,
429        })
430    }
431
432    /// Compute kernel matrix for non-centrosymmetric kernel
433    fn compute_matrix(&self) -> Mat<T> {
434        // Compute K(x, y) directly over full domain
435        let discretized = matrix_from_gauss_noncentrosymmetric(
436            &self.kernel,
437            &self.gauss_x,
438            &self.gauss_y,
439            &self.hints,
440        );
441
442        // Apply weights for SVE
443        discretized.apply_weights_for_sve()
444    }
445}
446
447impl<T, K> SVEStrategy<T> for NonCentrosymmSVE<T, K>
448where
449    T: CustomNumeric + Send + Sync + Clone + 'static,
450    K: AbstractKernel + KernelProperties + Clone,
451    K::SVEHintsType<T>: SVEHints<T> + Clone,
452{
453    fn matrices(&self) -> Vec<Mat<T>> {
454        // Single matrix for non-centrosymmetric kernel
455        vec![self.compute_matrix()]
456    }
457
458    fn postprocess(
459        &self,
460        u_list: Vec<Mat<T>>,
461        s_list: Vec<Vec<T>>,
462        v_list: Vec<Mat<T>>,
463    ) -> Result<SVEResult, Error> {
464        // Process single result using SamplingSVE. The functions are
465        // already on the full domain. Fix the sign gauge u_l(xmax) >= 0 as
466        // CentrosymmSVE does in merge_results.
467        let (u_polys, s, v_polys) = self
468            .sampling_sve
469            .postprocess_block(&u_list[0], &s_list[0], &v_list[0])?;
470        let (u_polys, v_polys) = canonicalize_signs(u_polys, v_polys);
471        // A matrix of rank 0 leaves empty vectors; SVEResult::new reports the
472        // empty `s`.
473        SVEResult::new(
474            PiecewiseLegendrePolyVector::from_polys_unchecked(u_polys),
475            s,
476            PiecewiseLegendrePolyVector::from_polys_unchecked(v_polys),
477            self.epsilon,
478        )
479    }
480}