1use 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
18pub trait SVEStrategy<T: CustomNumeric> {
20 fn matrices(&self) -> Vec<Mat<T>>;
22
23 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
36pub 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 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 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 fn postprocess_block(&self, u: &Mat<T>, s: &[T], v: &Mat<T>) -> Result<SvdBlock, Error> {
119 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 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 Ok((u_polys, s.iter().map(|&x| x.to_f64()).collect(), v_polys))
141 }
142}
143
144pub 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 #[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 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 pub fn new(kernel: K, epsilon: f64) -> Result<Self, Error> {
184 let hints = kernel.sve_hints::<T>(epsilon);
185
186 let segments_x = hints.segments_x();
188 let segments_y = hints.segments_y();
189 let n_gauss = hints.ngauss();
190
191 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 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 fn compute_reduced_matrix(&self, symmetry: SymmetryType) -> Mat<T> {
221 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 discretized.apply_weights_for_sve()
234 }
235
236 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 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 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 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 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_blocks(result_even_full, result_odd_full, self.epsilon)
290 }
291}
292
293#[derive(Debug, Clone)]
302struct FullDomainHints<H> {
303 inner: H,
304 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#[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 segments_x: Vec<T>,
371 segments_y: Vec<T>,
372 gauss_x: Rule<T>,
373 gauss_y: Rule<T>,
374
375 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 pub fn new(kernel: K, epsilon: f64) -> Result<Self, Error> {
392 let hints = FullDomainHints {
395 inner: kernel.sve_hints::<T>(epsilon),
396 half_domain: kernel.is_centrosymmetric(),
397 };
398
399 let segments_x = hints.segments_x();
401 let segments_y = hints.segments_y();
402 let n_gauss = hints.ngauss();
403
404 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 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 fn compute_matrix(&self) -> Mat<T> {
434 let discretized = matrix_from_gauss_noncentrosymmetric(
436 &self.kernel,
437 &self.gauss_x,
438 &self.gauss_y,
439 &self.hints,
440 );
441
442 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 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 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 SVEResult::new(
474 PiecewiseLegendrePolyVector::from_polys_unchecked(u_polys),
475 s,
476 PiecewiseLegendrePolyVector::from_polys_unchecked(v_polys),
477 self.epsilon,
478 )
479 }
480}