1use luma_io::lpk::LumaPack;
2use luma_tensor::{Device, DynTensor, GradStore, Scalar, Tensor, no_grad};
3
4use super::Optimizer;
5
6#[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>, second_moment: Tensor<D>, }
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 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(¶m) {
82 let g = g.clone();
83
84 first_moment.mul_scalar_(beta1)?;
86 first_moment.add_(&g.mul_scalar(1.0 - beta1)?)?;
87
88 second_moment.mul_scalar_(beta2)?;
90 second_moment.add_(&g.pow(2.0)?.mul_scalar(1.0 - beta2)?)?;
91
92 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#[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>, second_moment: Tensor<D>, }
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 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(¶m) {
228 let g = g.clone();
229
230 if weight_decay != 0.0 {
232 param.sub_(¶m.mul_scalar(lr * weight_decay)?)?;
233 }
234
235 first_moment.mul_scalar_(beta1)?;
237 first_moment.add_(&g.mul_scalar(1.0 - beta1)?)?;
238
239 second_moment.mul_scalar_(beta2)?;
241 second_moment.add_(&g.pow(2.0)?.mul_scalar(1.0 - beta2)?)?;
242
243 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}