1use anyhow::anyhow;
2use tch::IndexOp;
3use tch::Tensor;
4use std::sync::Arc;
5
6pub mod optimizers;
7
8pub enum Solver {
9 Euler { step: f64 },
10 RK4 { step: f64 },
11 ImplicitEuler { step: f64, optimizer: Arc<dyn optimizers::Optimizer> },
12 GLRK4 { step: f64, optimizer: Arc<dyn optimizers::Optimizer> },
13 RKF45 { rtol: f64, atol: f64, min_step: f64, safety_factor: f64 },
14 ROW1 { step: f64 }
15}
16
17impl Solver {
18 pub fn solve(
19 &self,
20 f: tch::CModule,
21 x_span: Tensor,
22 y0: Tensor
23 ) -> anyhow::Result<(Tensor, Tensor)> {
24 if x_span.size() != [2] {
25 return Err(anyhow!("x_span must be of shape [2] but it has shape {:?}", x_span.size().as_slice()));
26 }
27 if y0.size().len() != 1 {
28 return Err(anyhow!("y0 must be a one-dimensional tensor but it has {} dimensions", y0.size().len()));
29 }
30 if x_span.device() != y0.device() {
31 return Err(anyhow!("x_span and y0 must reside on the same device. Device of x_span is {:?}. Device of y0 is {:?}", x_span.device(), y0.device()));
32 }
33 if x_span.kind() != tch::Kind::Double && x_span.kind() != tch::Kind::Float && x_span.kind() != tch::Kind::BFloat16 && x_span.kind() != tch::Kind::Half {
34 return Err(anyhow!("x_span is of unsupported kind {:?}", x_span.kind()));
35 }
36 if y0.kind() != tch::Kind::Double && y0.kind() != tch::Kind::Float && y0.kind() != tch::Kind::BFloat16 && y0.kind() != tch::Kind::Half {
37 return Err(anyhow!("y0 is of unsupported kind {:?}", y0.kind()));
38 }
39 if x_span.kind() != y0.kind() {
40 return Err(anyhow!("x_span and y0 must be of the same kind. Kind of x_span is {:?}. Kind of y0 is {:?}", x_span.kind(), y0.kind()));
41 }
42
43 match self {
44 Self::Euler { step } => solve_euler(f, x_span, y0, *step),
45 Self::RK4 { step } => solve_rk4(f, x_span, y0, *step),
46 Self::ImplicitEuler { step, optimizer } => solve_implicit_euler(f, x_span, y0, *step, optimizer.as_ref()),
47 Self::GLRK4 { step, optimizer } => solve_glrk4(f, x_span, y0, *step, optimizer.as_ref()),
48 Self::RKF45 { rtol, atol, min_step, safety_factor } => solve_rkf45(f, x_span, y0, *rtol, *atol, *min_step, *safety_factor),
49 Self::ROW1 { step } => solve_row1(f, x_span, y0, *step)
50 }
51 }
52}
53
54fn solve_euler(
56 f: tch::CModule,
57 x_span: Tensor,
58 y0: Tensor,
59 step: f64,
60) -> anyhow::Result<(Tensor, Tensor)> {
61 let x_start = x_span.i(0);
62 let x_end = x_span.i(1);
63
64 let mut x = x_start.unsqueeze(0);
65 let mut y = y0.unsqueeze(0);
66
67 let mut all_x = vec![x.copy()];
68 let mut all_y = vec![y.copy()];
69
70 let mut current_step = step;
71 while x.lt_tensor(&x_end) == Tensor::from_slice(&[true]) {
72 let remaining = &x_end - &x.squeeze();
73 if remaining.double_value(&[]) < current_step {
74 current_step = remaining.double_value(&[]);
75 }
76
77 let dy = f.forward_ts(&[x.squeeze().copy(), y.squeeze().copy()])?;
78 let dy_rank = dy.size().len();
79 if dy_rank != 1 {
80 anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", dy_rank);
81 }
82
83 y = &y + current_step * &dy;
84 x = &x + current_step;
85
86 all_x.push(x.copy());
87 all_y.push(y.copy());
88 }
89
90 Ok((Tensor::cat(&all_x, 0), Tensor::cat(&all_y, 0)))
91}
92
93fn solve_rk4(
95 f: tch::CModule,
96 x_span: Tensor,
97 y0: Tensor,
98 step: f64,
99) -> anyhow::Result<(Tensor, Tensor)> {
100 let x_start = x_span.i(0);
101 let x_end = x_span.i(1);
102
103 let mut x = x_start.unsqueeze(0);
104 let mut y = y0.unsqueeze(0);
105
106 let mut all_x = vec![x.copy()];
107 let mut all_y = vec![y.copy()];
108
109 let mut current_step = step;
110 while x.lt_tensor(&x_end) == Tensor::from_slice(&[true]) {
111 let remaining = &x_end - &x.squeeze();
112 if remaining.double_value(&[]) < current_step {
113 current_step = remaining.double_value(&[]);
114 }
115
116 let k1 = f.forward_ts(&[x.squeeze().copy(), y.squeeze().copy()])?;
117 let k1_rank = k1.size().len();
118 if k1_rank != 1 {
119 anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", k1_rank);
120 }
121
122 let x_half: Tensor = &x + 0.5 * current_step;
123 let y_half: Tensor = &y + 0.5 * current_step * &k1;
124 let k2 = f.forward_ts(&[x_half.squeeze(), y_half.squeeze()])?;
125 let k2_rank = k2.size().len();
126 if k2_rank != 1 {
127 anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", k2_rank);
128 }
129
130 let x_half_again: Tensor = &x + 0.5 * current_step;
131 let y_half_again: Tensor = &y + 0.5 * current_step * &k2;
132 let k3 = f.forward_ts(&[x_half_again.squeeze(), y_half_again.squeeze()])?;
133 let k3_rank = k3.size().len();
134 if k3_rank != 1 {
135 anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", k3_rank);
136 }
137
138 let x_full = &x + current_step;
139 let y_full = &y + current_step * &k3;
140 let k4 = f.forward_ts(&[x_full.squeeze(), y_full.squeeze()])?;
141 let k4_rank = k4.size().len();
142 if k4_rank != 1 {
143 anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", k4_rank);
144 }
145
146 let step_div_6 = current_step / 6.0;
147 let y_next = &y + step_div_6 * (&k1 + 2.0 * &k2 + 2.0 * &k3 + &k4);
148
149 x = &x + current_step;
150 y = y_next;
151
152 all_x.push(x.copy());
153 all_y.push(y.copy());
154 }
155
156 Ok((Tensor::cat(&all_x, 0), Tensor::cat(&all_y, 0)))
157}
158
159fn solve_implicit_euler(
161 f: tch::CModule,
162 x_span: Tensor,
163 y0: Tensor,
164 step: f64,
165 optimizer: &dyn optimizers::Optimizer,
166) -> anyhow::Result<(Tensor, Tensor)> {
167 let x_start = x_span.i(0);
168 let x_end = x_span.i(1);
169
170 let mut x = x_start.unsqueeze(0);
171 let mut y = y0.unsqueeze(0);
172
173 let mut all_x = vec![x.copy()];
174 let mut all_y = vec![y.copy()];
175
176 let mut current_step = step;
177 while x.lt_tensor(&x_end) == Tensor::from_slice(&[true]) {
178 let remaining = &x_end - &x.squeeze();
179 if remaining.double_value(&[]) < current_step {
180 current_step = remaining.double_value(&[]);
181 }
182
183 let x_next = &x + current_step;
184 let y_prev = y.copy();
185
186 let y_next = optimizer.optimize(
187 &|y_next: &Tensor| {
188 let f_next = f
189 .forward_ts(&[x_next.squeeze().copy(), y_next.squeeze().copy()])
190 .unwrap();
191 let y_pred = &y_prev.squeeze() + current_step * &f_next;
192 (y_next - &y_pred).pow_tensor_scalar(2).sum(y_next.kind())
193 },
194 &(&y_prev.detach().squeeze()
195 + current_step * f.forward_ts(&[&x.squeeze(), &y_prev.squeeze()])?),
196 ).map_err( |err| {
197 anyhow!(format!("Optimizer failed with: {}", err))
198 })?;
199
200 y = y_next.unsqueeze(0);
201 x = x_next.copy();
202
203 all_x.push(x.copy());
204 all_y.push(y.copy());
205 }
206
207 Ok((Tensor::cat(&all_x, 0), Tensor::cat(&all_y, 0)))
208}
209
210fn solve_glrk4(
212 f: tch::CModule,
213 x_span: Tensor,
214 y0: Tensor,
215 step: f64,
216 optimizer: &dyn optimizers::Optimizer,
217) -> anyhow::Result<(Tensor, Tensor)> {
218 let x_start = x_span.i(0);
219 let x_end = x_span.i(1);
220
221 let mut x = x_start.unsqueeze(0);
222 let mut y = y0.unsqueeze(0);
223
224 let mut all_x = vec![x.copy()];
225 let mut all_y = vec![y.copy()];
226
227 let mut current_step = step;
228 while x.lt_tensor(&x_end) == Tensor::from_slice(&[true]) {
229 let remaining = &x_end - &x.squeeze();
230 if remaining.double_value(&[]) < current_step {
231 current_step = remaining.double_value(&[]);
232 }
233
234 let k = f.forward_ts(&[x.squeeze().copy(), y.squeeze().copy()])?;
235 let k_rank = k.size().len();
236 if k_rank != 1 {
237 anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", k_rank);
238 }
239
240 const C1: f64 = 0.2113248654f64;
241 const C2: f64 = 0.7886751346f64;
242 const A11: f64 = 0.25;
243 const A12: f64 = -0.03867513459f64;
244 const A21: f64 = 0.5386751346f64;
245 const A22: f64 = 0.25;
246
247 let first_k1k2_guess = Tensor::cat(
248 &[
249 f.forward_ts(&[
250 &x.squeeze() + C1 * current_step,
251 &y.squeeze() + C1 * current_step * &k,
252 ])?,
253 f.forward_ts(&[
254 &x.squeeze() + C2 * current_step,
255 &y.squeeze() + C2 * current_step * &k,
256 ])?,
257 ],
258 0,
259 );
260 let k1k2 = optimizer.optimize(
261 &|k1k2_guess| {
262 let diff1 = k1k2_guess.i(0..=1)
263 - f.forward_ts(&[
264 &x.squeeze() + C1 * current_step,
265 &y.squeeze()
266 + (A11 * k1k2_guess.i(0..=1) + A12 * k1k2_guess.i(2..=3))
267 * current_step,
268 ])
269 .unwrap();
270 let diff2 = k1k2_guess.i(2..=3)
271 - f.forward_ts(&[
272 &x.squeeze() + C2 * current_step,
273 &y.squeeze()
274 + (A21 * k1k2_guess.i(0..=1) + A22 * k1k2_guess.i(2..=3))
275 * current_step,
276 ])
277 .unwrap();
278
279 diff1.dot(&diff1) + diff2.dot(&diff2)
280 },
281 &first_k1k2_guess,
282 ).map_err( |err| {
283 anyhow!(format!("Optimizer failed with: {}", err))
284 })?;
285 assert!(k1k2.size().len() == 1);
286 assert!(k1k2.size()[0] == 4);
287
288 x = &x + current_step;
289 y = &y + current_step * (0.5 * k1k2.i(0..=1) + 0.5 * k1k2.i(2..=3));
290
291 all_x.push(x.copy());
292 all_y.push(y.copy());
293 }
294
295 Ok((Tensor::cat(&all_x, 0), Tensor::cat(&all_y, 0)))
296}
297
298fn solve_rkf45(
300 f: tch::CModule,
301 x_span: Tensor,
302 y0: Tensor,
303 rtol: f64,
304 atol: f64,
305 min_step: f64,
306 safety_factor: f64,
307) -> anyhow::Result<(Tensor, Tensor)> {
308 let x_start = x_span.i(0);
309 let x_end = x_span.i(1);
310
311 let mut x = x_start.unsqueeze(0);
312 let mut y = y0.unsqueeze(0);
313
314 let mut all_x = vec![x.copy()];
315 let mut all_y = vec![y.copy()];
316
317 let mut step = (&x_end - &x_start) * 0.1;
318 let safety_factor_tensor = Tensor::from(safety_factor);
319
320 while x.lt_tensor(&x_end) == Tensor::from_slice(&[true]) {
321 let remaining = &x_end - &x.squeeze();
322 if remaining.lt_tensor(&step) == Tensor::from(true) {
323 step = remaining.copy();
324 }
325
326 let k1 = f.forward_ts(&[x.squeeze().copy(), y.squeeze().copy()])?;
327 let k1_rank = k1.size().len();
328 if k1_rank != 1 {
329 anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", k1_rank);
330 }
331
332 let k2 = {
333 let x_step: Tensor = &x + 0.25 * &step;
334 let y_step: Tensor = &y + 0.25 * &step * &k1;
335 let k2_unchecked = f.forward_ts(&[x_step.squeeze(), y_step.squeeze()])?;
336 let k2_rank = k2_unchecked.size().len();
337 if k2_rank != 1 {
338 anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", k2_rank);
339 }
340
341 k2_unchecked
342 };
343
344 let k3 = {
345 let x_step: Tensor = &x + 0.375 * &step;
346 let y_step: Tensor = &y + (0.09375 * &step * &k1) + (0.28125 * &step * &k2);
347 let k3_unchecked = f.forward_ts(&[x_step.squeeze(), y_step.squeeze()])?;
348 let k3_rank = k3_unchecked.size().len();
349 if k3_rank != 1 {
350 anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", k3_rank);
351 }
352
353 k3_unchecked
354 };
355
356 let k4 = {
357 let x_step: Tensor = &x + (12.0 / 13.0) * &step;
358 let y_step: Tensor = &y
359 + (1932.0 / 2197.0 * &step * &k1)
360 + (-7200.0 / 2197.0 * &step * &k2)
361 + (7296.0 / 2197.0 * &step * &k3);
362 let k4_unchecked = f.forward_ts(&[x_step.squeeze(), y_step.squeeze()])?;
363 let k4_rank = k4_unchecked.size().len();
364 if k4_rank != 1 {
365 anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", k4_rank);
366 }
367
368 k4_unchecked
369 };
370
371 let k5 = {
372 let x_step: Tensor = &x + &step;
373 let y_step: Tensor = &y
374 + (439.0 / 216.0 * &step * &k1)
375 + (-8.0 * &step * &k2)
376 + (3680.0 / 513.0 * &step * &k3)
377 + (-845.0 / 4104.0 * &step * &k4);
378 let k5_unchecked = f.forward_ts(&[x_step.squeeze(), y_step.squeeze()])?;
379 let k5_rank = k5_unchecked.size().len();
380 if k5_rank != 1 {
381 anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", k5_rank);
382 }
383
384 k5_unchecked
385 };
386
387 let k6 = {
388 let x_step: Tensor = &x + 0.5 * &step;
389 let y_step: Tensor = &y
390 + (-8.0 / 27.0 * &step * &k1)
391 + (2.0 * &step * &k2)
392 + (-3544.0 / 2565.0 * &step * &k3)
393 + (1859.0 / 4104.0 * &step * &k4)
394 + (-11.0 / 40.0 * &step * &k5);
395 let k6_unchecked = f.forward_ts(&[x_step.squeeze(), y_step.squeeze()])?;
396 let k6_rank = k6_unchecked.size().len();
397 if k6_rank != 1 {
398 anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", k6_rank);
399 }
400
401 k6_unchecked
402 };
403
404 let next_y4: Tensor = &y
405 + &step
406 * ((25.0 / 216.0 * &k1)
407 + (1408.0 / 2565.0 * &k3)
408 + (2197.0 / 4104.0 * &k4)
409 + (-1.0 / 5.0 * &k5));
410 let next_y5: Tensor = &y
411 + &step
412 * ((16.0 / 135.0 * &k1)
413 + (6656.0 / 12825.0 * &k3)
414 + (28561.0 / 56430.0 * &k4)
415 + (-9.0 / 50.0 * &k5)
416 + (2.0 / 55.0 * &k6));
417
418 let d = (&next_y4 - &next_y5).abs();
419 let e = next_y5.abs() * rtol + atol;
420
421 let alpha_tensor = (e / d).sqrt().min();
422 let condition = &safety_factor_tensor * &alpha_tensor;
423
424 let condition_met = condition.lt(1.0);
425 let condition_met_bool: bool = condition_met == Tensor::from(true);
426
427 if condition_met_bool {
428 step = &step * &condition;
429 if step.double_value(&[]) < min_step {
430 return Err(anyhow!("Required step is smaller than minimal step"));
431 }
432 } else {
433 y = next_y4;
434 x = &x + &step;
435 all_x.push(x.copy());
436 all_y.push(y.copy());
437
438 let new_step = &step * &condition;
439 let max_step = &step * 5.0;
440 step = new_step.fmin(&max_step);
441 }
442 }
443
444 Ok((Tensor::cat(&all_x, 0), Tensor::cat(&all_y, 0)))
445}
446
447fn solve_row1(
449 f: tch::CModule,
450 x_span: Tensor,
451 y0: Tensor,
452 step: f64,
453) -> anyhow::Result<(Tensor, Tensor)> {
454 let x_start = x_span.i(0);
455 let x_end = x_span.i(1);
456
457 let mut x = x_start.unsqueeze(0);
458 let mut y = y0.unsqueeze(0);
459
460 let mut all_x = vec![x.copy()];
461 let mut all_y = vec![y.copy()];
462
463 while x.lt_tensor(&x_end) == Tensor::from_slice(&[true]) {
464 let remaining = &x_end - &x.squeeze();
465 let mut current_step = step;
466 if remaining.double_value(&[]) < step {
467 current_step = remaining.double_value(&[]);
468 }
469
470 let x_prev = x.copy();
471 let y_prev = y.copy().squeeze();
472
473 let jacobian = compute_jacobian(
474 |y| {
475 f.forward_ts(&[x_prev.squeeze().copy(), y.copy()])
476 .unwrap()
477 .squeeze()
478 },
479 &y_prev,
480 );
481 let f_current = f
482 .forward_ts(&[x_prev.squeeze().copy(), y_prev.copy()])?;
483 let f_current_rank = f_current.size().len();
484 if f_current_rank != 1 {
485 anyhow::bail!("Derivative CModule returned tensor of bad rank {}.", f_current_rank);
486 }
487
488 let n = jacobian.size()[0];
489 let eye = Tensor::eye(n, (tch::Kind::Float, jacobian.device()));
490 let step_j = current_step * &jacobian;
491 let inv_matrix = (eye - step_j).inverse();
492
493 let delta_y = inv_matrix.matmul(&f_current);
494 let y_next = y_prev + current_step * delta_y;
495
496 x = &x_prev + current_step;
497 y = y_next.unsqueeze(0);
498
499 all_x.push(x.copy());
500 all_y.push(y.copy());
501 }
502
503 Ok((Tensor::cat(&all_x, 0), Tensor::cat(&all_y, 0)))
504}
505
506fn compute_jacobian<F>(f: F, x: &Tensor) -> Tensor
508where
509 F: Fn(&Tensor) -> Tensor,
510{
511 assert_eq!(x.dim(), 1, "x must be 1-dimensional");
512 let mut x_with_grad = x.detach().copy().set_requires_grad(true);
513 let y = f(&x_with_grad);
514 assert_eq!(y.dim(), 1, "y must be 1-dimensional");
515
516 let y_size = y.size()[0];
517 let mut grads = Vec::new();
518
519 for i in 0..y_size {
520 let yi = y.i(i);
521 let grad = Tensor::run_backward(&[yi], &[&x_with_grad], true, false)[0].copy();
524 grads.push(grad.unsqueeze(0));
525 x_with_grad.zero_grad();
526 }
527
528 Tensor::cat(&grads, 0)
529}