Skip to main content

zenith_float_num/
ode.rs

1//! Fixed-step and adaptive ODE solvers on [`ExactNum`].
2
3use crate::defs::WORD_BIT_SIZE;
4use crate::Consts;
5use crate::ExactNum;
6use crate::ExactNumArray;
7use crate::RoundingMode;
8use alloc::vec::Vec;
9
10/// Maximum accepted steps for any solver in this module.
11pub const ODE_MAX_STEPS: usize = 65536;
12
13/// Default minimum step exponent: `h_min = 2^{ODE_MIN_STEP}`.
14pub const ODE_MIN_STEP: i32 = -256;
15
16/// Safety factor numerator for Dormand–Prince step updates (`9/10`).
17const ODE_FAC_NUM: i64 = 9;
18/// Safety factor denominator for Dormand–Prince step updates.
19const ODE_FAC_DEN: i64 = 10;
20/// Maximum step growth.
21const ODE_H_GROW: i64 = 5;
22/// Maximum step shrink is `1/ODE_H_SHRINK_DEN`.
23const ODE_H_SHRINK_DEN: i64 = 5;
24
25fn work_p(p: usize) -> usize {
26    p.saturating_add(WORD_BIT_SIZE)
27}
28
29fn finite(x: &ExactNum) -> bool {
30    !x.is_nan() && !x.is_inf()
31}
32
33fn frac(n: i64, d: i64, p: usize, rm: RoundingMode) -> ExactNum {
34    ExactNum::from_i64(n, p).div(&ExactNum::from_i64(d, p), p, rm)
35}
36
37/// Minimum step `2^{ODE_MIN_STEP}` at precision `p`.
38pub fn ode_min_step(p: usize, rm: RoundingMode) -> ExactNum {
39    ExactNum::from_u8(1, p).ldexp(ODE_MIN_STEP, p, rm)
40}
41
42fn to_row(p: usize, rm: RoundingMode, vals: &[ExactNum]) -> Option<ExactNumArray> {
43    let rounded: Vec<ExactNum> = vals
44        .iter()
45        .map(|x| {
46            let mut y = x.clone();
47            let _ = y.set_precision(p, rm);
48            y
49        })
50        .collect();
51    ExactNumArray::from_shape(p, 1, rounded.len(), &rounded)
52}
53
54fn push_pair(
55    ts: &mut Vec<ExactNum>,
56    ys: &mut Vec<ExactNum>,
57    t: ExactNum,
58    y: ExactNum,
59    cap: usize,
60) -> bool {
61    if ts.len() >= cap {
62        return false;
63    }
64    ts.push(t);
65    ys.push(y);
66    true
67}
68
69/// Classical RK4 for `y' = f(t, y)` on `[t0, t1]` with `n_steps` equal steps.
70///
71/// Returns row arrays `(t, y)` of length `n_steps+1`. `None` if `n_steps` is 0
72/// or greater than [`ODE_MAX_STEPS`], `t1 ≤ t0`, or a value is non-finite.
73pub fn rk4<F>(
74    mut f: F,
75    t0: &ExactNum,
76    y0: &ExactNum,
77    t1: &ExactNum,
78    n_steps: usize,
79    p: usize,
80    rm: RoundingMode,
81    cc: &mut Consts,
82) -> Option<(ExactNumArray, ExactNumArray)>
83where
84    F: FnMut(&ExactNum, &ExactNum, usize, RoundingMode, &mut Consts) -> ExactNum,
85{
86    if n_steps == 0 || n_steps > ODE_MAX_STEPS {
87        return None;
88    }
89    if !finite(t0) || !finite(y0) || !finite(t1) || t0.cmp(t1) != Some(-1) {
90        return None;
91    }
92    let wrk = work_p(p);
93    let n_f = ExactNum::from_u32(n_steps as u32, wrk);
94    let h = t1
95        .sub(t0, wrk, RoundingMode::None)
96        .div(&n_f, wrk, RoundingMode::None);
97    let two = ExactNum::from_u8(2, wrk);
98    let six = ExactNum::from_u8(6, wrk);
99    let half_h = h.div(&two, wrk, RoundingMode::None);
100    let mut t = t0.clone();
101    let mut y = y0.clone();
102    let _ = t.set_precision(wrk, RoundingMode::None);
103    let _ = y.set_precision(wrk, RoundingMode::None);
104    let mut ts = Vec::with_capacity(n_steps + 1);
105    let mut ys = Vec::with_capacity(n_steps + 1);
106    ts.push(t.clone());
107    ys.push(y.clone());
108    for k in 0..n_steps {
109        let k1 = f(&t, &y, wrk, RoundingMode::None, cc);
110        let y2 = y.add(
111            &half_h.mul(&k1, wrk, RoundingMode::None),
112            wrk,
113            RoundingMode::None,
114        );
115        let t2 = t.add(&half_h, wrk, RoundingMode::None);
116        let k2 = f(&t2, &y2, wrk, RoundingMode::None, cc);
117        let y3 = y.add(
118            &half_h.mul(&k2, wrk, RoundingMode::None),
119            wrk,
120            RoundingMode::None,
121        );
122        let k3 = f(&t2, &y3, wrk, RoundingMode::None, cc);
123        let y4 = y.add(
124            &h.mul(&k3, wrk, RoundingMode::None),
125            wrk,
126            RoundingMode::None,
127        );
128        let t4 = t.add(&h, wrk, RoundingMode::None);
129        let k4 = f(&t4, &y4, wrk, RoundingMode::None, cc);
130        if !finite(&k1) || !finite(&k2) || !finite(&k3) || !finite(&k4) {
131            return None;
132        }
133        let sum = k1
134            .add(
135                &two.mul(&k2, wrk, RoundingMode::None),
136                wrk,
137                RoundingMode::None,
138            )
139            .add(
140                &two.mul(&k3, wrk, RoundingMode::None),
141                wrk,
142                RoundingMode::None,
143            )
144            .add(&k4, wrk, RoundingMode::None);
145        y = y.add(
146            &h.div(&six, wrk, RoundingMode::None)
147                .mul(&sum, wrk, RoundingMode::None),
148            wrk,
149            RoundingMode::None,
150        );
151        t = t0.add(
152            &h.mul(
153                &ExactNum::from_u32((k + 1) as u32, wrk),
154                wrk,
155                RoundingMode::None,
156            ),
157            wrk,
158            RoundingMode::None,
159        );
160        if !finite(&t) || !finite(&y) {
161            return None;
162        }
163        ts.push(t.clone());
164        ys.push(y.clone());
165    }
166    Some((to_row(p, rm, &ts)?, to_row(p, rm, &ys)?))
167}
168
169/// Explicit Euler for `y' = f(t, y)` on `[t0, t1]` with `n_steps` equal steps.
170///
171/// Returns row arrays `(t, y)` of length `n_steps+1`. Same rejection as [`rk4`].
172pub fn euler<F>(
173    mut f: F,
174    t0: &ExactNum,
175    y0: &ExactNum,
176    t1: &ExactNum,
177    n_steps: usize,
178    p: usize,
179    rm: RoundingMode,
180) -> Option<(ExactNumArray, ExactNumArray)>
181where
182    F: FnMut(&ExactNum, &ExactNum, usize, RoundingMode) -> ExactNum,
183{
184    if n_steps == 0 || n_steps > ODE_MAX_STEPS {
185        return None;
186    }
187    if !finite(t0) || !finite(y0) || !finite(t1) || t0.cmp(t1) != Some(-1) {
188        return None;
189    }
190    let wrk = work_p(p);
191    let n_f = ExactNum::from_u32(n_steps as u32, wrk);
192    let h = t1
193        .sub(t0, wrk, RoundingMode::None)
194        .div(&n_f, wrk, RoundingMode::None);
195    let mut t = t0.clone();
196    let mut y = y0.clone();
197    let _ = t.set_precision(wrk, RoundingMode::None);
198    let _ = y.set_precision(wrk, RoundingMode::None);
199    let mut ts = Vec::with_capacity(n_steps + 1);
200    let mut ys = Vec::with_capacity(n_steps + 1);
201    ts.push(t.clone());
202    ys.push(y.clone());
203    for k in 0..n_steps {
204        let yp = f(&t, &y, wrk, RoundingMode::None);
205        if !finite(&yp) {
206            return None;
207        }
208        y = y.add(
209            &h.mul(&yp, wrk, RoundingMode::None),
210            wrk,
211            RoundingMode::None,
212        );
213        t = t0.add(
214            &h.mul(
215                &ExactNum::from_u32((k + 1) as u32, wrk),
216                wrk,
217                RoundingMode::None,
218            ),
219            wrk,
220            RoundingMode::None,
221        );
222        if !finite(&t) || !finite(&y) {
223            return None;
224        }
225        ts.push(t.clone());
226        ys.push(y.clone());
227    }
228    Some((to_row(p, rm, &ts)?, to_row(p, rm, &ys)?))
229}
230
231struct Dp45 {
232    c2: ExactNum,
233    c3: ExactNum,
234    c4: ExactNum,
235    c5: ExactNum,
236    a21: ExactNum,
237    a31: ExactNum,
238    a32: ExactNum,
239    a41: ExactNum,
240    a42: ExactNum,
241    a43: ExactNum,
242    a51: ExactNum,
243    a52: ExactNum,
244    a53: ExactNum,
245    a54: ExactNum,
246    a61: ExactNum,
247    a62: ExactNum,
248    a63: ExactNum,
249    a64: ExactNum,
250    a65: ExactNum,
251    b5_1: ExactNum,
252    b5_3: ExactNum,
253    b5_4: ExactNum,
254    b5_5: ExactNum,
255    b5_6: ExactNum,
256    b4_1: ExactNum,
257    b4_3: ExactNum,
258    b4_4: ExactNum,
259    b4_5: ExactNum,
260    b4_6: ExactNum,
261    b4_7: ExactNum,
262}
263
264impl Dp45 {
265    fn new(p: usize, rm: RoundingMode) -> Self {
266        Self {
267            c2: frac(1, 5, p, rm),
268            c3: frac(3, 10, p, rm),
269            c4: frac(4, 5, p, rm),
270            c5: frac(8, 9, p, rm),
271            a21: frac(1, 5, p, rm),
272            a31: frac(3, 40, p, rm),
273            a32: frac(9, 40, p, rm),
274            a41: frac(44, 45, p, rm),
275            a42: frac(-56, 15, p, rm),
276            a43: frac(32, 9, p, rm),
277            a51: frac(19372, 6561, p, rm),
278            a52: frac(-25360, 2187, p, rm),
279            a53: frac(64448, 6561, p, rm),
280            a54: frac(-212, 729, p, rm),
281            a61: frac(9017, 3168, p, rm),
282            a62: frac(-355, 33, p, rm),
283            a63: frac(46732, 5247, p, rm),
284            a64: frac(49, 176, p, rm),
285            a65: frac(-5103, 18656, p, rm),
286            b5_1: frac(35, 384, p, rm),
287            b5_3: frac(500, 1113, p, rm),
288            b5_4: frac(125, 192, p, rm),
289            b5_5: frac(-2187, 6784, p, rm),
290            b5_6: frac(11, 84, p, rm),
291            b4_1: frac(5179, 57600, p, rm),
292            b4_3: frac(7571, 16695, p, rm),
293            b4_4: frac(393, 640, p, rm),
294            b4_5: frac(-92097, 339200, p, rm),
295            b4_6: frac(187, 2100, p, rm),
296            b4_7: frac(1, 40, p, rm),
297        }
298    }
299}
300
301fn axpy(
302    y: &ExactNum,
303    h: &ExactNum,
304    a: &ExactNum,
305    k: &ExactNum,
306    p: usize,
307    rm: RoundingMode,
308) -> ExactNum {
309    y.add(&h.mul(a, p, rm).mul(k, p, rm), p, rm)
310}
311
312fn dp45_step<F>(
313    tab: &Dp45,
314    f: &mut F,
315    t: &ExactNum,
316    y: &ExactNum,
317    h: &ExactNum,
318    p: usize,
319    rm: RoundingMode,
320    cc: &mut Consts,
321) -> Option<(ExactNum, ExactNum)>
322where
323    F: FnMut(&ExactNum, &ExactNum, usize, RoundingMode, &mut Consts) -> ExactNum,
324{
325    let k1 = f(t, y, p, rm, cc);
326    let t2 = t.add(&h.mul(&tab.c2, p, rm), p, rm);
327    let y2 = axpy(y, h, &tab.a21, &k1, p, rm);
328    let k2 = f(&t2, &y2, p, rm, cc);
329
330    let t3 = t.add(&h.mul(&tab.c3, p, rm), p, rm);
331    let y3 = axpy(&axpy(y, h, &tab.a31, &k1, p, rm), h, &tab.a32, &k2, p, rm);
332    let k3 = f(&t3, &y3, p, rm, cc);
333
334    let t4 = t.add(&h.mul(&tab.c4, p, rm), p, rm);
335    let y4s = axpy(
336        &axpy(&axpy(y, h, &tab.a41, &k1, p, rm), h, &tab.a42, &k2, p, rm),
337        h,
338        &tab.a43,
339        &k3,
340        p,
341        rm,
342    );
343    let k4 = f(&t4, &y4s, p, rm, cc);
344
345    let t5 = t.add(&h.mul(&tab.c5, p, rm), p, rm);
346    let y5s = axpy(
347        &axpy(
348            &axpy(&axpy(y, h, &tab.a51, &k1, p, rm), h, &tab.a52, &k2, p, rm),
349            h,
350            &tab.a53,
351            &k3,
352            p,
353            rm,
354        ),
355        h,
356        &tab.a54,
357        &k4,
358        p,
359        rm,
360    );
361    let k5 = f(&t5, &y5s, p, rm, cc);
362
363    let t6 = t.add(h, p, rm);
364    let y6 = axpy(
365        &axpy(
366            &axpy(
367                &axpy(&axpy(y, h, &tab.a61, &k1, p, rm), h, &tab.a62, &k2, p, rm),
368                h,
369                &tab.a63,
370                &k3,
371                p,
372                rm,
373            ),
374            h,
375            &tab.a64,
376            &k4,
377            p,
378            rm,
379        ),
380        h,
381        &tab.a65,
382        &k5,
383        p,
384        rm,
385    );
386    let k6 = f(&t6, &y6, p, rm, cc);
387
388    let y5 = axpy(
389        &axpy(
390            &axpy(
391                &axpy(&axpy(y, h, &tab.b5_1, &k1, p, rm), h, &tab.b5_3, &k3, p, rm),
392                h,
393                &tab.b5_4,
394                &k4,
395                p,
396                rm,
397            ),
398            h,
399            &tab.b5_5,
400            &k5,
401            p,
402            rm,
403        ),
404        h,
405        &tab.b5_6,
406        &k6,
407        p,
408        rm,
409    );
410    let k7 = f(&t6, &y5, p, rm, cc);
411
412    let y4 = axpy(
413        &axpy(
414            &axpy(
415                &axpy(
416                    &axpy(&axpy(y, h, &tab.b4_1, &k1, p, rm), h, &tab.b4_3, &k3, p, rm),
417                    h,
418                    &tab.b4_4,
419                    &k4,
420                    p,
421                    rm,
422                ),
423                h,
424                &tab.b4_5,
425                &k5,
426                p,
427                rm,
428            ),
429            h,
430            &tab.b4_6,
431            &k6,
432            p,
433            rm,
434        ),
435        h,
436        &tab.b4_7,
437        &k7,
438        p,
439        rm,
440    );
441
442    if !finite(&y5) || !finite(&y4) {
443        return None;
444    }
445    Some((y5, y4))
446}
447
448/// Dormand–Prince RK5(4) with absolute / relative step control.
449///
450/// Accepts a step when `|y₅−y₄| ≤ atol + rtol max(|y|,|y₅|)`. Step size is
451/// updated by `(atol_scale / err)^{1/5}` with safety `9/10`, growth cap 5,
452/// shrink cap `1/5`. `None` if `t1 ≤ t0`, a value is non-finite, the step
453/// falls below [`ode_min_step`], or [`ODE_MAX_STEPS`] accepted steps are
454/// exhausted before `t1`.
455pub fn rk45_adaptive<F>(
456    mut f: F,
457    t0: &ExactNum,
458    y0: &ExactNum,
459    t1: &ExactNum,
460    atol: &ExactNum,
461    rtol: &ExactNum,
462    p: usize,
463    rm: RoundingMode,
464    cc: &mut Consts,
465) -> Option<(ExactNumArray, ExactNumArray)>
466where
467    F: FnMut(&ExactNum, &ExactNum, usize, RoundingMode, &mut Consts) -> ExactNum,
468{
469    if !finite(t0) || !finite(y0) || !finite(t1) || !finite(atol) || !finite(rtol) {
470        return None;
471    }
472    if t0.cmp(t1) != Some(-1) || atol.is_negative() || rtol.is_negative() {
473        return None;
474    }
475    let wrk = work_p(p);
476    let tab = Dp45::new(wrk, RoundingMode::None);
477    let hmin = ode_min_step(wrk, RoundingMode::None);
478    let fac = frac(ODE_FAC_NUM, ODE_FAC_DEN, wrk, RoundingMode::None);
479    let grow = ExactNum::from_i64(ODE_H_GROW, wrk);
480    let shrink = frac(1, ODE_H_SHRINK_DEN, wrk, RoundingMode::None);
481    let mut t = t0.clone();
482    let mut y = y0.clone();
483    let _ = t.set_precision(wrk, RoundingMode::None);
484    let _ = y.set_precision(wrk, RoundingMode::None);
485    let mut h = t1
486        .sub(t0, wrk, RoundingMode::None)
487        .ldexp(-4, wrk, RoundingMode::None);
488    let mut ts = Vec::new();
489    let mut ys = Vec::new();
490    if !push_pair(&mut ts, &mut ys, t.clone(), y.clone(), ODE_MAX_STEPS + 1) {
491        return None;
492    }
493    let mut accepted = 0usize;
494    let mut attempts = 0usize;
495    while t.cmp(t1) == Some(-1) {
496        if accepted >= ODE_MAX_STEPS || attempts >= ODE_MAX_STEPS.saturating_mul(4) {
497            return None;
498        }
499        attempts += 1;
500        let remain = t1.sub(&t, wrk, RoundingMode::None);
501        if h.cmp(&remain) == Some(1) {
502            h = remain;
503        }
504        if h.cmp(&hmin) == Some(-1) || h.is_zero() || h.is_negative() {
505            return None;
506        }
507        let (y5, y4) = dp45_step(&tab, &mut f, &t, &y, &h, wrk, RoundingMode::None, cc)?;
508        let err = y5.sub(&y4, wrk, RoundingMode::None).abs();
509        let ymax = if y.abs().cmp(&y5.abs()) == Some(1) { y.abs() } else { y5.abs() };
510        let scale = atol.add(
511            &rtol.mul(&ymax, wrk, RoundingMode::None),
512            wrk,
513            RoundingMode::None,
514        );
515        let ok = scale.is_zero()
516            || err.is_zero()
517            || err.cmp(&scale) == Some(-1)
518            || err.cmp(&scale) == Some(0);
519        if ok {
520            t = t.add(&h, wrk, RoundingMode::None);
521            y = y5;
522            if t.cmp(t1) == Some(1) {
523                t = t1.clone();
524            }
525            if !push_pair(&mut ts, &mut ys, t.clone(), y.clone(), ODE_MAX_STEPS + 1) {
526                return None;
527            }
528            accepted += 1;
529        }
530        let ratio = if err.is_zero() {
531            grow.clone()
532        } else {
533            let q = scale
534                .div(&err, wrk, RoundingMode::None)
535                .nth_root(5, wrk, RoundingMode::None);
536            let raw = fac.mul(&q, wrk, RoundingMode::None);
537            if raw.cmp(&grow) == Some(1) {
538                grow.clone()
539            } else if raw.cmp(&shrink) == Some(-1) {
540                shrink.clone()
541            } else {
542                raw
543            }
544        };
545        h = h.mul(&ratio, wrk, RoundingMode::None);
546        if !ok && !finite(&h) {
547            return None;
548        }
549    }
550    Some((to_row(p, rm, &ts)?, to_row(p, rm, &ys)?))
551}
552
553#[cfg(test)]
554mod tests {
555    use super::*;
556    use crate::Consts;
557
558    fn gold_p() -> (usize, RoundingMode) {
559        (256, RoundingMode::ToEven)
560    }
561
562    /// Global RK4 error on `y'=-y` with `h=10^{-3}` is `O(h^4) ≈ 10^{-12}`.
563    const RK4_EXP_ERR_DIGITS: isize = 12;
564    /// Adaptive gold the step cap can actually meet (plan `1e-50` needs ~10¹² steps).
565    const RK45_ATOL_DIGITS: isize = 12;
566    const EULER_OH_STEPS: usize = 1000;
567
568    #[test]
569    fn ode_rk4_rk45_euler() {
570        let (p, rm) = gold_p();
571        let mut cc = Consts::new().expect("consts");
572        let zero = ExactNum::new(p);
573        let one = ExactNum::from_u8(1, p);
574        let ten = ExactNum::from_u8(10, p);
575
576        let (_t, y) = rk4(
577            |_t, y, _p, _rm, _cc| y.neg(),
578            &zero,
579            &one,
580            &one,
581            EULER_OH_STEPS,
582            p,
583            rm,
584            &mut cc,
585        )
586        .expect("rk4");
587        let yend = y.get(y.len() - 1).expect("yend").clone();
588        let einv = one.neg().exp(p, rm, &mut cc);
589        let err = yend.sub(&einv, p, rm).abs();
590        let tol12 = one.div(&ten.powsi(RK4_EXP_ERR_DIGITS, p, rm), p, rm);
591        assert!(err.is_zero() || err.cmp(&tol12) == Some(-1));
592
593        let atol = one.div(&ten.powsi(RK45_ATOL_DIGITS, p, rm), p, rm);
594        let rtol = ExactNum::new(p);
595        let (_t45, y45) = rk45_adaptive(
596            |_t, y, _p, _rm, _cc| y.neg(),
597            &zero,
598            &one,
599            &one,
600            &atol,
601            &rtol,
602            p,
603            rm,
604            &mut cc,
605        )
606        .expect("rk45");
607        let y45e = y45.get(y45.len() - 1).expect("y45e").clone();
608        let err45 = y45e.sub(&einv, p, rm).abs();
609        assert!(err45.is_zero() || err45.cmp(&atol) == Some(-1) || err45.cmp(&atol) == Some(0));
610
611        let (_te, ye) = euler(
612            |_t, y, _p, _rm| y.clone(),
613            &zero,
614            &one,
615            &one,
616            EULER_OH_STEPS,
617            p,
618            rm,
619        )
620        .expect("euler 1000");
621        let (_te2, ye2) = euler(
622            |_t, y, _p, _rm| y.clone(),
623            &zero,
624            &one,
625            &one,
626            EULER_OH_STEPS * 2,
627            p,
628            rm,
629        )
630        .expect("euler 2000");
631        assert_eq!(ye.len(), EULER_OH_STEPS + 1);
632        assert_eq!(ye2.len(), EULER_OH_STEPS * 2 + 1);
633        let two = ExactNum::from_u8(2, p);
634        let ee = one.exp(p, rm, &mut cc);
635        let e1 = ye.get(ye.len() - 1).expect("e1").sub(&ee, p, rm).abs();
636        let e2 = ye2.get(ye2.len() - 1).expect("e2").sub(&ee, p, rm).abs();
637        let two_h_bound = two.div(&ExactNum::from_u32(EULER_OH_STEPS as u32, p), p, rm);
638        assert!(e1.cmp(&two_h_bound) == Some(-1));
639        assert!(e2.cmp(&e1) == Some(-1));
640
641        assert!(rk4(
642            |_t, y, _p, _rm, _cc| y.neg(),
643            &one,
644            &one,
645            &zero,
646            4,
647            p,
648            rm,
649            &mut cc
650        )
651        .is_none());
652    }
653}