mlx_rs/optimizers/
adadelta.rs1use 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 #[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 #[builder(ty_override = f32)]
33 pub lr: Array,
34
35 #[builder(optional, ty_override = f32, default = AdaDelta::DEFAULT_RHO)]
38 pub rho: Array,
39
40 #[builder(optional, ty_override = f32, default = AdaDelta::DEFAULT_EPS)]
43 pub eps: Array,
44
45 #[builder(ignore)]
47 pub state: State<(Array, Array)>,
48 }
49}
50
51fn 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 pub const DEFAULT_RHO: f32 = 0.9;
75
76 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}