1use crate::defs::WORD_BIT_SIZE;
4use crate::Consts;
5use crate::ExactNum;
6use crate::ExactNumArray;
7use crate::RoundingMode;
8use alloc::vec::Vec;
9
10pub const ODE_MAX_STEPS: usize = 65536;
12
13pub const ODE_MIN_STEP: i32 = -256;
15
16const ODE_FAC_NUM: i64 = 9;
18const ODE_FAC_DEN: i64 = 10;
20const ODE_H_GROW: i64 = 5;
22const 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
37pub 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
69pub 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
169pub 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
448pub 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 const RK4_EXP_ERR_DIGITS: isize = 12;
564 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}