1use anyhow::anyhow;
2use std::fmt;
3use std::sync::Arc;
4use tch::IndexOp;
5use tch::Tensor;
6
7pub mod optimizers;
8
9#[cfg(test)]
10mod tests;
11
12pub enum Solver {
13 Euler {
14 step: f64,
15 },
16 RK4 {
17 step: f64,
18 },
19 ImplicitEuler {
20 step: f64,
21 optimizer: Arc<dyn optimizers::Optimizer>,
22 },
23 GLRK4 {
24 step: f64,
25 optimizer: Arc<dyn optimizers::Optimizer>,
26 },
27 RKF45 {
28 rtol: f64,
29 atol: f64,
30 min_step: f64,
31 safety_factor: f64,
32 },
33 ROW1 {
34 step: f64,
35 },
36}
37
38impl Solver {
39 pub fn solve(
40 &self,
41 f: tch::CModule,
42 x_span: (f64, f64),
43 y0: Tensor,
44 ) -> anyhow::Result<(Tensor, Tensor)> {
45 let kind = y0.kind();
46 let device = y0.device();
47
48 let y0_size = y0.size();
50 if y0_size.len() != 1 {
51 return Err(anyhow!(
52 "y0 must be a one-dimensional tensor but it has {} dimensions",
53 y0_size.len()
54 ));
55 }
56 if kind != tch::Kind::Double
57 && kind != tch::Kind::Float
58 && kind != tch::Kind::BFloat16
59 && kind != tch::Kind::Half
60 {
61 return Err(anyhow!("y0 is of unsupported kind {:?}", y0.kind()));
62 }
63
64 let dy = f.forward_ts(&[
66 Tensor::from(x_span.0).to_kind(kind).to_device(device),
67 y0.copy(),
68 ])?;
69 let dy_size = dy.size();
70 if dy_size.len() != 1 {
71 return Err(anyhow!(
72 "Function `f` returns tensor of rank {}, expected one-dimensional tensor",
73 dy_size.len()
74 ));
75 }
76 if dy_size[0] != y0_size[0] {
77 return Err(anyhow!(
78 "Function `f` returns vector of length {}, expected vector of length {} (same as y0)",
79 dy_size[0], y0_size[0]
80 ));
81 }
82 let dy_device = dy.device();
83 if dy_device != device {
84 return Err(anyhow!(
85 "Function `f` returns tensor on device {:?}, expected tensor to be on device {:?} (same as y0)",
86 dy_device, device
87 ));
88 }
89 let dy_kind = dy.kind();
90 if dy_kind != kind {
91 return Err(anyhow!(
92 "Function `f` returns tensor of kind {:?}, expected tensor to be of kind {:?} (same as y0)",
93 dy_kind, kind
94 ));
95 }
96
97 match self {
98 Self::Euler { step } => solve_euler(f, x_span, y0, *step),
99 Self::RK4 { step } => solve_rk4(f, x_span, y0, *step),
100 Self::ImplicitEuler { step, optimizer } => {
101 solve_implicit_euler(f, x_span, y0, *step, optimizer.as_ref())
102 }
103 Self::GLRK4 { step, optimizer } => {
104 solve_glrk4(f, x_span, y0, *step, optimizer.as_ref())
105 }
106 Self::RKF45 {
107 rtol,
108 atol,
109 min_step,
110 safety_factor,
111 } => solve_rkf45(f, x_span, y0, *rtol, *atol, *min_step, *safety_factor),
112 Self::ROW1 { step } => solve_row1(f, x_span, y0, *step),
113 }
114 }
115}
116
117impl fmt::Display for Solver {
118 fn fmt(&self, f: &mut fmt::Formatter) -> fmt::Result {
119 match self {
120 Solver::Euler { step } => write!(f, "Euler(step={})", step),
121 Solver::RK4 { step } => write!(f, "RK4(step={})", step),
122 Solver::ImplicitEuler { step, optimizer } => {
123 write!(f, "ImplicitEuler(step={}, optimizer={})", step, optimizer)
124 }
125 Solver::GLRK4 { step, optimizer } => {
126 write!(f, "GLRK4(step={}, optimizer={})", step, optimizer)
127 }
128 Solver::RKF45 {
129 rtol,
130 atol,
131 min_step,
132 safety_factor,
133 } => write!(
134 f,
135 "RKF45(rtol={}, atol={}, min_step={}, safety_factor={})",
136 rtol, atol, min_step, safety_factor
137 ),
138 Solver::ROW1 { step } => write!(f, "ROW1(step={})", step),
139 }
140 }
141}
142
143fn solve_euler(
145 f: tch::CModule,
146 x_span: (f64, f64),
147 y0: Tensor,
148 step: f64,
149) -> anyhow::Result<(Tensor, Tensor)> {
150 let device = y0.device();
151 let kind = y0.kind();
152
153 let x_start = x_span.0;
154 let x_end = x_span.1;
155
156 let mut x = x_start;
157 let mut y = y0.copy();
158
159 let mut all_x = vec![x];
160 let mut all_y = vec![y.copy()];
161
162 let mut current_step = step;
163 while x < x_end {
164 let remaining = x_end - x;
165 if remaining < current_step {
166 current_step = remaining;
167 }
168
169 let dy = f.forward_ts(&[
170 Tensor::from(x).to_kind(kind).to_device(device),
171 y.copy(),
172 ])?;
173 let dy_size = dy.size();
174 let dy_rank = dy_size.len();
175 if dy_rank != 1 {
176 anyhow::bail!(
177 "Derivative CModule returned tensor of bad rank {}.",
178 dy_rank
179 );
180 }
181 if dy_size[0] != y0.size()[0] {
182 anyhow::bail!(
183 "Derivative CModule returned vector of bad length {}.",
184 dy_size[0]
185 );
186 }
187
188 y = &y + current_step * &dy;
189 x = &x + current_step;
190
191 all_x.push(x);
192 all_y.push(y.copy());
193 }
194
195 Ok((
196 Tensor::from_slice(&all_x).to_kind(kind).to_device(device),
197 Tensor::stack(&all_y, 0),
198 ))
199}
200
201fn solve_rk4(
203 f: tch::CModule,
204 x_span: (f64, f64),
205 y0: Tensor,
206 step: f64,
207) -> anyhow::Result<(Tensor, Tensor)> {
208 let device = y0.device();
209 let kind = y0.kind();
210
211 let x_start = x_span.0;
212 let x_end = x_span.1;
213
214 let mut x = x_start;
215 let mut y = y0.copy();
216
217 let mut all_x = vec![x];
218 let mut all_y = vec![y.copy()];
219
220 let mut current_step = step;
221 while x < x_end {
222 let remaining = x_end - x;
223 if remaining < current_step {
224 current_step = remaining;
225 }
226
227 let k1 = f.forward_ts(&[
228 Tensor::from(x).to_kind(kind).to_device(device),
229 y.copy(),
230 ])?;
231 let k1_size = k1.size();
232 let k1_rank = k1_size.len();
233 if k1_rank != 1 {
234 anyhow::bail!(
235 "Derivative CModule returned tensor of bad rank {}.",
236 k1_rank
237 );
238 }
239 if k1_size[0] != y0.size()[0] {
240 anyhow::bail!(
241 "Derivative CModule returned vector of bad length {}.",
242 k1_size[0]
243 );
244 }
245
246 let x_half = x + 0.5 * current_step;
247 let y_half: Tensor = &y + 0.5 * current_step * &k1;
248 let k2 = f.forward_ts(&[
249 Tensor::from(x_half).to_kind(kind).to_device(device),
250 y_half,
251 ])?;
252 let k2_size = k2.size();
253 let k2_rank = k2_size.len();
254 if k2_rank != 1 {
255 anyhow::bail!(
256 "Derivative CModule returned tensor of bad rank {}.",
257 k2_rank
258 );
259 }
260 if k2_size[0] != y0.size()[0] {
261 anyhow::bail!(
262 "Derivative CModule returned vector of bad length {}.",
263 k2_size[0]
264 );
265 }
266
267 let x_half_again = x + 0.5 * current_step;
268 let y_half_again: Tensor = &y + 0.5 * current_step * &k2;
269 let k3 = f.forward_ts(&[
270 Tensor::from(x_half_again).to_kind(kind).to_device(device),
271 y_half_again,
272 ])?;
273 let k3_size = k3.size();
274 let k3_rank = k3_size.len();
275 if k3_rank != 1 {
276 anyhow::bail!(
277 "Derivative CModule returned tensor of bad rank {}.",
278 k3_rank
279 );
280 }
281 if k3_size[0] != y0.size()[0] {
282 anyhow::bail!(
283 "Derivative CModule returned vector of bad length {}.",
284 k3_size[0]
285 );
286 }
287
288 let x_full = x + current_step;
289 let y_full = &y + current_step * &k3;
290 let k4 = f.forward_ts(&[
291 Tensor::from(x_full).to_kind(kind).to_device(device),
292 y_full,
293 ])?;
294 let k4_size = k4.size();
295 let k4_rank = k4_size.len();
296 if k4_rank != 1 {
297 anyhow::bail!(
298 "Derivative CModule returned tensor of bad rank {}.",
299 k4_rank
300 );
301 }
302 if k4_size[0] != y0.size()[0] {
303 anyhow::bail!(
304 "Derivative CModule returned vector of bad length {}.",
305 k4_size[0]
306 );
307 }
308
309 let step_div_6 = current_step / 6.0;
310 let y_next = &y + step_div_6 * (&k1 + 2.0 * &k2 + 2.0 * &k3 + &k4);
311
312 x = &x + current_step;
313 y = y_next;
314
315 all_x.push(x);
316 all_y.push(y.copy());
317 }
318
319 Ok((
320 Tensor::from_slice(&all_x).to_kind(kind).to_device(device),
321 Tensor::stack(&all_y, 0),
322 ))
323}
324
325fn solve_implicit_euler(
327 f: tch::CModule,
328 x_span: (f64, f64),
329 y0: Tensor,
330 step: f64,
331 optimizer: &dyn optimizers::Optimizer,
332) -> anyhow::Result<(Tensor, Tensor)> {
333 let device = y0.device();
334 let kind = y0.kind();
335
336 let x_start = x_span.0;
337 let x_end = x_span.1;
338
339 let mut x = x_start;
340 let mut y = y0.copy();
341
342 let mut all_x = vec![x];
343 let mut all_y = vec![y.copy()];
344
345 let mut current_step = step;
346 while x < x_end {
347 let remaining = x_end - x;
348 if remaining < current_step {
349 current_step = remaining;
350 }
351
352 let x_next = &x + current_step;
353 let y_prev = y.copy();
354
355 let y_next = optimizer
356 .optimize(
357 &|y_next: &Tensor| {
358 let f_next = f
359 .forward_ts(&[
360 Tensor::from(x_next).to_kind(kind).to_device(device),
361 y_next.copy(),
362 ])
363 .unwrap();
364 let y_pred = &y_prev + current_step * &f_next;
365 (y_next - &y_pred).pow_tensor_scalar(2).sum(y_next.kind())
366 },
367 &(&y_prev.detach()
368 + current_step
369 * f.forward_ts(&[
370 &Tensor::from(x).to_kind(kind).to_device(device),
371 &y_prev,
372 ])?),
373 )
374 .map_err(|err| anyhow!(format!("Optimizer failed with: {}", err)))?;
375
376 y = y_next.copy();
377 x = x_next;
378
379 all_x.push(x);
380 all_y.push(y.copy());
381 }
382
383 Ok((
384 Tensor::from_slice(&all_x).to_kind(kind).to_device(device),
385 Tensor::stack(&all_y, 0),
386 ))
387}
388
389fn solve_glrk4(
391 f: tch::CModule,
392 x_span: (f64, f64),
393 y0: Tensor,
394 step: f64,
395 optimizer: &dyn optimizers::Optimizer,
396) -> anyhow::Result<(Tensor, Tensor)> {
397 let device = y0.device();
398 let kind = y0.kind();
399
400 let x_start = x_span.0;
401 let x_end = x_span.1;
402
403 let mut x = x_start;
404 let mut y = y0.copy();
405 let y_length = y.size()[0];
406
407 let mut all_x = vec![x];
408 let mut all_y = vec![y.copy()];
409
410 let mut current_step = step;
411 while x < x_end {
412 let remaining = x_end - x;
413 if remaining < current_step {
414 current_step = remaining;
415 }
416
417 let k = f.forward_ts(&[
418 Tensor::from(x).to_kind(kind).to_device(device),
419 y.copy(),
420 ])?;
421 let k_size = k.size();
422 let k_rank = k_size.len();
423 if k_rank != 1 {
424 anyhow::bail!(
425 "Derivative CModule returned tensor of bad rank {}.",
426 k_rank
427 );
428 }
429 if k_size[0] != y0.size()[0] {
430 anyhow::bail!(
431 "Derivative CModule returned vector of bad length {}.",
432 k_size[0]
433 );
434 }
435
436 const C1: f64 = 0.2113248654f64;
437 const C2: f64 = 0.7886751346f64;
438 const A11: f64 = 0.25;
439 const A12: f64 = -0.03867513459f64;
440 const A21: f64 = 0.5386751346f64;
441 const A22: f64 = 0.25;
442
443 let first_k1k2_guess = Tensor::cat(
444 &[
445 f.forward_ts(&[
446 Tensor::from(x + C1 * current_step)
447 .to_kind(kind)
448 .to_device(device),
449 &y + C1 * current_step * &k,
450 ])?,
451 f.forward_ts(&[
452 Tensor::from(x + C2 * current_step)
453 .to_kind(kind)
454 .to_device(device),
455 &y + C2 * current_step * &k,
456 ])?,
457 ],
458 0,
459 );
460 let k1k2 = optimizer
461 .optimize(
462 &|k1k2_guess| {
463 let diff1 = k1k2_guess.i(0..y_length)
464 - f.forward_ts(&[
465 Tensor::from(x + C1 * current_step)
466 .to_kind(kind)
467 .to_device(device),
468 &y
469 + (A11 * k1k2_guess.i(0..y_length) + A12 * k1k2_guess.i(y_length..2*y_length))
470 * current_step,
471 ])
472 .unwrap();
473 let diff2 = k1k2_guess.i(y_length..2*y_length)
474 - f.forward_ts(&[
475 Tensor::from(x + C2 * current_step)
476 .to_kind(kind)
477 .to_device(device),
478 &y
479 + (A21 * k1k2_guess.i(0..y_length) + A22 * k1k2_guess.i(y_length..2*y_length))
480 * current_step,
481 ])
482 .unwrap();
483
484 diff1.dot(&diff1) + diff2.dot(&diff2)
485 },
486 &first_k1k2_guess,
487 )
488 .map_err(|err| anyhow!(format!("Optimizer failed with: {}", err)))?;
489
490 x = x + current_step;
491 y = &y + current_step * (0.5 * k1k2.i(0..y_length) + 0.5 * k1k2.i(y_length..2*y_length));
492
493 all_x.push(x);
494 all_y.push(y.copy());
495 }
496
497 Ok((
498 Tensor::from_slice(&all_x).to_kind(kind).to_device(device),
499 Tensor::stack(&all_y, 0),
500 ))
501}
502
503fn solve_rkf45(
505 f: tch::CModule,
506 x_span: (f64, f64),
507 y0: Tensor,
508 rtol: f64,
509 atol: f64,
510 min_step: f64,
511 safety_factor: f64,
512) -> anyhow::Result<(Tensor, Tensor)> {
513 let device = y0.device();
514 let kind = y0.kind();
515
516 let x_start = x_span.0;
517 let x_end = x_span.1;
518
519 let mut x = x_start;
520 let mut y = y0.copy();
521
522 let mut all_x = vec![x];
523 let mut all_y = vec![y.copy()];
524
525 let mut step = (x_end - x_start) * 0.1;
526
527 while x < x_end {
528 let remaining = x_end - x;
529 if remaining < step {
530 step = remaining;
531 }
532
533 let k1 = f.forward_ts(&[
534 Tensor::from(x).to_kind(kind).to_device(device),
535 y.copy(),
536 ])?;
537 let k1_size = k1.size();
538 let k1_rank = k1_size.len();
539 if k1_rank != 1 {
540 anyhow::bail!(
541 "Derivative CModule returned tensor of bad rank {}.",
542 k1_rank
543 );
544 }
545 if k1_size[0] != y0.size()[0] {
546 anyhow::bail!(
547 "Derivative CModule returned vector of bad length {}.",
548 k1_size[0]
549 );
550 }
551
552 let k2 = {
553 let x_step = x + 0.25 * step;
554 let y_step: Tensor = &y + 0.25 * &step * &k1;
555 let k2_unchecked = f.forward_ts(&[
556 Tensor::from(x_step).to_kind(kind).to_device(device),
557 y_step,
558 ])?;
559 let k2_size = k2_unchecked.size();
560 let k2_rank = k2_size.len();
561 if k2_rank != 1 {
562 anyhow::bail!(
563 "Derivative CModule returned tensor of bad rank {}.",
564 k2_rank
565 );
566 }
567 if k2_size[0] != y0.size()[0] {
568 anyhow::bail!(
569 "Derivative CModule returned vector of bad length {}.",
570 k2_size[0]
571 );
572 }
573
574 k2_unchecked
575 };
576
577 let k3 = {
578 let x_step = x + 0.375 * step;
579 let y_step: Tensor = &y + (0.09375 * &step * &k1) + (0.28125 * &step * &k2);
580 let k3_unchecked = f.forward_ts(&[
581 Tensor::from(x_step).to_kind(kind).to_device(device),
582 y_step,
583 ])?;
584 let k3_size = k3_unchecked.size();
585 let k3_rank = k3_size.len();
586 if k3_rank != 1 {
587 anyhow::bail!(
588 "Derivative CModule returned tensor of bad rank {}.",
589 k3_rank
590 );
591 }
592 if k3_size[0] != y0.size()[0] {
593 anyhow::bail!(
594 "Derivative CModule returned vector of bad length {}.",
595 k3_size[0]
596 );
597 }
598
599 k3_unchecked
600 };
601
602 let k4 = {
603 let x_step = x + (12.0 / 13.0) * step;
604 let y_step: Tensor = &y
605 + (1932.0 / 2197.0 * &step * &k1)
606 + (-7200.0 / 2197.0 * &step * &k2)
607 + (7296.0 / 2197.0 * &step * &k3);
608 let k4_unchecked = f.forward_ts(&[
609 Tensor::from(x_step).to_kind(kind).to_device(device),
610 y_step,
611 ])?;
612 let k4_size = k4_unchecked.size();
613 let k4_rank = k4_size.len();
614 if k4_rank != 1 {
615 anyhow::bail!(
616 "Derivative CModule returned tensor of bad rank {}.",
617 k4_rank
618 );
619 }
620 if k4_size[0] != y0.size()[0] {
621 anyhow::bail!(
622 "Derivative CModule returned vector of bad length {}.",
623 k4_size[0]
624 );
625 }
626
627 k4_unchecked
628 };
629
630 let k5 = {
631 let x_step = x + step;
632 let y_step: Tensor = &y
633 + (439.0 / 216.0 * &step * &k1)
634 + (-8.0 * &step * &k2)
635 + (3680.0 / 513.0 * &step * &k3)
636 + (-845.0 / 4104.0 * &step * &k4);
637 let k5_unchecked = f.forward_ts(&[
638 Tensor::from(x_step).to_kind(kind).to_device(device),
639 y_step,
640 ])?;
641 let k5_size = k5_unchecked.size();
642 let k5_rank = k5_size.len();
643 if k5_rank != 1 {
644 anyhow::bail!(
645 "Derivative CModule returned tensor of bad rank {}.",
646 k5_rank
647 );
648 }
649 if k5_size[0] != y0.size()[0] {
650 anyhow::bail!(
651 "Derivative CModule returned vector of bad length {}.",
652 k5_size[0]
653 );
654 }
655
656 k5_unchecked
657 };
658
659 let k6 = {
660 let x_step = x + 0.5 * step;
661 let y_step: Tensor = &y
662 + (-8.0 / 27.0 * &step * &k1)
663 + (2.0 * &step * &k2)
664 + (-3544.0 / 2565.0 * &step * &k3)
665 + (1859.0 / 4104.0 * &step * &k4)
666 + (-11.0 / 40.0 * &step * &k5);
667 let k6_unchecked = f.forward_ts(&[
668 Tensor::from(x_step).to_kind(kind).to_device(device),
669 y_step,
670 ])?;
671 let k6_size = k6_unchecked.size();
672 let k6_rank = k6_size.len();
673 if k6_rank != 1 {
674 anyhow::bail!(
675 "Derivative CModule returned tensor of bad rank {}.",
676 k6_rank
677 );
678 }
679 if k6_size[0] != y0.size()[0] {
680 anyhow::bail!(
681 "Derivative CModule returned vector of bad length {}.",
682 k6_size[0]
683 );
684 }
685
686 k6_unchecked
687 };
688
689 let next_y4: Tensor = &y
690 + step
691 * ((25.0 / 216.0 * &k1)
692 + (1408.0 / 2565.0 * &k3)
693 + (2197.0 / 4104.0 * &k4)
694 + (-1.0 / 5.0 * &k5));
695 let next_y5: Tensor = &y
696 + step
697 * ((16.0 / 135.0 * &k1)
698 + (6656.0 / 12825.0 * &k3)
699 + (28561.0 / 56430.0 * &k4)
700 + (-9.0 / 50.0 * &k5)
701 + (2.0 / 55.0 * &k6));
702
703 let d = (&next_y4 - &next_y5).abs();
704 let e = next_y5.abs() * rtol + atol;
705
706 let alpha = (e / d).sqrt().min().double_value(&[]);
707 let condition = safety_factor * alpha;
708
709 if condition < 1f64 {
710 step = step * condition;
711 if step < min_step {
712 return Err(anyhow!("Required step is smaller than minimal step"));
713 }
714 } else {
715 y = next_y4;
716 x = &x + &step;
717 all_x.push(x);
718 all_y.push(y.copy());
719
720 let new_step = step * condition;
721 let max_step = step * 5.0;
722 step = match new_step.partial_cmp(&max_step) {
723 Some(std::cmp::Ordering::Less) => new_step,
724 Some(_) => max_step,
725 None => return Err(anyhow!("Runtime error")),
726 };
727 }
728 }
729
730 Ok((
731 Tensor::from_slice(&all_x).to_kind(kind).to_device(device),
732 Tensor::stack(&all_y, 0),
733 ))
734}
735
736fn solve_row1(
738 f: tch::CModule,
739 x_span: (f64, f64),
740 y0: Tensor,
741 step: f64,
742) -> anyhow::Result<(Tensor, Tensor)> {
743 let device = y0.device();
744 let kind = y0.kind();
745
746 let x_start = x_span.0;
747 let x_end = x_span.1;
748
749 let mut x = x_start;
750 let mut y = y0.copy();
751
752 let mut all_x = vec![x];
753 let mut all_y = vec![y.copy()];
754
755 while x < x_end {
756 let remaining = x_end - x;
757 let mut current_step = step;
758 if remaining < step {
759 current_step = remaining;
760 }
761
762 let x_prev = x;
763 let y_prev = y.copy();
764
765 let jacobian = compute_jacobian(
766 |y| {
767 f.forward_ts(&[
768 Tensor::from(x_prev).to_kind(kind).to_device(device),
769 y.copy(),
770 ])
771 .unwrap()
772 },
773 &y_prev,
774 );
775 let f_current = f.forward_ts(&[
776 Tensor::from(x_prev).to_kind(kind).to_device(device),
777 y_prev.copy(),
778 ])?;
779 let f_current_size = f_current.size();
780 let f_current_rank = f_current_size.len();
781 if f_current_rank != 1 {
782 anyhow::bail!(
783 "Derivative CModule returned tensor of bad rank {}.",
784 f_current_rank
785 );
786 }
787 if f_current_size[0] != y0.size()[0] {
788 anyhow::bail!(
789 "Derivative CModule returned vector of bad length {}.",
790 f_current_size[0]
791 );
792 }
793
794 let n = jacobian.size()[0];
795 let eye = Tensor::eye(n, (tch::Kind::Float, jacobian.device()));
796 let step_j = current_step * &jacobian;
797 let inv_matrix = (eye - step_j).inverse();
798
799 let delta_y = inv_matrix.matmul(&f_current);
800 let y_next = y_prev + current_step * delta_y;
801
802 x = &x_prev + current_step;
803 y = y_next.copy();
804
805 all_x.push(x);
806 all_y.push(y.copy());
807 }
808
809 Ok((
810 Tensor::from_slice(&all_x).to_kind(kind).to_device(device),
811 Tensor::stack(&all_y, 0),
812 ))
813}
814
815fn compute_jacobian<F>(f: F, x: &Tensor) -> Tensor
817where
818 F: Fn(&Tensor) -> Tensor,
819{
820 assert_eq!(x.dim(), 1, "x must be 1-dimensional");
821 let mut x_with_grad = x.detach().copy().set_requires_grad(true);
822 let y = f(&x_with_grad);
823 assert_eq!(y.dim(), 1, "y must be 1-dimensional");
824
825 let y_size = y.size()[0];
826 let mut grads = Vec::new();
827
828 for i in 0..y_size {
829 let yi = y.i(i);
830 let grad = Tensor::run_backward(&[yi], &[&x_with_grad], true, false)[0].copy();
831 grads.push(grad);
832 x_with_grad.zero_grad();
833 }
834
835 Tensor::stack(&grads, 0)
836}