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}