Skip to main content

sparse_ir_basis/
kernelmatrix.rs

1//! Kernel matrix discretization for SparseIR
2//!
3//! This module provides functionality to discretize kernels using Gauss quadrature
4//! rules and store them as matrices for numerical computation.
5
6use crate::gauss::Rule;
7use crate::kernel::{AbstractKernel, CentrosymmKernel, KernelProperties, SymmetryType};
8use crate::matrix::Mat;
9use crate::numeric::CustomNumeric;
10use std::fmt::Debug;
11
12/// This structure stores a discrete kernel matrix along with the corresponding
13/// Gauss quadrature rules for x and y coordinates. This enables easy application
14/// of weights for SVE computation and maintains the relationship between matrix
15/// elements and their corresponding quadrature points.
16#[derive(Debug, Clone)]
17pub struct DiscretizedKernel<T> {
18    /// Discrete kernel matrix
19    pub matrix: Mat<T>,
20    /// Gauss quadrature rule for x coordinates
21    pub gauss_x: Rule<T>,
22    /// Gauss quadrature rule for y coordinates
23    pub gauss_y: Rule<T>,
24    /// X-axis segment boundaries (from SVEHints)
25    pub segments_x: Vec<T>,
26    /// Y-axis segment boundaries (from SVEHints)
27    pub segments_y: Vec<T>,
28}
29
30impl<T: CustomNumeric + Clone> DiscretizedKernel<T> {
31    /// Create a new DiscretizedKernel
32    pub fn new(
33        matrix: Mat<T>,
34        gauss_x: Rule<T>,
35        gauss_y: Rule<T>,
36        segments_x: Vec<T>,
37        segments_y: Vec<T>,
38    ) -> Self {
39        Self {
40            matrix,
41            gauss_x,
42            gauss_y,
43            segments_x,
44            segments_y,
45        }
46    }
47
48    /// Create a new DiscretizedKernel without segments (legacy)
49    pub fn new_legacy(matrix: Mat<T>, gauss_x: Rule<T>, gauss_y: Rule<T>) -> Self {
50        Self {
51            matrix,
52            gauss_x: gauss_x.clone(),
53            gauss_y: gauss_y.clone(),
54            segments_x: vec![gauss_x.a, gauss_x.b],
55            segments_y: vec![gauss_y.a, gauss_y.b],
56        }
57    }
58
59    /// Delegate to matrix methods
60    pub fn is_empty(&self) -> bool {
61        self.matrix.is_empty()
62    }
63
64    pub fn nrows(&self) -> usize {
65        self.matrix.shape().0
66    }
67
68    pub fn ncols(&self) -> usize {
69        self.matrix.shape().1
70    }
71
72    pub fn iter(&self) -> impl Iterator<Item = &T> {
73        self.matrix.iter()
74    }
75
76    /// Apply weights for SVE computation
77    ///
78    /// This applies the square root of Gauss weights to the matrix,
79    /// which is required before performing SVD for SVE computation.
80    /// The original matrix remains unchanged.
81    pub fn apply_weights_for_sve(&self) -> Mat<T> {
82        let mut weighted_matrix = self.matrix.clone();
83        let shape = *weighted_matrix.shape();
84
85        // Apply square root of x-direction weights to rows
86        for i in 0..self.gauss_x.x.len() {
87            let weight_sqrt = self.gauss_x.w[i].sqrt();
88            for j in 0..shape.1 {
89                weighted_matrix[[i, j]] = weighted_matrix[[i, j]] * weight_sqrt;
90            }
91        }
92
93        // Apply square root of y-direction weights to columns
94        for j in 0..self.gauss_y.x.len() {
95            let weight_sqrt = self.gauss_y.w[j].sqrt();
96            for i in 0..shape.0 {
97                weighted_matrix[[i, j]] = weighted_matrix[[i, j]] * weight_sqrt;
98            }
99        }
100
101        weighted_matrix
102    }
103
104    /// Remove weights from matrix (inverse of apply_weights_for_sve)
105    pub fn remove_weights_from_sve(&mut self) {
106        let shape = *self.matrix.shape();
107
108        // Remove weights from U matrix (x-direction)
109        for i in 0..self.gauss_x.x.len() {
110            let weight_sqrt = self.gauss_x.w[i].sqrt();
111            for j in 0..shape.1 {
112                self.matrix[[i, j]] = self.matrix[[i, j]] / weight_sqrt;
113            }
114        }
115
116        // Remove weights from V matrix (y-direction)
117        for j in 0..self.gauss_y.x.len() {
118            let weight_sqrt = self.gauss_y.w[j].sqrt();
119            for i in 0..shape.0 {
120                self.matrix[[i, j]] = self.matrix[[i, j]] / weight_sqrt;
121            }
122        }
123    }
124
125    /// Get the number of Gauss points in x direction
126    pub fn n_gauss_x(&self) -> usize {
127        self.gauss_x.x.len()
128    }
129
130    /// Get the number of Gauss points in y direction
131    pub fn n_gauss_y(&self) -> usize {
132        self.gauss_y.x.len()
133    }
134}
135
136/// Compute matrix from Gauss quadrature rules with segments from SVEHints
137///
138/// This function evaluates the kernel at all combinations of Gauss points
139/// and returns a DiscretizedKernel containing the matrix, quadrature rules, and segments.
140pub fn matrix_from_gauss_with_segments<
141    T: CustomNumeric + Clone + Send + Sync,
142    K: CentrosymmKernel + KernelProperties,
143    H: crate::kernel::SVEHints<T>,
144>(
145    kernel: &K,
146    gauss_x: &Rule<T>,
147    gauss_y: &Rule<T>,
148    symmetry: SymmetryType,
149    hints: &H,
150) -> DiscretizedKernel<T> {
151    let segments_x = hints.segments_x();
152    let segments_y = hints.segments_y();
153
154    // TODO: Fix range checking for composite Gauss rules
155    // For now, skip range checking to allow testing
156    /*
157    // Check that Gauss points are within [0, xmax] and [0, ymax]
158    let kernel_xmax = kernel.xmax();
159    let kernel_ymax = kernel.ymax();
160    let tolerance = 1e-12;
161
162    // Check x points are in [0, xmax]
163    for &x in &gauss_x.x {
164        let x_f64 = x.to_f64();
165        assert!(
166            x_f64 >= -tolerance && x_f64 <= kernel_xmax + tolerance,
167            "Gauss x point {} is outside [0, {}]", x_f64, kernel_xmax
168        );
169    }
170
171    // Check y points are in [0, ymax]
172    for &y in &gauss_y.x {
173        let y_f64 = y.to_f64();
174        assert!(
175            y_f64 >= -tolerance && y_f64 <= kernel_ymax + tolerance,
176            "Gauss y point {} is outside [0, {}]", y_f64, kernel_ymax
177        );
178    }
179    */
180
181    let n = gauss_x.x.len();
182    let m = gauss_y.x.len();
183    let mut result = Mat::<T>::from_elem([n, m], T::zero());
184
185    // Evaluate kernel at all combinations of Gauss points
186    for i in 0..n {
187        for j in 0..m {
188            let x = gauss_x.x[i];
189            let y = gauss_y.x[j];
190            result[[i, j]] = kernel.compute_reduced(x, y, symmetry);
191        }
192    }
193
194    DiscretizedKernel::new(
195        result,
196        gauss_x.clone(),
197        gauss_y.clone(),
198        segments_x,
199        segments_y,
200    )
201}
202
203/// Compute matrix from Gauss quadrature rules (legacy version without segments)
204///
205/// This function evaluates the kernel at all combinations of Gauss points
206/// and returns a DiscretizedKernel containing the matrix and quadrature rules.
207pub fn matrix_from_gauss<T: CustomNumeric + Clone, K: CentrosymmKernel + KernelProperties>(
208    kernel: &K,
209    gauss_x: &Rule<T>,
210    gauss_y: &Rule<T>,
211    symmetry: SymmetryType,
212) -> DiscretizedKernel<T> {
213    // Check that Gauss points are within [0, xmax] and [0, ymax]
214    let kernel_xmax = kernel.xmax();
215    let kernel_ymax = kernel.ymax();
216    let tolerance = 1e-12;
217
218    // Check x points are in [0, xmax]
219    for &x in &gauss_x.x {
220        let x_f64 = x.to_f64();
221        assert!(
222            x_f64 >= -tolerance && x_f64 <= kernel_xmax + tolerance,
223            "Gauss x point {} is outside [0, {}]",
224            x_f64,
225            kernel_xmax
226        );
227    }
228
229    // Check y points are in [0, ymax]
230    for &y in &gauss_y.x {
231        let y_f64 = y.to_f64();
232        assert!(
233            y_f64 >= -tolerance && y_f64 <= kernel_ymax + tolerance,
234            "Gauss y point {} is outside [0, {}]",
235            y_f64,
236            kernel_ymax
237        );
238    }
239
240    let n = gauss_x.x.len();
241    let m = gauss_y.x.len();
242    let mut result = Mat::<T>::from_elem([n, m], T::zero());
243
244    // Evaluate kernel at all combinations of Gauss points
245    for i in 0..n {
246        for j in 0..m {
247            let x = gauss_x.x[i];
248            let y = gauss_y.x[j];
249
250            // Use T type directly for kernel computation
251            // Note: gauss_x and gauss_y should already be scaled to [0, 1] interval
252            result[[i, j]] = kernel.compute_reduced(x, y, symmetry);
253        }
254    }
255
256    DiscretizedKernel::new_legacy(result, gauss_x.clone(), gauss_y.clone())
257}
258
259/// Compute matrix from Gauss quadrature rules for non-centrosymmetric kernels
260///
261/// This function evaluates the kernel directly at all combinations of Gauss points
262/// without exploiting symmetry. It works with the full domain [-xmax, xmax] × [-ymax, ymax].
263///
264/// # Arguments
265///
266/// * `kernel` - The kernel implementing AbstractKernel
267/// * `gauss_x` - Gauss quadrature rule for x coordinates (full domain)
268/// * `gauss_y` - Gauss quadrature rule for y coordinates (full domain)
269/// * `hints` - SVE hints providing segment information
270///
271/// # Returns
272///
273/// DiscretizedKernel containing the matrix, quadrature rules, and segments
274pub fn matrix_from_gauss_noncentrosymmetric<
275    T: CustomNumeric + Clone + Send + Sync,
276    K: AbstractKernel + KernelProperties,
277    H: crate::kernel::SVEHints<T>,
278>(
279    kernel: &K,
280    gauss_x: &Rule<T>,
281    gauss_y: &Rule<T>,
282    hints: &H,
283) -> DiscretizedKernel<T> {
284    let segments_x = hints.segments_x();
285    let segments_y = hints.segments_y();
286
287    let n = gauss_x.x.len();
288    let m = gauss_y.x.len();
289    let mut result = Mat::<T>::from_elem([n, m], T::zero());
290
291    // Evaluate kernel directly at all combinations of Gauss points
292    for i in 0..n {
293        for j in 0..m {
294            let x = gauss_x.x[i];
295            let y = gauss_y.x[j];
296
297            // Direct kernel evaluation (no symmetry exploitation)
298            result[[i, j]] = kernel.compute(x, y);
299        }
300    }
301
302    DiscretizedKernel::new(
303        result,
304        gauss_x.clone(),
305        gauss_y.clone(),
306        segments_x,
307        segments_y,
308    )
309}
310
311#[cfg(test)]
312#[path = "kernelmatrix_tests.rs"]
313mod tests;