Skip to main content

mlx_rs/optimizers/
adadelta.rs

1use std::rc::Rc;
2
3use crate::{
4    array,
5    ops::sqrt,
6    utils::{get_mut_or_insert_with, Updatable},
7    Array,
8};
9use mlx_internal_macros::{generate_builder, Buildable};
10
11use crate::error::AdaDeltaBuildError;
12
13use super::*;
14
15generate_builder! {
16    /// The AdaDelta optimizer with a learning rate.
17    ///
18    /// The default `rho` is `0.9`, matching Python MLX. Earlier mlx-rs releases used `0.99`.
19    ///
20    /// Please refer to the original paper for more details:
21    ///
22    /// [1]: Zeiler, M.D., 2012. ADADELTA: an adaptive learning rate method. arXiv preprint arXiv:1212.5701.
23    #[derive(Debug, Clone, Buildable)]
24    #[buildable(root = crate)]
25    #[builder(
26        build_with = build_adadelta,
27        err = AdaDeltaBuildError,
28        root = crate
29    )]
30    pub struct AdaDelta {
31        /// The learning rate
32        #[builder(ty_override = f32)]
33        pub lr: Array,
34
35        /// The coefficient used for computing a running average of squared gradients. Defaults to
36        /// [`AdaDelta::DEFAULT_RHO`] (`0.9`, matching Python MLX).
37        #[builder(optional, ty_override = f32, default = AdaDelta::DEFAULT_RHO)]
38        pub rho: Array,
39
40        /// The epsilon added to the denominator to improve numerical stability. Default to
41        /// [`AdaDelta::DEFAULT_EPS`].
42        #[builder(optional, ty_override = f32, default = AdaDelta::DEFAULT_EPS)]
43        pub eps: Array,
44
45        /// Inner state
46        #[builder(ignore)]
47        pub state: State<(Array, Array)>,
48    }
49}
50
51/// Builds a new [`AdaDelta`] optimizer
52fn build_adadelta(builder: AdaDeltaBuilder) -> Result<AdaDelta, AdaDeltaBuildError> {
53    let rho = builder.rho;
54    let eps = builder.eps;
55
56    if rho < 0.0 {
57        return Err(AdaDeltaBuildError::NegativeRho);
58    }
59
60    if eps < 0.0 {
61        return Err(AdaDeltaBuildError::NegativeEps);
62    }
63
64    Ok(AdaDelta {
65        lr: array!(builder.lr),
66        rho: array!(rho),
67        eps: array!(eps),
68        state: State::new(),
69    })
70}
71
72impl AdaDelta {
73    /// Default value for `rho`, matching Python MLX.
74    pub const DEFAULT_RHO: f32 = 0.9;
75
76    /// Default value for `eps`
77    pub const DEFAULT_EPS: f32 = 1e-6;
78}
79
80impl Optimizer for AdaDelta {
81    type State = State<(Array, Array)>;
82
83    fn state(&self) -> &Self::State {
84        &self.state
85    }
86
87    fn state_mut(&mut self) -> &mut Self::State {
88        &mut self.state
89    }
90
91    fn update_single(
92        &mut self,
93        key: &Rc<str>,
94        gradient: &Array,
95        parameter: &mut Array,
96    ) -> crate::error::Result<()> {
97        let (v, u) = get_mut_or_insert_with(&mut self.state, key, || (array!(0.0), array!(0.0)));
98
99        let one_minus_rho = array!(1.0).subtract(&self.rho)?;
100        let first_term = self.rho.multiply(&v)?;
101        let second_term = one_minus_rho.multiply(gradient.square()?)?;
102        let v_new = first_term.add(&second_term)?;
103
104        let num = sqrt(&u.add(&self.eps)?)?;
105        let den = sqrt(&v_new.add(&self.eps)?)?;
106        let d = num.divide(&den)?.multiply(gradient)?;
107        let first_term = self.rho.multiply(&u)?;
108        let second_term = one_minus_rho.multiply(d.square()?)?;
109        let u_new = first_term.add(&second_term)?;
110
111        let param_new = parameter.subtract(self.lr.multiply(d)?)?;
112
113        *parameter = param_new;
114
115        *v = v_new;
116        *u = u_new;
117
118        Ok(())
119    }
120}
121
122impl Updatable for AdaDelta {
123    optimizer_updatable_state_methods!();
124}
125
126impl_updatable_for_mut_optimizer!(AdaDelta);
127
128#[cfg(test)]
129mod tests {
130    use super::AdaDelta;
131
132    #[test]
133    fn default_rho_matches_python_mlx() {
134        let optimizer = AdaDelta::new(0.1_f32).unwrap();
135
136        assert_eq!(optimizer.rho.item_exact::<f32>(), 0.9);
137    }
138}