sparse_ir_basis/sve/compute.rs
1//! Main SVE computation functions
2
3use crate::error::{Error, require_accuracy, require_nonzero_size};
4use crate::fpu_check::FpuGuard;
5use crate::kernel::{AbstractKernel, CentrosymmKernel, KernelProperties, SVEHints};
6use crate::matrix::Mat;
7use crate::numeric::CustomNumeric;
8
9use super::result::SVEResult;
10use super::strategy::{CentrosymmSVE, NonCentrosymmSVE, SVEStrategy};
11use super::types::{SVDStrategy, TworkType, safe_epsilon};
12
13/// Default relative cutoff for singular value truncation: `2 * T::epsilon()`
14///
15/// This is `2^-51` (about 4.44e-16) for `f64` and `2^-104` (about 4.93e-32)
16/// for `Df64`. Both SVE paths use it when `cutoff` is `None`.
17/// Convention-matched with libsparseir, whose `pre_postprocess` in
18/// `backend/cxx/include/sparseir/impl/sve_impl.ipp` (commit 4bc58ea) uses
19/// `T(2) * std::numeric_limits<T>::epsilon()` when the cutoff is NaN.
20fn default_cutoff<T: CustomNumeric>() -> T {
21 T::from_f64_unchecked(2.0) * T::epsilon()
22}
23
24/// Release unused memory back to the OS.
25///
26/// SVE computation allocates large temporary buffers for SVD.
27/// After computation, these are freed but the allocator may retain them.
28/// This function asks the allocator to return unused memory to the OS.
29///
30/// Supported platforms:
31/// - macOS: uses `malloc_zone_pressure_relief`
32/// - Linux (glibc): uses `malloc_trim`
33/// - Other platforms: no-op (memory is still freed, just not returned to OS immediately)
34#[inline]
35fn release_unused_memory() {
36 #[cfg(target_os = "macos")]
37 {
38 unsafe extern "C" {
39 fn malloc_zone_pressure_relief(zone: *mut std::ffi::c_void, goal: usize) -> usize;
40 }
41 unsafe { malloc_zone_pressure_relief(std::ptr::null_mut(), 0) };
42 }
43
44 // Only use malloc_trim on Linux with glibc (not musl)
45 #[cfg(all(target_os = "linux", target_env = "gnu"))]
46 {
47 unsafe extern "C" {
48 fn malloc_trim(pad: usize) -> i32;
49 }
50 unsafe { malloc_trim(0) };
51 }
52}
53
54/// Main SVE computation function for centrosymmetric kernels
55///
56/// Automatically chooses the appropriate SVE strategy based on kernel properties
57/// and working precision based on epsilon.
58///
59/// # Arguments
60///
61/// * `kernel` - The centrosymmetric kernel to expand
62/// * `epsilon` - Required accuracy, in (0, 1). `None` selects the best
63/// accuracy of the working precision (about 1.6e-16 in Float64X2).
64/// * `cutoff` - Relative tolerance for singular value truncation: singular
65/// values smaller than `cutoff` times the largest singular value are
66/// discarded. `None` selects `2 * machine epsilon` of the working precision,
67/// about 4.44e-16 for Float64 and 4.93e-32 for Float64X2 (the libsparseir
68/// default). The SVD of each even/odd block already discards singular values
69/// below `2 * machine epsilon` times that block's largest singular value, so
70/// a smaller `cutoff` has little effect. Must be in [0, 1].
71/// * `max_num_svals` - Maximum number of singular values to keep
72/// * `twork` - Working precision type (Auto for automatic selection)
73///
74/// # Returns
75///
76/// SVEResult containing singular functions and values
77///
78/// # Errors
79///
80/// * [`Error::InvalidParameter`] if `epsilon` is not in (0, 1), `cutoff` is
81/// not in [0, 1], or `max_num_svals` is `Some(0)`; checked before any work
82/// * [`Error::NonFiniteInput`] if the discretized kernel has a NaN or
83/// infinite entry (e.g. a cutoff Λ so small that 1/Λ overflows)
84/// * [`Error::DecompositionFailed`] if an SVD fails
85///
86/// # FPU State Warning
87///
88/// This function checks for dangerous FPU settings (Flush-to-Zero and Denormals-Are-Zero)
89/// that can cause incorrect results. If detected, it temporarily corrects the FPU state
90/// and prints a warning. If you see this warning, add `-fp-model precise` flag when
91/// compiling with Intel Fortran.
92pub fn compute_sve<K>(
93 kernel: K,
94 epsilon: Option<f64>,
95 cutoff: Option<f64>,
96 max_num_svals: Option<usize>,
97 twork: TworkType,
98) -> Result<SVEResult, Error>
99where
100 K: CentrosymmKernel + KernelProperties + Clone + 'static,
101{
102 check_sve_parameters(epsilon, cutoff, max_num_svals)?;
103
104 // Protect computation from dangerous FPU settings (FZ/DAZ)
105 // This temporarily disables FZ/DAZ and restores them after computation
106 let _fpu_guard = FpuGuard::new_protect_computation();
107
108 // Determine safe epsilon and working precision
109 let (safe_epsilon, twork_actual, _svd_strategy) =
110 safe_epsilon(epsilon, twork, SVDStrategy::Auto);
111
112 // Dispatch based on working precision
113 let result = match twork_actual {
114 TworkType::Float64 => {
115 compute_sve_with_precision::<f64, K>(kernel, safe_epsilon, cutoff, max_num_svals)
116 }
117 TworkType::Float64X2 => compute_sve_with_precision::<crate::Df64, K>(
118 kernel,
119 safe_epsilon,
120 cutoff.map(crate::Df64::from),
121 max_num_svals,
122 ),
123 _ => panic!("Invalid TworkType: {:?}", twork_actual),
124 };
125
126 // Release temporary memory back to OS after SVE computation
127 release_unused_memory();
128
129 result
130}
131
132/// Main SVE computation function for general kernels (centrosymmetric or non-centrosymmetric)
133///
134/// Discretizes the kernel on its full domain `[-xmax, xmax] × [-ymax, ymax]`
135/// with [`NonCentrosymmSVE`], which only requires [`AbstractKernel`].
136///
137/// Centrosymmetric kernels are expanded correctly as well: their half-domain
138/// [`SVEHints`] segments are mirrored onto the full domain, so the singular
139/// values agree with [`compute_sve`] up to rounding. The even/odd block
140/// structure is not exploited, however: the SVD is taken of one matrix with
141/// twice as many rows and columns as each block used by [`compute_sve`]
142/// (asymptotically about four times the work), the singular functions carry no
143/// parity tag, and the singular functions of (nearly) degenerate singular
144/// values may mix the even and odd sectors. For kernels implementing
145/// [`CentrosymmKernel`], such as [`LogisticKernel`](crate::kernel::LogisticKernel)
146/// and [`RegularizedBoseKernel`](crate::kernel::RegularizedBoseKernel), prefer
147/// [`compute_sve`].
148///
149/// This function cannot select [`CentrosymmSVE`] by itself even when
150/// [`AbstractKernel::is_centrosymmetric`] returns true: that strategy needs
151/// the reduced kernels of [`CentrosymmKernel::compute_reduced`], and a
152/// `K: AbstractKernel` bound cannot be refined to `K: CentrosymmKernel`
153/// without trait specialization, which stable Rust does not provide.
154///
155/// # Arguments
156///
157/// * `kernel` - The kernel to expand (can be centrosymmetric or non-centrosymmetric)
158/// * `epsilon` - Required accuracy, in (0, 1). `None` selects the best
159/// accuracy of the working precision (about 1.6e-16 in Float64X2).
160/// * `cutoff` - Relative tolerance for singular value truncation: singular
161/// values smaller than `cutoff` times the largest singular value are
162/// discarded. `None` selects `2 * machine epsilon` of the working precision,
163/// about 4.44e-16 for Float64 and 4.93e-32 for Float64X2, as in
164/// [`compute_sve`]. The SVD of the full-domain matrix already discards
165/// singular values below `2 * machine epsilon` times the largest one, so a
166/// smaller `cutoff` has no effect. Must be in [0, 1].
167/// * `max_num_svals` - Maximum number of singular values to keep
168/// * `twork` - Working precision type (Auto for automatic selection)
169///
170/// # Returns
171///
172/// SVEResult containing singular functions and values
173///
174/// # Errors
175///
176/// * [`Error::InvalidParameter`] if `epsilon` is not in (0, 1), `cutoff` is
177/// not in [0, 1], or `max_num_svals` is `Some(0)`; checked before any work
178/// * [`Error::NonFiniteInput`] if the discretized kernel has a NaN or
179/// infinite entry (e.g. a cutoff Λ so small that 1/Λ overflows)
180/// * [`Error::DecompositionFailed`] if an SVD fails
181///
182/// # FPU State Warning
183///
184/// This function checks for dangerous FPU settings (Flush-to-Zero and Denormals-Are-Zero)
185/// that can cause incorrect results. If detected, it temporarily corrects the FPU state
186/// and prints a warning. If you see this warning, add `-fp-model precise` flag when
187/// compiling with Intel Fortran.
188pub fn compute_sve_general<K>(
189 kernel: K,
190 epsilon: Option<f64>,
191 cutoff: Option<f64>,
192 max_num_svals: Option<usize>,
193 twork: TworkType,
194) -> Result<SVEResult, Error>
195where
196 K: AbstractKernel + KernelProperties + Clone + 'static,
197{
198 check_sve_parameters(epsilon, cutoff, max_num_svals)?;
199
200 // Protect computation from dangerous FPU settings (FZ/DAZ)
201 // This temporarily disables FZ/DAZ and restores them after computation
202 let _fpu_guard = FpuGuard::new_protect_computation();
203
204 // Determine safe epsilon and working precision
205 let (safe_epsilon, twork_actual, _svd_strategy) =
206 safe_epsilon(epsilon, twork, SVDStrategy::Auto);
207
208 // Dispatch based on working precision
209 let result = match twork_actual {
210 TworkType::Float64 => compute_sve_general_with_precision::<f64, K>(
211 kernel,
212 safe_epsilon,
213 cutoff,
214 max_num_svals,
215 ),
216 TworkType::Float64X2 => compute_sve_general_with_precision::<crate::Df64, K>(
217 kernel,
218 safe_epsilon,
219 cutoff.map(crate::Df64::from),
220 max_num_svals,
221 ),
222 _ => panic!("Invalid TworkType: {:?}", twork_actual),
223 };
224
225 // Release temporary memory back to OS after SVE computation
226 release_unused_memory();
227
228 result
229}
230
231/// Check the parameters shared by [`compute_sve`] and [`compute_sve_general`]
232fn check_sve_parameters(
233 epsilon: Option<f64>,
234 cutoff: Option<f64>,
235 max_num_svals: Option<usize>,
236) -> Result<(), Error> {
237 require_accuracy("epsilon", epsilon)?;
238 if let Some(c) = cutoff {
239 if !(0.0..=1.0).contains(&c) {
240 return Err(Error::InvalidParameter {
241 name: "cutoff",
242 value: format!("{c:?}"),
243 reason: "must be in [0, 1]".to_string(),
244 });
245 }
246 }
247 require_nonzero_size("max_num_svals", max_num_svals)
248}
249
250/// Compute SVE with specific precision type
251fn compute_sve_with_precision<T, K>(
252 kernel: K,
253 epsilon: f64,
254 cutoff: Option<T>,
255 max_num_svals: Option<usize>,
256) -> Result<SVEResult, Error>
257where
258 T: CustomNumeric + Send + Sync + Clone + 'static,
259 K: CentrosymmKernel + KernelProperties + Clone + 'static,
260 K::SVEHintsType<T>: SVEHints<T> + Clone,
261{
262 // 1. Determine SVE strategy (automatically chooses CentrosymmSVE for centrosymmetric kernels)
263 let sve = determine_sve::<T, K>(kernel, epsilon)?;
264
265 // 2. Compute matrices
266 let matrices = sve.matrices();
267
268 // 3. Compute SVD for each matrix
269 let mut u_list = Vec::new();
270 let mut s_list = Vec::new();
271 let mut v_list = Vec::new();
272
273 for matrix in matrices.iter() {
274 let (u, s, v) = crate::tsvd::compute_svd_dtensor(matrix)?;
275 u_list.push(u);
276 s_list.push(s);
277 v_list.push(v);
278 }
279
280 // 4. Truncate based on cutoff (default: 2 * T::epsilon(), see default_cutoff)
281 let rtol_t = cutoff.unwrap_or_else(default_cutoff::<T>);
282 let (u_trunc, s_trunc, v_trunc) = truncate(u_list, s_list, v_list, rtol_t, max_num_svals);
283
284 // 5. Post-process to create SVEResult
285 sve.postprocess(u_trunc, s_trunc, v_trunc)
286}
287
288/// Compute SVE with specific precision type for general kernels
289fn compute_sve_general_with_precision<T, K>(
290 kernel: K,
291 epsilon: f64,
292 cutoff: Option<T>,
293 max_num_svals: Option<usize>,
294) -> Result<SVEResult, Error>
295where
296 T: CustomNumeric + Send + Sync + Clone + 'static,
297 K: AbstractKernel + KernelProperties + Clone + 'static,
298 K::SVEHintsType<T>: SVEHints<T> + Clone,
299{
300 // 1. Determine SVE strategy (full-domain NonCentrosymmSVE)
301 let sve = determine_sve_general::<T, K>(kernel, epsilon)?;
302
303 // 2. Compute matrices
304 let matrices = sve.matrices();
305
306 // 3. Compute SVD for each matrix
307 let mut u_list = Vec::new();
308 let mut s_list = Vec::new();
309 let mut v_list = Vec::new();
310
311 for matrix in matrices.iter() {
312 let (u, s, v) = crate::tsvd::compute_svd_dtensor(matrix)?;
313 u_list.push(u);
314 s_list.push(s);
315 v_list.push(v);
316 }
317
318 // 4. Truncate based on cutoff (default: 2 * T::epsilon(), see default_cutoff)
319 let rtol_t = cutoff.unwrap_or_else(default_cutoff::<T>);
320 let (u_trunc, s_trunc, v_trunc) = truncate(u_list, s_list, v_list, rtol_t, max_num_svals);
321
322 // 5. Post-process to create SVEResult
323 sve.postprocess(u_trunc, s_trunc, v_trunc)
324}
325
326/// Determine the appropriate SVE strategy
327///
328/// For centrosymmetric kernels, uses CentrosymmSVE for efficient computation
329/// by exploiting even/odd symmetry.
330fn determine_sve<T, K>(kernel: K, epsilon: f64) -> Result<Box<dyn SVEStrategy<T>>, Error>
331where
332 T: CustomNumeric + Send + Sync + Clone + 'static,
333 K: CentrosymmKernel + KernelProperties + Clone + 'static,
334 K::SVEHintsType<T>: SVEHints<T> + Clone,
335{
336 // CentrosymmKernel trait implies centrosymmetric
337 Ok(Box::new(CentrosymmSVE::new(kernel, epsilon)?))
338}
339
340/// Determine the SVE strategy for general kernels
341///
342/// Always uses [`NonCentrosymmSVE`], which handles centrosymmetric kernels by
343/// mirroring their half-domain hint segments onto the full domain.
344/// [`CentrosymmSVE`] would require `K: CentrosymmKernel`, which cannot be
345/// recovered from the `K: AbstractKernel` bound (see [`compute_sve_general`]).
346fn determine_sve_general<T, K>(kernel: K, epsilon: f64) -> Result<Box<dyn SVEStrategy<T>>, Error>
347where
348 T: CustomNumeric + Send + Sync + Clone + 'static,
349 K: AbstractKernel + KernelProperties + Clone + 'static,
350 K::SVEHintsType<T>: SVEHints<T> + Clone,
351{
352 Ok(Box::new(NonCentrosymmSVE::new(kernel, epsilon)?))
353}
354
355/// Truncate SVD results based on cutoff and maximum size
356///
357/// # Arguments
358///
359/// * `u_list` - List of U matrices
360/// * `s_list` - List of singular value vectors
361/// * `v_list` - List of V matrices
362/// * `rtol` - Relative tolerance for truncation
363/// * `max_num_svals` - Maximum number of singular values to keep
364///
365/// # Returns
366///
367/// Tuple of (truncated_u_list, truncated_s_list, truncated_v_list)
368pub(crate) fn truncate<T: CustomNumeric>(
369 u_list: Vec<Mat<T>>,
370 s_list: Vec<Vec<T>>,
371 v_list: Vec<Mat<T>>,
372 rtol: T,
373 max_num_svals: Option<usize>,
374) -> (Vec<Mat<T>>, Vec<Vec<T>>, Vec<Mat<T>>) {
375 let zero = T::zero();
376
377 // Validate
378 if let Some(max) = max_num_svals {
379 if max == 0 {
380 panic!("max_num_svals must be positive");
381 }
382 }
383 if rtol < zero || rtol > T::from_f64_unchecked(1.0) {
384 panic!("rtol must be in [0, 1]");
385 }
386
387 // Find global maximum singular value
388 let mut all_svals = Vec::new();
389 for s in &s_list {
390 all_svals.extend(s.iter().copied());
391 }
392
393 let max_sval = all_svals
394 .iter()
395 .max_by(|a, b| a.partial_cmp(b).unwrap_or(std::cmp::Ordering::Equal))
396 .copied()
397 .unwrap_or(zero);
398
399 // Determine cutoff
400 let cutoff = if let Some(max_count) = max_num_svals {
401 if max_count < all_svals.len() {
402 let mut sorted = all_svals.clone();
403 sorted.sort_by(|a, b| b.partial_cmp(a).unwrap_or(std::cmp::Ordering::Equal));
404 let nth = sorted[max_count - 1];
405 if rtol * max_sval > nth {
406 rtol * max_sval
407 } else {
408 nth
409 }
410 } else {
411 rtol * max_sval
412 }
413 } else {
414 rtol * max_sval
415 };
416
417 // Truncate each result
418 let mut u_trunc = Vec::new();
419 let mut s_trunc = Vec::new();
420 let mut v_trunc = Vec::new();
421
422 for i in 0..s_list.len() {
423 let s = &s_list[i];
424 let u = &u_list[i];
425 let v = &v_list[i];
426
427 // Count singular values above cutoff
428 let mut n_keep = 0;
429 for &val in s.iter() {
430 if val >= cutoff {
431 n_keep += 1;
432 }
433 }
434
435 // Slice U: keep first n_keep columns. Preserve one output block per
436 // input block even when all singular values in that block are removed.
437 let u_shape = *u.shape();
438 let u_sliced = Mat::<T>::from_fn([u_shape.0, n_keep], |idx| u[[idx[0], idx[1]]]);
439 u_trunc.push(u_sliced);
440
441 s_trunc.push(s[..n_keep].to_vec());
442
443 // Slice V: keep first n_keep columns
444 let v_shape = *v.shape();
445 let v_sliced = Mat::<T>::from_fn([v_shape.0, n_keep], |idx| v[[idx[0], idx[1]]]);
446 v_trunc.push(v_sliced);
447 }
448
449 (u_trunc, s_trunc, v_trunc)
450}
451
452#[cfg(test)]
453mod tests {
454 use super::*;
455
456 /// The default cutoff is exactly `2 * machine epsilon` of each working
457 /// precision (issue #249); `Df64` epsilon is `f64::EPSILON^2 / 2 = 2^-105`.
458 #[test]
459 fn test_default_cutoff_is_two_machine_epsilon() {
460 assert_eq!(default_cutoff::<f64>(), 2.0 * f64::EPSILON);
461 assert_eq!(default_cutoff::<f64>(), 2f64.powi(-51));
462
463 let df64 = default_cutoff::<crate::Df64>();
464 assert_eq!(df64, crate::Df64::from(2f64.powi(-104)));
465 assert_eq!((df64.hi(), df64.lo()), (2f64.powi(-104), 0.0));
466 }
467
468 #[test]
469 fn test_truncate_by_rtol() {
470 let u = vec![Mat::<f64>::from_elem([3, 3], 1.0)];
471 let s = vec![vec![10.0, 5.0, 0.1]];
472 let v = vec![Mat::<f64>::from_elem([3, 3], 1.0)];
473
474 // rtol = 0.1, max_sval = 10.0, cutoff = 1.0
475 // Keep values >= 1.0: [10.0, 5.0]
476 let (_, s_trunc, _) = truncate(u, s, v, 0.1, None);
477
478 assert_eq!(s_trunc[0].len(), 2);
479 assert_eq!(s_trunc[0][0], 10.0);
480 assert_eq!(s_trunc[0][1], 5.0);
481 }
482
483 #[test]
484 fn test_truncate_by_max_size() {
485 let u = vec![Mat::<f64>::from_elem([3, 3], 1.0)];
486 let s = vec![vec![10.0, 5.0, 2.0]];
487 let v = vec![Mat::<f64>::from_elem([3, 3], 1.0)];
488
489 // max_num_svals = 2
490 let (_, s_trunc, _) = truncate(u, s, v, 0.0, Some(2));
491
492 assert_eq!(s_trunc[0].len(), 2);
493 }
494
495 #[test]
496 #[should_panic(expected = "max_num_svals must be positive")]
497 fn test_truncate_rejects_zero_max_size() {
498 let u = vec![Mat::<f64>::from_elem([1, 1], 1.0)];
499 let s = vec![vec![1.0]];
500 let v = vec![Mat::<f64>::from_elem([1, 1], 1.0)];
501
502 truncate(u, s, v, 0.0, Some(0));
503 }
504
505 #[test]
506 fn test_truncate_preserves_empty_blocks() {
507 let u = vec![
508 Mat::<f64>::from_elem([2, 1], 1.0),
509 Mat::<f64>::from_elem([2, 1], 2.0),
510 ];
511 let s = vec![vec![10.0], vec![1.0]];
512 let v = vec![
513 Mat::<f64>::from_elem([2, 1], 1.0),
514 Mat::<f64>::from_elem([2, 1], 2.0),
515 ];
516
517 let (u_trunc, s_trunc, v_trunc) = truncate(u, s, v, 0.5, None);
518
519 assert_eq!(s_trunc.len(), 2);
520 assert_eq!(s_trunc[0], vec![10.0]);
521 assert!(s_trunc[1].is_empty());
522 assert_eq!(*u_trunc[1].shape(), (2, 0));
523 assert_eq!(*v_trunc[1].shape(), (2, 0));
524 }
525
526 #[test]
527 #[should_panic(expected = "rtol must be in [0, 1]")]
528 fn test_truncate_invalid_rtol() {
529 let u = vec![Mat::<f64>::from_elem([1, 1], 1.0)];
530 let s = vec![vec![1.0]];
531 let v = vec![Mat::<f64>::from_elem([1, 1], 1.0)];
532
533 truncate(u, s, v, 1.5, None);
534 }
535}