manopt_rs/optimizers/
many_steps.rs1use crate::prelude::*;
6use burn::optim::{LearningRate, SimpleOptimizer};
7
8pub trait LessSimpleOptimizer<B: Backend>: SimpleOptimizer<B> {
13 fn many_steps<const D: usize>(
14 &self,
15 lr_function: impl FnMut(usize) -> LearningRate,
16 num_steps: usize,
17 grad_function: impl FnMut(Tensor<B, D>) -> Tensor<B, D>,
18 tensor: Tensor<B, D>,
19 state: Option<Self::State<D>>,
20 ) -> (Tensor<B, D>, Option<Self::State<D>>);
21}
22
23impl<B: Backend, T: SimpleOptimizer<B>> LessSimpleOptimizer<B> for T {
28 #[inline]
29 fn many_steps<const D: usize>(
30 &self,
31 mut lr_function: impl FnMut(usize) -> LearningRate,
32 num_steps: usize,
33 mut grad_function: impl FnMut(Tensor<B, D>) -> Tensor<B, D>,
34 mut tensor: Tensor<B, D>,
35 mut state: Option<Self::State<D>>,
36 ) -> (Tensor<B, D>, Option<Self::State<D>>) {
37 for step in 0..num_steps {
39 let cur_grad = grad_function(tensor.clone());
41 let cur_lr = lr_function(step);
43 let (new_x, new_state) = self.step(cur_lr, tensor.clone(), cur_grad, state);
45 tensor = new_x.detach().require_grad();
46 state = new_state;
47 }
48 (tensor, state)
49 }
50}