1use crate::error::{OptimError, Result};
4use crate::optimizers::Optimizer;
5use crate::parameter_groups::{
6 GroupManager, GroupedOptimizer, ParameterGroup, ParameterGroupConfig,
7};
8use scirs2_core::ndarray::{Array, Dimension, ScalarOperand};
9use scirs2_core::numeric::Float;
10use std::collections::HashMap;
11use std::fmt::Debug;
12
13#[derive(Debug)]
45pub struct GroupedAdam<A: Float + Send + Sync, D: Dimension> {
46 defaultlr: A,
48 default_beta1: A,
50 default_beta2: A,
52 default_weight_decay: A,
54 epsilon: A,
56 amsgrad: bool,
58 group_manager: GroupManager<A, D>,
60 step: usize,
62 group_steps: HashMap<usize, usize>,
64 implicit_group: Option<usize>,
69}
70
71impl<A: Float + ScalarOperand + Debug + Send + Sync, D: Dimension + Send + Sync> GroupedAdam<A, D> {
72 pub fn new(defaultlr: A) -> Self {
74 Self {
75 defaultlr,
76 default_beta1: A::from(0.9).expect("GroupedAdam: default beta1 (0.9) must fit in A"),
77 default_beta2: A::from(0.999)
78 .expect("GroupedAdam: default beta2 (0.999) must fit in A"),
79 default_weight_decay: A::zero(),
80 epsilon: A::from(1e-8).expect("GroupedAdam: default epsilon (1e-8) must fit in A"),
81 amsgrad: false,
82 group_manager: GroupManager::new(),
83 step: 0,
84 group_steps: HashMap::new(),
85 implicit_group: None,
86 }
87 }
88
89 pub fn group_step_count(&self, groupid: usize) -> usize {
91 self.group_steps.get(&groupid).copied().unwrap_or(0)
92 }
93
94 pub fn clear_groups(&mut self) {
96 self.group_manager = GroupManager::new();
97 self.group_steps.clear();
98 self.implicit_group = None;
99 self.step = 0;
100 }
101
102 pub fn with_beta1(mut self, beta1: A) -> Self {
104 self.default_beta1 = beta1;
105 self
106 }
107
108 pub fn with_beta2(mut self, beta2: A) -> Self {
110 self.default_beta2 = beta2;
111 self
112 }
113
114 pub fn with_weight_decay(mut self, weight_decay: A) -> Self {
116 self.default_weight_decay = weight_decay;
117 self
118 }
119
120 pub fn with_amsgrad(mut self) -> Self {
122 self.amsgrad = true;
123 self
124 }
125
126 fn init_group_state(&mut self, groupid: usize) -> Result<()> {
128 let group = self.group_manager.get_group_mut(groupid)?;
129
130 if group.state.is_empty() {
131 let mut m_t = Vec::new();
132 let mut v_t = Vec::new();
133 let mut v_hat_max = Vec::new();
134
135 for param in &group.params {
136 m_t.push(Array::zeros(param.raw_dim()));
137 v_t.push(Array::zeros(param.raw_dim()));
138 if self.amsgrad {
139 v_hat_max.push(Array::zeros(param.raw_dim()));
140 }
141 }
142
143 group.state.insert("m_t".to_string(), m_t);
144 group.state.insert("v_t".to_string(), v_t);
145 if self.amsgrad {
146 group.state.insert("v_hat_max".to_string(), v_hat_max);
147 }
148 }
149
150 Ok(())
151 }
152
153 fn step_group_internal(
158 &mut self,
159 groupid: usize,
160 group_step: usize,
161 gradients: &[Array<A, D>],
162 ) -> Result<Vec<Array<A, D>>> {
163 let t = i32::try_from(group_step).map_err(|_| {
164 OptimError::InvalidConfig(
165 "Timestep too large for bias correction calculation".to_string(),
166 )
167 })?;
168
169 self.init_group_state(groupid)?;
171
172 let group = self.group_manager.get_group_mut(groupid)?;
173
174 if gradients.len() != group.params.len() {
175 return Err(OptimError::InvalidConfig(format!(
176 "Number of gradients ({}) doesn't match number of parameters ({})",
177 gradients.len(),
178 group.params.len()
179 )));
180 }
181
182 let lr = group.learning_rate(self.defaultlr);
184 let beta1 = group.get_custom_param("beta1", self.default_beta1);
185 let beta2 = group.get_custom_param("beta2", self.default_beta2);
186 let weightdecay = group.weight_decay(self.default_weight_decay);
187
188 let mut updated_params = Vec::new();
189
190 for i in 0..group.params.len() {
192 let param = &group.params[i];
193 let grad = &gradients[i];
194
195 let grad_with_decay = if weightdecay > A::zero() {
197 grad + &(param * weightdecay)
198 } else {
199 grad.clone()
200 };
201
202 let updated = {
204 let m_t = group.state.get_mut("m_t").ok_or_else(|| {
206 OptimError::InvalidConfig("missing 'm_t' state for group".to_string())
207 })?;
208 m_t[i] = &m_t[i] * beta1 + &grad_with_decay * (A::one() - beta1);
209 let m_hat = &m_t[i] / (A::one() - beta1.powi(t));
210
211 let v_t = group.state.get_mut("v_t").ok_or_else(|| {
213 OptimError::InvalidConfig("missing 'v_t' state for group".to_string())
214 })?;
215 v_t[i] = &v_t[i] * beta2 + &grad_with_decay * &grad_with_decay * (A::one() - beta2);
216 let v_hat = &v_t[i] / (A::one() - beta2.powi(t));
217
218 if self.amsgrad {
220 let v_hat_max = group.state.get_mut("v_hat_max").ok_or_else(|| {
221 OptimError::InvalidConfig("missing 'v_hat_max' state for group".to_string())
222 })?;
223 v_hat_max[i].zip_mut_with(&v_hat, |a, &b| *a = a.max(b));
224 param - &(&m_hat * lr / (&v_hat_max[i].mapv(|x| x.sqrt()) + self.epsilon))
225 } else {
226 param - &(&m_hat * lr / (&v_hat.mapv(|x| x.sqrt()) + self.epsilon))
227 }
228 };
229
230 updated_params.push(updated);
231 }
232
233 group.params = updated_params.clone();
235
236 Ok(updated_params)
237 }
238}
239
240impl<A: Float + ScalarOperand + Debug + Send + Sync, D: Dimension + Send + Sync>
241 GroupedOptimizer<A, D> for GroupedAdam<A, D>
242{
243 fn add_group(
244 &mut self,
245 params: Vec<Array<A, D>>,
246 config: ParameterGroupConfig<A>,
247 ) -> Result<usize> {
248 Ok(self.group_manager.add_group(params, config))
249 }
250
251 fn get_group(&self, groupid: usize) -> Result<&ParameterGroup<A, D>> {
252 self.group_manager.get_group(groupid)
253 }
254
255 fn get_group_mut(&mut self, groupid: usize) -> Result<&mut ParameterGroup<A, D>> {
256 self.group_manager.get_group_mut(groupid)
257 }
258
259 fn groups(&self) -> &[ParameterGroup<A, D>] {
260 self.group_manager.groups()
261 }
262
263 fn groups_mut(&mut self) -> &mut [ParameterGroup<A, D>] {
264 self.group_manager.groups_mut()
265 }
266
267 fn step_group(
268 &mut self,
269 groupid: usize,
270 gradients: &[Array<A, D>],
271 ) -> Result<Vec<Array<A, D>>> {
272 self.step = self.step.saturating_add(1);
273 let group_step = {
274 let counter = self.group_steps.entry(groupid).or_insert(0);
275 *counter = counter.saturating_add(1);
276 *counter
277 };
278 self.step_group_internal(groupid, group_step, gradients)
279 }
280
281 fn set_group_learning_rate(&mut self, groupid: usize, lr: A) -> Result<()> {
282 let group = self.group_manager.get_group_mut(groupid)?;
283 group.config.learning_rate = Some(lr);
284 Ok(())
285 }
286
287 fn set_group_weight_decay(&mut self, groupid: usize, wd: A) -> Result<()> {
288 let group = self.group_manager.get_group_mut(groupid)?;
289 group.config.weight_decay = Some(wd);
290 Ok(())
291 }
292}
293
294impl<A: Float + ScalarOperand + Debug + Send + Sync, D: Dimension + Send + Sync> Optimizer<A, D>
296 for GroupedAdam<A, D>
297{
298 fn step(&mut self, params: &Array<A, D>, gradients: &Array<A, D>) -> Result<Array<A, D>> {
299 if params.shape() != gradients.shape() {
300 return Err(OptimError::DimensionMismatch(format!(
301 "Incompatible shapes: parameters have shape {:?}, gradients have shape {:?}",
302 params.shape(),
303 gradients.shape()
304 )));
305 }
306
307 let reusable = self
311 .implicit_group
312 .filter(|id| self.group_manager.get_group(*id).is_ok());
313
314 let groupid = if let Some(id) = reusable {
315 let group = self.group_manager.get_group_mut(id)?;
316 let shape_changed =
317 group.params.len() != 1 || group.params[0].raw_dim() != params.raw_dim();
318 if shape_changed {
319 group.params = vec![params.clone()];
321 group.state.clear();
322 self.group_steps.insert(id, 0);
323 } else {
324 group.params[0] = params.clone();
325 }
326 id
327 } else {
328 let id = self.add_group(vec![params.clone()], ParameterGroupConfig::new())?;
329 self.implicit_group = Some(id);
330 self.group_steps.insert(id, 0);
331 id
332 };
333
334 let result = self.step_group(groupid, std::slice::from_ref(gradients))?;
335
336 result.into_iter().next().ok_or_else(|| {
337 OptimError::InvalidConfig("grouped Adam step produced no parameters".to_string())
338 })
339 }
340
341 fn get_learning_rate(&self) -> A {
342 self.defaultlr
343 }
344
345 fn set_learning_rate(&mut self, learning_rate: A) {
346 self.defaultlr = learning_rate;
347 }
348}
349
350#[cfg(test)]
351mod tests {
352 use super::*;
353 use scirs2_core::ndarray::Array1;
354
355 #[test]
356 fn test_grouped_adam_creation() {
357 let optimizer: GroupedAdam<f64, scirs2_core::ndarray::Ix1> = GroupedAdam::new(0.001);
358 assert_eq!(optimizer.defaultlr, 0.001);
359 assert_eq!(optimizer.default_beta1, 0.9);
360 assert_eq!(optimizer.default_beta2, 0.999);
361 }
362
363 #[test]
364 fn test_grouped_adam_multiple_groups() {
365 let mut optimizer = GroupedAdam::new(0.001);
366
367 let params1 = vec![Array1::from_vec(vec![1.0, 2.0])];
369 let config1 = ParameterGroupConfig::new().with_learning_rate(0.01);
370 let group1 = optimizer
371 .add_group(params1, config1)
372 .expect("add_group succeeds in test_grouped_adam_multiple_groups");
373
374 let params2 = vec![Array1::from_vec(vec![3.0, 4.0, 5.0])];
376 let config2 = ParameterGroupConfig::new().with_learning_rate(0.0001);
377 let group2 = optimizer
378 .add_group(params2, config2)
379 .expect("add_group succeeds in test_grouped_adam_multiple_groups");
380
381 let grads1 = vec![Array1::from_vec(vec![0.1, 0.2])];
383 let updated1 = optimizer
384 .step_group(group1, &grads1)
385 .expect("step_group succeeds in test_grouped_adam_multiple_groups");
386
387 let grads2 = vec![Array1::from_vec(vec![0.3, 0.4, 0.5])];
389 let updated2 = optimizer
390 .step_group(group2, &grads2)
391 .expect("step_group succeeds in test_grouped_adam_multiple_groups");
392
393 assert!(updated1[0][0] < 1.0); assert!(updated2[0][0] > 2.9); }
397
398 #[test]
399 fn test_grouped_adam_custom_betas() {
400 let mut optimizer = GroupedAdam::new(0.001);
401
402 let params = vec![Array1::from_vec(vec![1.0, 2.0])];
404 let config = ParameterGroupConfig::new()
405 .with_custom_param("beta1".to_string(), 0.8)
406 .with_custom_param("beta2".to_string(), 0.99);
407 let group = optimizer
408 .add_group(params, config)
409 .expect("optimizer.add_group succeeds in test_grouped_adam_custom_betas");
410
411 let group_ref = optimizer
413 .get_group(group)
414 .expect("optimizer.get_group succeeds in test_grouped_adam_custom_betas");
415 assert_eq!(group_ref.get_custom_param("beta1", 0.0), 0.8);
416 assert_eq!(group_ref.get_custom_param("beta2", 0.0), 0.99);
417 }
418
419 #[test]
420 fn test_grouped_adam_clear() {
421 let mut optimizer = GroupedAdam::new(0.001);
422
423 let params1 = vec![Array1::zeros(2)];
425 let config1 = ParameterGroupConfig::new();
426 optimizer
427 .add_group(params1, config1)
428 .expect("add_group succeeds in test_grouped_adam_clear");
429
430 assert_eq!(optimizer.groups().len(), 1);
431
432 optimizer.clear_groups();
434
435 assert_eq!(optimizer.groups().len(), 0);
436 assert_eq!(optimizer.step, 0);
437 }
438
439 #[test]
443 fn test_grouped_adam_step_reuses_single_group() {
444 let mut optimizer: GroupedAdam<f64, scirs2_core::ndarray::Ix1> = GroupedAdam::new(0.1);
445
446 let mut params = Array1::from_vec(vec![0.0f64]);
447 let gradients = Array1::from_vec(vec![1.0f64]);
448
449 for _ in 0..100 {
450 params = optimizer
451 .step(¶ms, &gradients)
452 .expect("step should succeed");
453 }
454
455 assert_eq!(optimizer.groups().len(), 1);
457 assert_eq!(optimizer.group_step_count(0), 100);
458
459 let group = optimizer.get_group(0).expect("implicit group must exist");
461 let m_t = group
462 .state
463 .get("m_t")
464 .expect("first moment state must exist");
465 assert!(
467 (m_t[0][0] - 1.0).abs() < 1e-3,
468 "first moment did not accumulate: {}",
469 m_t[0][0]
470 );
471 }
472
473 #[test]
475 fn test_grouped_adam_first_step_uses_t_one() {
476 let mut optimizer: GroupedAdam<f64, scirs2_core::ndarray::Ix1> = GroupedAdam::new(0.1);
477
478 let params = Array1::from_vec(vec![0.0f64]);
479 let gradients = Array1::from_vec(vec![1.0f64]);
480
481 let updated = optimizer
482 .step(¶ms, &gradients)
483 .expect("step should succeed");
484
485 assert!(
487 (updated[0] + 0.1).abs() < 1e-9,
488 "expected -0.1, got {}",
489 updated[0]
490 );
491 }
492
493 #[test]
495 fn test_grouped_adam_per_group_step_counters() {
496 let mut optimizer = GroupedAdam::new(0.1);
497
498 let group_a = optimizer
499 .add_group(
500 vec![Array1::from_vec(vec![0.0f64])],
501 ParameterGroupConfig::new(),
502 )
503 .expect("add group a");
504 let group_b = optimizer
505 .add_group(
506 vec![Array1::from_vec(vec![0.0f64])],
507 ParameterGroupConfig::new(),
508 )
509 .expect("add group b");
510
511 let grads = vec![Array1::from_vec(vec![1.0f64])];
512
513 for _ in 0..5 {
515 optimizer.step_group(group_a, &grads).expect("step group a");
516 }
517
518 let first_b = optimizer.step_group(group_b, &grads).expect("step group b");
519
520 assert_eq!(optimizer.group_step_count(group_a), 5);
521 assert_eq!(optimizer.group_step_count(group_b), 1);
522
523 assert!(
525 (first_b[0][0] + 0.1).abs() < 1e-9,
526 "group B was contaminated by group A's clock: {}",
527 first_b[0][0]
528 );
529 }
530}