Skip to main content

luma_optim/
adam.rs

1use luma_io::lpk::LumaPack;
2use luma_tensor::{Device, DynTensor, GradStore, Scalar, Tensor, no_grad};
3
4use super::Optimizer;
5
6// ============================================================================
7//   Adam
8// ============================================================================
9
10#[derive(Clone, Debug)]
11pub struct AdamConfig {
12    pub lr: f64,
13    pub beta1: f64,
14    pub beta2: f64,
15    pub eps: f64,
16}
17
18impl Default for AdamConfig {
19    fn default() -> Self {
20        Self { lr: 1e-3, beta1: 0.9, beta2: 0.999, eps: 1e-8 }
21    }
22}
23
24struct AdamParam<D: Device> {
25    param: Tensor<D>,
26    first_moment: Tensor<D>,  // m
27    second_moment: Tensor<D>, // v
28}
29
30pub struct Adam<D: Device> {
31    params: Vec<AdamParam<D>>,
32    step_t: usize,
33    config: AdamConfig,
34}
35
36impl<D: Device> Adam<D> {
37    pub fn new(params: impl Into<Vec<Tensor<D>>>, config: AdamConfig) -> luma_tensor::Result<Self> {
38        let params = params
39            .into()
40            .into_iter()
41            .map(|param| {
42                let first_moment = param.zeros_like()?;
43                let second_moment = param.zeros_like()?;
44                Ok(AdamParam { param, first_moment, second_moment })
45            })
46            .collect::<luma_tensor::Result<Vec<_>>>()?;
47        Ok(Self { params, step_t: 0, config })
48    }
49}
50
51impl<D: Device> Optimizer for Adam<D> {
52    type Device = D;
53
54    fn get_lr(&self) -> f64 {
55        self.config.lr
56    }
57
58    fn set_lr(&mut self, lr: f64) {
59        self.config.lr = lr;
60    }
61
62    /// ```text
63    ///   m = β₁·m + (1-β₁)·g
64    ///   v = β₂·v + (1-β₂)·g²
65    ///   m̂ = m / (1-β₁ᵗ)    v̂ = v / (1-β₂ᵗ)
66    ///   param -= lr * m̂ / (√v̂ + ε)
67    /// ```
68    fn step(&mut self, grads: &GradStore<Self::Device>) -> luma_tensor::Result<()> {
69        no_grad!();
70        self.step_t += 1;
71
72        let lr = self.config.lr;
73        let beta1 = self.config.beta1;
74        let beta2 = self.config.beta2;
75        let eps = self.config.eps;
76
77        let bias_m = 1.0 - beta1.powi(self.step_t as i32);
78        let bias_v = 1.0 - beta2.powi(self.step_t as i32);
79
80        for AdamParam { param, first_moment, second_moment } in self.params.iter_mut() {
81            if let Some(g) = grads.get(&param) {
82                let g = g.clone();
83
84                // m = β₁·m + (1-β₁)·g
85                first_moment.mul_scalar_(beta1)?;
86                first_moment.add_(&g.mul_scalar(1.0 - beta1)?)?;
87
88                // v = β₂·v + (1-β₂)·g²
89                second_moment.mul_scalar_(beta2)?;
90                second_moment.add_(&g.pow(2.0)?.mul_scalar(1.0 - beta2)?)?;
91
92                // bias-corrected estimates
93                let m_hat = first_moment.mul_scalar(1.0 / bias_m)?;
94                let v_hat = second_moment.mul_scalar(1.0 / bias_v)?;
95                let denom = v_hat.sqrt()?.add_scalar(eps)?;
96                param.sub_(&m_hat.div(&denom)?.mul_scalar(lr)?)?;
97            }
98        }
99
100        Ok(())
101    }
102
103    fn state_dict(&self) -> luma_tensor::Result<LumaPack<Self::Device>> {
104        let mut pack = LumaPack::new();
105        for (i, p) in self.params.iter().enumerate() {
106            pack.tensors.insert(format!("{i}.first_moment"), DynTensor::Float(p.first_moment.clone()));
107            pack.tensors.insert(format!("{i}.second_moment"), DynTensor::Float(p.second_moment.clone()));
108        }
109        pack.scalars.insert("lr".into(), Scalar::F64(self.config.lr));
110        pack.scalars.insert("beta1".into(), Scalar::F64(self.config.beta1));
111        pack.scalars.insert("beta2".into(), Scalar::F64(self.config.beta2));
112        pack.scalars.insert("eps".into(), Scalar::F64(self.config.eps));
113        pack.scalars.insert("step_t".into(), Scalar::I32(self.step_t as i32));
114        Ok(pack)
115    }
116
117    fn load_state_dict(&mut self, pack: &LumaPack<Self::Device>) -> luma_tensor::Result<()> {
118        if let Some(v) = pack.scalars.get("lr").and_then(|s| s.to_f64()) {
119            self.config.lr = v;
120        }
121        if let Some(v) = pack.scalars.get("beta1").and_then(|s| s.to_f64()) {
122            self.config.beta1 = v;
123        }
124        if let Some(v) = pack.scalars.get("beta2").and_then(|s| s.to_f64()) {
125            self.config.beta2 = v;
126        }
127        if let Some(v) = pack.scalars.get("eps").and_then(|s| s.to_f64()) {
128            self.config.eps = v;
129        }
130        if let Some(v) = pack.scalars.get("step_t").and_then(|s| s.to_i64()) {
131            self.step_t = v as usize;
132        }
133        for (i, p) in self.params.iter_mut().enumerate() {
134            if let Some(dt) = pack.tensors.get(&format!("{i}.first_moment")) {
135                if let Some(src) = dt.as_float() {
136                    p.first_moment.copy_(src)?;
137                }
138            }
139            if let Some(dt) = pack.tensors.get(&format!("{i}.second_moment")) {
140                if let Some(src) = dt.as_float() {
141                    p.second_moment.copy_(src)?;
142                }
143            }
144        }
145        Ok(())
146    }
147}
148
149// ============================================================================
150//   AdamW
151// ============================================================================
152
153#[derive(Clone, Debug)]
154pub struct AdamWConfig {
155    pub lr: f64,
156    pub beta1: f64,
157    pub beta2: f64,
158    pub eps: f64,
159    pub weight_decay: f64,
160}
161
162impl Default for AdamWConfig {
163    fn default() -> Self {
164        Self { lr: 1e-3, beta1: 0.9, beta2: 0.999, eps: 1e-8, weight_decay: 1e-2 }
165    }
166}
167
168struct AdamWParam<D: Device> {
169    param: Tensor<D>,
170    first_moment: Tensor<D>,  // m
171    second_moment: Tensor<D>, // v
172}
173
174pub struct AdamW<D: Device> {
175    params: Vec<AdamWParam<D>>,
176    step_t: usize,
177    config: AdamWConfig,
178}
179
180impl<D: Device> AdamW<D> {
181    pub fn new(params: impl Into<Vec<Tensor<D>>>, config: AdamWConfig) -> luma_tensor::Result<Self> {
182        let params = params
183            .into()
184            .into_iter()
185            .map(|param| {
186                let first_moment = param.zeros_like()?;
187                let second_moment = param.zeros_like()?;
188                Ok(AdamWParam { param, first_moment, second_moment })
189            })
190            .collect::<luma_tensor::Result<Vec<_>>>()?;
191        Ok(Self { params, step_t: 0, config })
192    }
193}
194
195impl<D: Device> Optimizer for AdamW<D> {
196    type Device = D;
197
198    fn get_lr(&self) -> f64 {
199        self.config.lr
200    }
201
202    fn set_lr(&mut self, lr: f64) {
203        self.config.lr = lr;
204    }
205
206    /// ```text
207    ///   param -= lr * weight_decay * param          (decoupled)
208    ///   m = β₁·m + (1-β₁)·g
209    ///   v = β₂·v + (1-β₂)·g²
210    ///   m̂ = m / (1-β₁ᵗ)    v̂ = v / (1-β₂ᵗ)
211    ///   param -= lr * m̂ / (√v̂ + ε)
212    /// ```
213    fn step(&mut self, grads: &GradStore<Self::Device>) -> luma_tensor::Result<()> {
214        no_grad!();
215        self.step_t += 1;
216
217        let lr = self.config.lr;
218        let beta1 = self.config.beta1;
219        let beta2 = self.config.beta2;
220        let eps = self.config.eps;
221        let weight_decay = self.config.weight_decay;
222
223        let bias_m = 1.0 - beta1.powi(self.step_t as i32);
224        let bias_v = 1.0 - beta2.powi(self.step_t as i32);
225
226        for AdamWParam { param, first_moment, second_moment } in self.params.iter_mut() {
227            if let Some(g) = grads.get(&param) {
228                let g = g.clone();
229
230                // decoupled weight decay
231                if weight_decay != 0.0 {
232                    param.sub_(&param.mul_scalar(lr * weight_decay)?)?;
233                }
234
235                // m = β₁·m + (1-β₁)·g
236                first_moment.mul_scalar_(beta1)?;
237                first_moment.add_(&g.mul_scalar(1.0 - beta1)?)?;
238
239                // v = β₂·v + (1-β₂)·g²
240                second_moment.mul_scalar_(beta2)?;
241                second_moment.add_(&g.pow(2.0)?.mul_scalar(1.0 - beta2)?)?;
242
243                // bias-corrected estimates
244                let m_hat = first_moment.mul_scalar(1.0 / bias_m)?;
245                let v_hat = second_moment.mul_scalar(1.0 / bias_v)?;
246                let denom = v_hat.sqrt()?.add_scalar(eps)?;
247                param.sub_(&m_hat.div(&denom)?.mul_scalar(lr)?)?;
248            }
249        }
250
251        Ok(())
252    }
253
254    fn state_dict(&self) -> luma_tensor::Result<LumaPack<Self::Device>> {
255        let mut pack = LumaPack::new();
256        for (i, p) in self.params.iter().enumerate() {
257            pack.tensors.insert(format!("{i}.first_moment"), DynTensor::Float(p.first_moment.clone()));
258            pack.tensors.insert(format!("{i}.second_moment"), DynTensor::Float(p.second_moment.clone()));
259        }
260        pack.scalars.insert("lr".into(), Scalar::F64(self.config.lr));
261        pack.scalars.insert("beta1".into(), Scalar::F64(self.config.beta1));
262        pack.scalars.insert("beta2".into(), Scalar::F64(self.config.beta2));
263        pack.scalars.insert("eps".into(), Scalar::F64(self.config.eps));
264        pack.scalars.insert("weight_decay".into(), Scalar::F64(self.config.weight_decay));
265        pack.scalars.insert("step_t".into(), Scalar::I32(self.step_t as i32));
266        Ok(pack)
267    }
268
269    fn load_state_dict(&mut self, pack: &LumaPack<Self::Device>) -> luma_tensor::Result<()> {
270        if let Some(v) = pack.scalars.get("lr").and_then(|s| s.to_f64()) {
271            self.config.lr = v;
272        }
273        if let Some(v) = pack.scalars.get("beta1").and_then(|s| s.to_f64()) {
274            self.config.beta1 = v;
275        }
276        if let Some(v) = pack.scalars.get("beta2").and_then(|s| s.to_f64()) {
277            self.config.beta2 = v;
278        }
279        if let Some(v) = pack.scalars.get("eps").and_then(|s| s.to_f64()) {
280            self.config.eps = v;
281        }
282        if let Some(v) = pack.scalars.get("weight_decay").and_then(|s| s.to_f64()) {
283            self.config.weight_decay = v;
284        }
285        if let Some(v) = pack.scalars.get("step_t").and_then(|s| s.to_i64()) {
286            self.step_t = v as usize;
287        }
288        for (i, p) in self.params.iter_mut().enumerate() {
289            if let Some(dt) = pack.tensors.get(&format!("{i}.first_moment")) {
290                if let Some(src) = dt.as_float() {
291                    p.first_moment.copy_(src)?;
292                }
293            }
294            if let Some(dt) = pack.tensors.get(&format!("{i}.second_moment")) {
295                if let Some(src) = dt.as_float() {
296                    p.second_moment.copy_(src)?;
297                }
298            }
299        }
300        Ok(())
301    }
302}