Skip to main content

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}