Skip to main content

sparse_ir_basis/sve/
types.rs

1//! Type definitions for SVE computation
2
3use simba::scalar::ComplexField;
4
5/// Working precision type for SVE computations
6///
7/// Values match the C-API constants defined in sparseir.h
8#[derive(Debug, Clone, Copy, PartialEq, Eq)]
9pub enum TworkType {
10    /// Use double precision (64-bit)
11    Float64 = 0, // SPIR_TWORK_FLOAT64
12    /// Use extended precision (128-bit double-double)
13    Float64X2 = 1, // SPIR_TWORK_FLOAT64X2
14    /// Automatically choose precision based on epsilon
15    Auto = -1, // SPIR_TWORK_AUTO
16}
17
18/// SVD computation strategy
19///
20/// Values match the C-API constants defined in sparseir.h
21#[derive(Debug, Clone, Copy, PartialEq, Eq)]
22pub enum SVDStrategy {
23    /// Fast computation
24    Fast = 0, // SPIR_SVDSTRAT_FAST
25    /// Accurate computation
26    Accurate = 1, // SPIR_SVDSTRAT_ACCURATE
27    /// Automatically choose strategy
28    Auto = -1, // SPIR_SVDSTRAT_AUTO
29}
30
31/// Determine safe epsilon and working precision
32///
33/// This function determines the safe epsilon value based on the working precision,
34/// and automatically selects the working precision if TworkType::Auto is specified.
35///
36/// # Arguments
37///
38/// * `epsilon` - Required accuracy (non-negative); `None` selects the best
39///   accuracy of the working precision
40/// * `twork` - Working precision type (Auto for automatic selection)
41/// * `svd_strategy` - SVD computation strategy (Auto for automatic selection)
42///
43/// # Returns
44///
45/// Tuple of (safe_epsilon, actual_twork, actual_svd_strategy)
46///
47/// # Panics
48///
49/// Panics if epsilon is negative or NaN. [`compute_sve`](crate::sve::compute_sve)
50/// checks it first.
51pub(crate) fn safe_epsilon(
52    epsilon: Option<f64>,
53    twork: TworkType,
54    svd_strategy: SVDStrategy,
55) -> (f64, TworkType, SVDStrategy) {
56    // Check for a negative or NaN epsilon (following the C++ implementation
57    // for negative values; NaN used to select the automatic accuracy)
58    if let Some(eps) = epsilon.filter(|eps| !(*eps >= 0.0)) {
59        panic!("eps_required must be non-negative, got {eps:?}");
60    }
61
62    // First, choose the working dtype based on the eps required
63    let twork_actual = match twork {
64        TworkType::Auto => match epsilon {
65            Some(eps) if eps >= 1e-8 => TworkType::Float64,
66            _ => TworkType::Float64X2, // MAX_DTYPE equivalent
67        },
68        other => other,
69    };
70
71    // Next, work out the actual epsilon.
72    // The precision floor is the smallest epsilon achievable with the chosen
73    // working type.  The returned epsilon is the *larger* of the user's
74    // request and this floor so that (a) we never promise more accuracy than
75    // the arithmetic can deliver and (b) a user who asks for *less* accuracy
76    // actually gets what they asked for.
77    let precision_floor = match twork_actual {
78        TworkType::Float64 => {
79            // This is technically a bit too low (the true value is about 1.5e-8),
80            // but it's not too far off and easier to remember for the user.
81            1e-8
82        }
83        TworkType::Float64X2 => {
84            // sqrt(Df64 epsilon) ≈ sqrt(2.465e-32) ≈ 1.57e-16
85            use crate::numeric::CustomNumeric;
86            crate::Df64::epsilon().sqrt().to_f64()
87        }
88        _ => 1e-8,
89    };
90    let safe_eps = epsilon.map_or(precision_floor, |eps| eps.max(precision_floor));
91
92    // Work out the SVD strategy to be used
93    let svd_strategy_actual = match svd_strategy {
94        SVDStrategy::Auto => match epsilon {
95            // TODO: Add warning output like C++
96            Some(eps) if eps < safe_eps => SVDStrategy::Accurate,
97            _ => SVDStrategy::Fast,
98        },
99        other => other,
100    };
101
102    (safe_eps, twork_actual, svd_strategy_actual)
103}
104
105#[cfg(test)]
106mod tests {
107    use super::*;
108
109    #[test]
110    fn test_safe_epsilon_auto_float64() {
111        // epsilon=1e-7 > floor=1e-8 → safe_eps should honour the user's request
112        let (safe_eps, twork, _) = safe_epsilon(Some(1e-7), TworkType::Auto, SVDStrategy::Auto);
113        assert_eq!(twork, TworkType::Float64);
114        assert_eq!(safe_eps, 1e-7);
115    }
116
117    #[test]
118    fn test_safe_epsilon_auto_float64x2() {
119        // epsilon=1e-10 > floor≈1.57e-16 → safe_eps should be the user's epsilon
120        let (safe_eps, twork, _) = safe_epsilon(Some(1e-10), TworkType::Auto, SVDStrategy::Auto);
121        assert_eq!(twork, TworkType::Float64X2);
122        assert_eq!(safe_eps, 1e-10);
123    }
124
125    #[test]
126    fn test_safe_epsilon_explicit_precision() {
127        // epsilon=1e-7 > floor≈1.57e-16 → safe_eps should honour the user's epsilon
128        let (safe_eps, twork, _) =
129            safe_epsilon(Some(1e-7), TworkType::Float64X2, SVDStrategy::Auto);
130        assert_eq!(twork, TworkType::Float64X2);
131        assert_eq!(safe_eps, 1e-7);
132    }
133
134    #[test]
135    fn test_svd_strategy_auto_accurate() {
136        // epsilon = 1e-20 < 1.57e-16 (safe_eps for Float64X2) → Accurate
137        let (_, _, strategy) = safe_epsilon(Some(1e-20), TworkType::Auto, SVDStrategy::Auto);
138        assert_eq!(strategy, SVDStrategy::Accurate);
139    }
140
141    #[test]
142    fn test_svd_strategy_auto_fast() {
143        let (_, _, strategy) = safe_epsilon(Some(1e-7), TworkType::Auto, SVDStrategy::Auto);
144        assert_eq!(strategy, SVDStrategy::Fast);
145    }
146
147    #[test]
148    #[should_panic(expected = "eps_required must be non-negative")]
149    fn test_negative_epsilon_panics() {
150        safe_epsilon(Some(-1.0), TworkType::Auto, SVDStrategy::Auto);
151    }
152
153    /// NaN is not an accuracy: it used to select the automatic accuracy.
154    /// compute_sve rejects it first; the public helper panics like it does
155    /// for a negative epsilon.
156    #[test]
157    #[should_panic(expected = "eps_required must be non-negative, got NaN")]
158    fn test_nan_epsilon_panics() {
159        safe_epsilon(Some(f64::NAN), TworkType::Auto, SVDStrategy::Auto);
160    }
161}