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;