1use scirs2_core::ndarray::{Array, Dimension, IxDyn, ScalarOperand, Zip};
6use scirs2_core::numeric::Float;
7use std::fmt::Debug;
8
9use crate::error::{OptimError, Result};
10use crate::optimizers::Optimizer;
11
12#[derive(Debug, Clone)]
47pub struct AdamW<A: Float + ScalarOperand + Debug> {
48 learning_rate: A,
50 beta1: A,
52 beta2: A,
54 epsilon: A,
56 weight_decay: A,
58 m: Option<Vec<Array<A, IxDyn>>>,
60 v: Option<Vec<Array<A, IxDyn>>>,
62 t: Vec<usize>,
67}
68
69impl<A: Float + ScalarOperand + Debug + Send + Sync> AdamW<A> {
70 pub fn new(learning_rate: A) -> Self {
76 Self {
77 learning_rate,
78 beta1: A::from(0.9).expect("AdamW: default beta1 (0.9) must fit in A"),
79 beta2: A::from(0.999).expect("AdamW: default beta2 (0.999) must fit in A"),
80 epsilon: A::from(1e-8).expect("AdamW: default epsilon (1e-8) must fit in A"),
81 weight_decay: A::from(0.01).expect("AdamW: default weight_decay (0.01) must fit in A"),
83 m: None,
84 v: None,
85 t: Vec::new(),
86 }
87 }
88
89 pub fn new_with_config(
99 learning_rate: A,
100 beta1: A,
101 beta2: A,
102 epsilon: A,
103 weight_decay: A,
104 ) -> Self {
105 Self {
106 learning_rate,
107 beta1,
108 beta2,
109 epsilon,
110 weight_decay,
111 m: None,
112 v: None,
113 t: Vec::new(),
114 }
115 }
116
117 pub fn set_beta1(&mut self, beta1: A) -> &mut Self {
119 self.beta1 = beta1;
120 self
121 }
122
123 pub fn get_beta1(&self) -> A {
125 self.beta1
126 }
127
128 pub fn set_beta2(&mut self, beta2: A) -> &mut Self {
130 self.beta2 = beta2;
131 self
132 }
133
134 pub fn get_beta2(&self) -> A {
136 self.beta2
137 }
138
139 pub fn set_epsilon(&mut self, epsilon: A) -> &mut Self {
141 self.epsilon = epsilon;
142 self
143 }
144
145 pub fn get_epsilon(&self) -> A {
147 self.epsilon
148 }
149
150 pub fn set_weight_decay(&mut self, weight_decay: A) -> &mut Self {
152 self.weight_decay = weight_decay;
153 self
154 }
155
156 pub fn get_weight_decay(&self) -> A {
158 self.weight_decay
159 }
160
161 pub fn learning_rate(&self) -> A {
163 self.learning_rate
164 }
165
166 pub fn set_lr(&mut self, lr: A) {
168 self.learning_rate = lr;
169 }
170
171 pub fn reset(&mut self) {
173 self.m = None;
174 self.v = None;
175 self.t.clear();
176 }
177
178 pub fn timestep(&self, index: usize) -> usize {
182 self.t.get(index).copied().unwrap_or(0)
183 }
184
185 fn advance_state(&mut self, index: usize, dim: &IxDyn) -> Result<usize> {
187 let m = self.m.get_or_insert_with(Vec::new);
188 let v = self.v.get_or_insert_with(Vec::new);
189 while m.len() <= index {
190 m.push(Array::zeros(dim.clone()));
191 }
192 while v.len() <= index {
193 v.push(Array::zeros(dim.clone()));
194 }
195 while self.t.len() <= index {
196 self.t.push(0);
197 }
198
199 if m[index].raw_dim() != *dim || v[index].raw_dim() != *dim {
201 m[index] = Array::zeros(dim.clone());
202 v[index] = Array::zeros(dim.clone());
203 self.t[index] = 0;
204 }
205
206 let next = self.t[index].checked_add(1).ok_or_else(|| {
207 OptimError::InvalidConfig(
208 "Timestep counter overflow - too many optimization steps".to_string(),
209 )
210 })?;
211 self.t[index] = next;
212 Ok(next)
213 }
214
215 pub fn step_inplace_indexed<D: Dimension>(
220 &mut self,
221 index: usize,
222 params: &mut Array<A, D>,
223 gradients: &Array<A, D>,
224 ) -> Result<()> {
225 if params.shape() != gradients.shape() {
226 return Err(OptimError::DimensionMismatch(format!(
227 "Incompatible shapes: parameters have shape {:?}, gradients have shape {:?}",
228 params.shape(),
229 gradients.shape()
230 )));
231 }
232
233 let dim = params.raw_dim().into_dyn();
234 let t = self.advance_state(index, &dim)?;
235 let exp = i32::try_from(t).map_err(|_| {
236 OptimError::InvalidConfig(
237 "Timestep too large for bias correction calculation".to_string(),
238 )
239 })?;
240
241 let beta1 = self.beta1;
242 let beta2 = self.beta2;
243 let lr = self.learning_rate;
244 let eps = self.epsilon;
245 let one = A::one();
246 let bias_correction1 = one - beta1.powi(exp);
247 let bias_correction2 = one - beta2.powi(exp);
248 let weight_decay_factor = one - lr * self.weight_decay;
250
251 let m = self
252 .m
253 .as_mut()
254 .ok_or_else(|| OptimError::InvalidConfig("AdamW state not initialized".to_string()))?;
255 let v = self
256 .v
257 .as_mut()
258 .ok_or_else(|| OptimError::InvalidConfig("AdamW state not initialized".to_string()))?;
259
260 let mut params_view = params.view_mut().into_dyn();
261 let gradients_view = gradients.view().into_dyn();
262
263 Zip::from(&mut params_view)
264 .and(&gradients_view)
265 .and(&mut m[index])
266 .and(&mut v[index])
267 .for_each(|p, &g, m_i, v_i| {
268 *m_i = *m_i * beta1 + g * (one - beta1);
269 *v_i = *v_i * beta2 + g * g * (one - beta2);
270 let m_hat = *m_i / bias_correction1;
271 let v_hat = *v_i / bias_correction2;
272 *p = *p * weight_decay_factor - lr * m_hat / (v_hat.sqrt() + eps);
273 });
274
275 Ok(())
276 }
277
278 pub fn step_inplace<D: Dimension>(
280 &mut self,
281 params: &mut Array<A, D>,
282 gradients: &Array<A, D>,
283 ) -> Result<()> {
284 self.step_inplace_indexed(0, params, gradients)
285 }
286
287 pub fn step_indexed<D: Dimension>(
292 &mut self,
293 index: usize,
294 params: &Array<A, D>,
295 gradients: &Array<A, D>,
296 ) -> Result<Array<A, D>> {
297 let mut updated = params.to_owned();
298 self.step_inplace_indexed(index, &mut updated, gradients)?;
299 Ok(updated)
300 }
301}
302
303impl<A, D> Optimizer<A, D> for AdamW<A>
304where
305 A: Float + ScalarOperand + Debug + Send + Sync,
306 D: Dimension,
307{
308 fn step(&mut self, params: &Array<A, D>, gradients: &Array<A, D>) -> Result<Array<A, D>> {
309 self.step_indexed(0, params, gradients)
310 }
311
312 fn step_list(
313 &mut self,
314 params_list: &[&Array<A, D>],
315 gradients_list: &[&Array<A, D>],
316 ) -> Result<Vec<Array<A, D>>> {
317 if params_list.len() != gradients_list.len() {
318 return Err(OptimError::InvalidConfig(format!(
319 "Number of parameter arrays ({}) does not match number of gradient arrays ({})",
320 params_list.len(),
321 gradients_list.len()
322 )));
323 }
324
325 let mut results = Vec::with_capacity(params_list.len());
326 for (index, (params, grads)) in params_list.iter().zip(gradients_list.iter()).enumerate() {
327 results.push(self.step_indexed(index, params, grads)?);
328 }
329 Ok(results)
330 }
331
332 fn get_learning_rate(&self) -> A {
333 self.learning_rate
334 }
335
336 fn set_learning_rate(&mut self, learning_rate: A) {
337 self.learning_rate = learning_rate;
338 }
339}
340
341#[cfg(test)]
342mod tests {
343 use super::*;
344 use scirs2_core::ndarray::Array1;
345
346 #[test]
347 fn test_adamw_step() {
348 let params = Array1::zeros(3);
350 let gradients = Array1::from_vec(vec![0.1, 0.2, 0.3]);
351
352 let mut optimizer = AdamW::new(0.01);
354
355 let new_params = optimizer
357 .step(¶ms, &gradients)
358 .expect("optimizer.step succeeds in test_adamw_step");
359
360 assert!(new_params.iter().all(|&x| x != 0.0));
362
363 for param in new_params.iter() {
366 assert!(*param < 0.0);
367 }
368 }
369
370 #[test]
371 fn test_adamw_multiple_steps() {
372 let mut params = Array1::zeros(3);
374 let gradients = Array1::from_vec(vec![0.1, 0.2, 0.3]);
375
376 let mut optimizer = AdamW::new_with_config(
378 0.01, 0.9, 0.999, 1e-8, 0.1, );
380
381 for _ in 0..10 {
383 params = optimizer
384 .step(¶ms, &gradients)
385 .expect("optimizer.step succeeds in test_adamw_multiple_steps");
386 }
387
388 for (i, param) in params.iter().enumerate() {
390 assert!(*param < 0.0);
392 if i > 0 {
393 assert!(param < ¶ms[i - 1]);
395 }
396 }
397 }
398
399 #[test]
418 fn test_adamw_reset() {
419 let params = Array1::zeros(3);
421 let gradients = Array1::from_vec(vec![0.1, 0.2, 0.3]);
422
423 let mut optimizer = AdamW::new(0.01);
425
426 optimizer
428 .step(¶ms, &gradients)
429 .expect("optimizer.step succeeds in test_adamw_reset");
430 assert_eq!(optimizer.timestep(0), 1);
431 assert!(optimizer.m.is_some());
432 assert!(optimizer.v.is_some());
433
434 optimizer.reset();
436 assert_eq!(optimizer.timestep(0), 0);
437 assert!(optimizer.m.is_none());
438 assert!(optimizer.v.is_none());
439 }
440}