Skip to main content

qualia_core_db/solvers/calculus/
ode_adaptive.rs

1//! Caller-buffered adaptive Dormand-Prince ODE integration.
2
3#[derive(Debug, Clone, Copy, PartialEq)]
4pub enum OdeError {
5    InvalidDomain,
6    DimensionMismatch,
7    WorkspaceTooSmall { required: usize, available: usize },
8    NonFiniteDerivative,
9    StepUnderflow,
10    StepLimitExceeded,
11}
12
13#[repr(C)]
14#[derive(Debug, Clone, Copy, PartialEq)]
15pub struct AdaptiveOdeConfig {
16    pub absolute_tolerance: f64,
17    pub relative_tolerance: f64,
18    pub initial_step: f64,
19    pub minimum_step: f64,
20    pub maximum_step: f64,
21    pub max_steps: u32,
22}
23
24impl Default for AdaptiveOdeConfig {
25    fn default() -> Self {
26        Self {
27            absolute_tolerance: 1e-9,
28            relative_tolerance: 1e-7,
29            initial_step: 1e-3,
30            minimum_step: 1e-14,
31            maximum_step: 1.0,
32            max_steps: 100_000,
33        }
34    }
35}
36
37#[repr(C)]
38#[derive(Debug, Clone, Copy, PartialEq)]
39pub struct AdaptiveOdeResult {
40    pub final_time: f64,
41    pub accepted_steps: u32,
42    pub rejected_steps: u32,
43    pub derivative_evaluations: u32,
44    pub last_error_norm: f64,
45    pub last_step: f64,
46}
47
48pub const fn dopri5_workspace_len(dimension: usize) -> Option<usize> {
49    dimension.checked_mul(8)
50}
51
52pub fn integrate_dopri5<F>(
53    derivative: F,
54    state: &mut [f64],
55    t0: f64,
56    t_final: f64,
57    config: AdaptiveOdeConfig,
58    workspace: &mut [f64],
59) -> Result<AdaptiveOdeResult, OdeError>
60where
61    F: Fn(f64, &[f64], &mut [f64]) -> Result<(), OdeError>,
62{
63    let dimension = state.len();
64    let required = dopri5_workspace_len(dimension).ok_or(OdeError::InvalidDomain)?;
65    if dimension == 0
66        || !t0.is_finite()
67        || !t_final.is_finite()
68        || t_final < t0
69        || state.iter().any(|value| !value.is_finite())
70        || !valid_config(config)
71    {
72        return Err(OdeError::InvalidDomain);
73    }
74    if workspace.len() < required {
75        return Err(OdeError::WorkspaceTooSmall {
76            required,
77            available: workspace.len(),
78        });
79    }
80    if t_final == t0 {
81        return Ok(AdaptiveOdeResult {
82            final_time: t0,
83            accepted_steps: 0,
84            rejected_steps: 0,
85            derivative_evaluations: 0,
86            last_error_norm: 0.0,
87            last_step: 0.0,
88        });
89    }
90
91    let (k1, rest) = workspace.split_at_mut(dimension);
92    let (k2, rest) = rest.split_at_mut(dimension);
93    let (k3, rest) = rest.split_at_mut(dimension);
94    let (k4, rest) = rest.split_at_mut(dimension);
95    let (k5, rest) = rest.split_at_mut(dimension);
96    let (k6, rest) = rest.split_at_mut(dimension);
97    let (k7, rest) = rest.split_at_mut(dimension);
98    let temp = &mut rest[..dimension];
99
100    let mut time = t0;
101    let mut step = config.initial_step.min(config.maximum_step);
102    let mut accepted = 0;
103    let mut rejected = 0;
104    let mut evaluations = 0;
105    let mut last_error = f64::INFINITY;
106    let mut last_step = 0.0;
107
108    for _ in 0..config.max_steps {
109        if time >= t_final {
110            return Ok(AdaptiveOdeResult {
111                final_time: time,
112                accepted_steps: accepted,
113                rejected_steps: rejected,
114                derivative_evaluations: evaluations,
115                last_error_norm: last_error,
116                last_step,
117            });
118        }
119        step = step.min(t_final - time).min(config.maximum_step);
120        if step < config.minimum_step || time + step == time {
121            return Err(OdeError::StepUnderflow);
122        }
123
124        eval(&derivative, time, state, k1)?;
125        stage(state, temp, step, &[(1.0 / 5.0, k1)]);
126        eval(&derivative, time + step / 5.0, temp, k2)?;
127
128        stage(state, temp, step, &[(3.0 / 40.0, k1), (9.0 / 40.0, k2)]);
129        eval(&derivative, time + 3.0 * step / 10.0, temp, k3)?;
130
131        stage(
132            state,
133            temp,
134            step,
135            &[(44.0 / 45.0, k1), (-56.0 / 15.0, k2), (32.0 / 9.0, k3)],
136        );
137        eval(&derivative, time + 4.0 * step / 5.0, temp, k4)?;
138
139        stage(
140            state,
141            temp,
142            step,
143            &[
144                (19372.0 / 6561.0, k1),
145                (-25360.0 / 2187.0, k2),
146                (64448.0 / 6561.0, k3),
147                (-212.0 / 729.0, k4),
148            ],
149        );
150        eval(&derivative, time + 8.0 * step / 9.0, temp, k5)?;
151
152        stage(
153            state,
154            temp,
155            step,
156            &[
157                (9017.0 / 3168.0, k1),
158                (-355.0 / 33.0, k2),
159                (46732.0 / 5247.0, k3),
160                (49.0 / 176.0, k4),
161                (-5103.0 / 18656.0, k5),
162            ],
163        );
164        eval(&derivative, time + step, temp, k6)?;
165
166        stage(
167            state,
168            temp,
169            step,
170            &[
171                (35.0 / 384.0, k1),
172                (500.0 / 1113.0, k3),
173                (125.0 / 192.0, k4),
174                (-2187.0 / 6784.0, k5),
175                (11.0 / 84.0, k6),
176            ],
177        );
178        eval(&derivative, time + step, temp, k7)?;
179        evaluations += 7;
180
181        let mut error_norm = 0.0_f64;
182        for index in 0..dimension {
183            let fifth = state[index]
184                + step
185                    * (35.0 / 384.0 * k1[index]
186                        + 500.0 / 1113.0 * k3[index]
187                        + 125.0 / 192.0 * k4[index]
188                        - 2187.0 / 6784.0 * k5[index]
189                        + 11.0 / 84.0 * k6[index]);
190            let fourth = state[index]
191                + step
192                    * (5179.0 / 57600.0 * k1[index]
193                        + 7571.0 / 16695.0 * k3[index]
194                        + 393.0 / 640.0 * k4[index]
195                        - 92097.0 / 339200.0 * k5[index]
196                        + 187.0 / 2100.0 * k6[index]
197                        + 1.0 / 40.0 * k7[index]);
198            let scale = config.absolute_tolerance
199                + config.relative_tolerance * state[index].abs().max(fifth.abs());
200            error_norm = error_norm.max((fifth - fourth).abs() / scale);
201            temp[index] = fifth;
202        }
203        if !error_norm.is_finite() || temp.iter().any(|value| !value.is_finite()) {
204            return Err(OdeError::NonFiniteDerivative);
205        }
206
207        last_error = error_norm;
208        let factor = if error_norm == 0.0 {
209            5.0
210        } else {
211            (0.9 * error_norm.powf(-0.2)).clamp(0.2, 5.0)
212        };
213        if error_norm <= 1.0 {
214            state.copy_from_slice(temp);
215            time += step;
216            last_step = step;
217            accepted += 1;
218        } else {
219            rejected += 1;
220        }
221        step = (step * factor).clamp(config.minimum_step, config.maximum_step);
222    }
223
224    Err(OdeError::StepLimitExceeded)
225}
226
227fn valid_config(config: AdaptiveOdeConfig) -> bool {
228    config.absolute_tolerance.is_finite()
229        && config.absolute_tolerance > 0.0
230        && config.relative_tolerance.is_finite()
231        && config.relative_tolerance >= 0.0
232        && config.initial_step.is_finite()
233        && config.initial_step > 0.0
234        && config.minimum_step.is_finite()
235        && config.minimum_step > 0.0
236        && config.maximum_step.is_finite()
237        && config.maximum_step >= config.minimum_step
238        && config.max_steps > 0
239}
240
241fn eval<F>(derivative: &F, time: f64, state: &[f64], output: &mut [f64]) -> Result<(), OdeError>
242where
243    F: Fn(f64, &[f64], &mut [f64]) -> Result<(), OdeError>,
244{
245    derivative(time, state, output)?;
246    if output.len() != state.len() {
247        return Err(OdeError::DimensionMismatch);
248    }
249    if output.iter().any(|value| !value.is_finite()) {
250        return Err(OdeError::NonFiniteDerivative);
251    }
252    Ok(())
253}
254
255fn stage(state: &[f64], output: &mut [f64], step: f64, terms: &[(f64, &[f64])]) {
256    for index in 0..state.len() {
257        let mut value = state[index];
258        for (coefficient, derivative) in terms {
259            value += step * coefficient * derivative[index];
260        }
261        output[index] = value;
262    }
263}
264
265#[cfg(test)]
266mod tests {
267    use super::*;
268
269    fn decay(_time: f64, state: &[f64], output: &mut [f64]) -> Result<(), OdeError> {
270        if state.len() != output.len() {
271            return Err(OdeError::DimensionMismatch);
272        }
273        for (out, value) in output.iter_mut().zip(state) {
274            *out = -*value;
275        }
276        Ok(())
277    }
278
279    #[test]
280    fn dopri5_lands_exactly_at_final_time_with_tolerance_control() {
281        let mut state = [1.0, 2.0, 3.0];
282        let mut workspace = [0.0; 24];
283        let report = integrate_dopri5(
284            decay,
285            &mut state,
286            0.0,
287            2.0,
288            AdaptiveOdeConfig::default(),
289            &mut workspace,
290        )
291        .unwrap();
292        assert_eq!(report.final_time, 2.0);
293        for (index, value) in state.iter().enumerate() {
294            let expected = (index + 1) as f64 * (-2.0_f64).exp();
295            assert!((value - expected).abs() < 3e-8);
296        }
297        assert!(report.accepted_steps > 0);
298    }
299
300    #[test]
301    fn dopri5_rejects_workspace_and_non_finite_derivatives() {
302        let mut state = [1.0, 2.0];
303        let mut too_small = [0.0; 15];
304        assert_eq!(
305            integrate_dopri5(
306                decay,
307                &mut state,
308                0.0,
309                1.0,
310                AdaptiveOdeConfig::default(),
311                &mut too_small,
312            ),
313            Err(OdeError::WorkspaceTooSmall {
314                required: 16,
315                available: 15
316            })
317        );
318
319        let mut workspace = [0.0; 16];
320        assert_eq!(
321            integrate_dopri5(
322                |_t, _y, out| {
323                    out.fill(f64::NAN);
324                    Ok(())
325                },
326                &mut state,
327                0.0,
328                1.0,
329                AdaptiveOdeConfig::default(),
330                &mut workspace,
331            ),
332            Err(OdeError::NonFiniteDerivative)
333        );
334    }
335}