Skip to main content

ruda_optim/optim/simple/
base.rs

1
2use crate::LearningRate;
3use ruda_model::record::Record;
4use ruda_model::tensor::{Tensor, backend::Backend};
5
6/// Simple optimizer is an opinionated trait to simplify the process of implementing an
7/// optimizer.
8///
9/// Implementations don't have to handle missing gradients, loading and exporting records, navigate the
10/// module parameter structure, handle tracked and untracked tensors, and the likes.
11pub trait SimpleOptimizer<B>: Send + Sync + Clone
12where
13    B: Backend,
14{
15    /// The state of the optimizer. It also implements [record](Record), so that it can be saved.
16    type State<const D: usize>: Record<B> + Clone + 'static;
17
18    /// The optimizer step is performed for one tensor at a time with its gradient and state.
19    ///
20    /// Note that the state is passed as parameter, so implementations don't have to handle
21    /// the saving and loading of recorded states.
22    fn step<const D: usize>(
23        &self,
24        lr: LearningRate,
25        tensor: Tensor<B, D>,
26        grad: Tensor<B, D>,
27        state: Option<Self::State<D>>,
28    ) -> (Tensor<B, D>, Option<Self::State<D>>);
29
30    /// Change the device of the state.
31    ///
32    /// This function will be called accordingly to have the state on the same device as the
33    /// gradient and the tensor when the [step](SimpleOptimizer::step) function is called.
34    fn to_device<const D: usize>(state: Self::State<D>, device: &B::Device) -> Self::State<D>;
35}