1use crate::{
9 Optimizer, OptimizerError, OptimizerResult, OptimizerState, ParamGroup, ParamGroupState,
10};
11use parking_lot::RwLock;
12use std::collections::HashMap;
13use std::sync::Arc;
14use torsh_core::error::TorshError;
15use torsh_tensor::Tensor;
16
17fn param_key(param: &Arc<RwLock<Tensor>>) -> String {
19 format!("param_{:p}", Arc::as_ptr(param))
20}
21
22fn deep_copy(tensor: &Tensor) -> OptimizerResult<Tensor> {
27 let data = tensor.to_vec().map_err(OptimizerError::TensorError)?;
28 Tensor::from_data(data, tensor.shape().dims().to_vec(), tensor.device())
29 .map_err(OptimizerError::TensorError)
30}
31
32pub struct AdvancedAdam {
34 pub lr: f64,
35 pub beta1: f64,
36 pub beta2: f64,
37 pub eps: f64,
38 pub weight_decay: f64,
39 pub amsgrad: bool,
40
41 pub param_groups: Vec<ParamGroup>,
43
44 pub state: HashMap<String, AdamState>,
46 pub step_count: u64,
47
48 pub adaptive_lr: bool,
50 pub gradient_clipping: Option<f64>,
51 pub warmup_steps: Option<u64>,
52}
53
54#[derive(Debug, Clone)]
55pub struct AdamState {
56 pub exp_avg: Tensor,
57 pub exp_avg_sq: Tensor,
58 pub max_exp_avg_sq: Option<Tensor>,
59}
60
61impl AdvancedAdam {
62 pub fn new(lr: f64) -> Self {
64 Self {
65 lr,
66 beta1: 0.9,
67 beta2: 0.999,
68 eps: 1e-8,
69 weight_decay: 0.0,
70 amsgrad: false,
71 param_groups: Vec::new(),
72 state: HashMap::new(),
73 step_count: 0,
74 adaptive_lr: false,
75 gradient_clipping: None,
76 warmup_steps: None,
77 }
78 }
79
80 pub fn with_params(lr: f64, params: Vec<Arc<RwLock<Tensor>>>) -> Self {
82 let mut optimizer = Self::new(lr);
83 optimizer
84 .param_groups
85 .push(ParamGroup::new(params, lr as f32));
86 optimizer
87 }
88
89 pub fn with_amsgrad(mut self) -> Self {
91 self.amsgrad = true;
92 self
93 }
94
95 pub fn with_weight_decay(mut self, weight_decay: f64) -> Self {
97 self.weight_decay = weight_decay;
98 self
99 }
100
101 pub fn with_adaptive_lr(mut self) -> Self {
107 self.adaptive_lr = true;
108 self
109 }
110
111 pub fn with_gradient_clipping(mut self, max_norm: f64) -> Self {
113 self.gradient_clipping = Some(max_norm);
114 self
115 }
116
117 pub fn with_warmup(mut self, warmup_steps: u64) -> Self {
119 self.warmup_steps = Some(warmup_steps);
120 self
121 }
122
123 fn schedule_scale(&self) -> f64 {
128 let step = self.step_count.max(1) as f64;
129 let warmup = self.warmup_steps.unwrap_or(0);
130
131 let warmup_scale = if warmup > 0 {
132 (step / warmup as f64).min(1.0)
133 } else {
134 1.0
135 };
136
137 let decay_scale = if self.adaptive_lr {
138 let reference = warmup.max(1) as f64;
139 if step > reference {
140 (reference / step).sqrt()
141 } else {
142 1.0
143 }
144 } else {
145 1.0
146 };
147
148 warmup_scale * decay_scale
149 }
150}
151
152impl Optimizer for AdvancedAdam {
153 fn step(&mut self) -> OptimizerResult<()> {
154 self.step_count += 1;
155 let step = self.step_count as i32;
156 let scale = self.schedule_scale();
157 let bias_correction1 = 1.0 - self.beta1.powi(step);
158 let bias_correction2 = 1.0 - self.beta2.powi(step);
159
160 let groups: Vec<(f32, Vec<Arc<RwLock<Tensor>>>)> = self
163 .param_groups
164 .iter()
165 .map(|group| (group.lr, group.params.clone()))
166 .collect();
167
168 for (group_lr, params) in groups {
169 let effective_lr = (group_lr as f64 * scale) as f32;
170
171 for param_arc in params {
172 let mut param = param_arc.write();
173 let Some(mut grad) = param.grad() else {
174 continue;
175 };
176
177 if let Some(max_norm) = self.gradient_clipping {
179 let norm =
180 grad.norm()
181 .map_err(OptimizerError::TensorError)?
182 .item()
183 .map_err(OptimizerError::TensorError)? as f64;
184 if norm > max_norm && norm > 0.0 {
185 grad = grad
186 .mul_scalar((max_norm / norm) as f32)
187 .map_err(OptimizerError::TensorError)?;
188 }
189 }
190
191 if self.weight_decay != 0.0 {
193 let decay = param
194 .mul_scalar(self.weight_decay as f32)
195 .map_err(OptimizerError::TensorError)?;
196 grad = grad.add(&decay).map_err(OptimizerError::TensorError)?;
197 }
198
199 let key = param_key(¶m_arc);
200 if !self.state.contains_key(&key) {
201 let zeros = torsh_tensor::creation::zeros_like(¶m)
202 .map_err(OptimizerError::TensorError)?;
203 self.state.insert(
204 key.clone(),
205 AdamState {
206 exp_avg: zeros.clone(),
207 exp_avg_sq: zeros.clone(),
208 max_exp_avg_sq: if self.amsgrad { Some(zeros) } else { None },
209 },
210 );
211 }
212 let state = self
213 .state
214 .get_mut(&key)
215 .expect("state was just inserted for this key");
216
217 let grad_term = grad
219 .mul_scalar(1.0 - self.beta1 as f32)
220 .map_err(OptimizerError::TensorError)?;
221 state
222 .exp_avg
223 .mul_scalar_(self.beta1 as f32)
224 .map_err(OptimizerError::TensorError)?;
225 state
226 .exp_avg
227 .add_(&grad_term)
228 .map_err(OptimizerError::TensorError)?;
229
230 let grad_sq = grad.mul_op(&grad).map_err(OptimizerError::TensorError)?;
232 let grad_sq_term = grad_sq
233 .mul_scalar(1.0 - self.beta2 as f32)
234 .map_err(OptimizerError::TensorError)?;
235 state
236 .exp_avg_sq
237 .mul_scalar_(self.beta2 as f32)
238 .map_err(OptimizerError::TensorError)?;
239 state
240 .exp_avg_sq
241 .add_(&grad_sq_term)
242 .map_err(OptimizerError::TensorError)?;
243
244 let corrected_exp_avg = state
245 .exp_avg
246 .div_scalar(bias_correction1 as f32)
247 .map_err(OptimizerError::TensorError)?;
248 let corrected_exp_avg_sq = state
249 .exp_avg_sq
250 .div_scalar(bias_correction2 as f32)
251 .map_err(OptimizerError::TensorError)?;
252
253 let denom_source = if let Some(max_exp_avg_sq) = state.max_exp_avg_sq.as_mut() {
255 let new_max = max_exp_avg_sq
256 .maximum(&corrected_exp_avg_sq)
257 .map_err(OptimizerError::TensorError)?;
258 *max_exp_avg_sq = new_max;
259 max_exp_avg_sq.clone()
260 } else {
261 corrected_exp_avg_sq
262 };
263
264 let denom = denom_source
265 .sqrt()
266 .map_err(OptimizerError::TensorError)?
267 .add_scalar(self.eps as f32)
268 .map_err(OptimizerError::TensorError)?;
269
270 let update = corrected_exp_avg
271 .div(&denom)
272 .map_err(OptimizerError::TensorError)?
273 .mul_scalar(effective_lr)
274 .map_err(OptimizerError::TensorError)?;
275
276 crate::param_update::sub_assign(&mut param, &update)
277 .map_err(OptimizerError::TensorError)?;
278 }
279 }
280
281 Ok(())
282 }
283
284 fn zero_grad(&mut self) {
285 for group in &self.param_groups {
286 for param in &group.params {
287 param.write().zero_grad();
288 }
289 }
290 }
291
292 fn get_lr(&self) -> Vec<f32> {
293 if self.param_groups.is_empty() {
294 vec![self.lr as f32]
295 } else {
296 self.param_groups.iter().map(|group| group.lr).collect()
297 }
298 }
299
300 fn set_lr(&mut self, lr: f32) {
301 self.lr = lr as f64;
302 for group in &mut self.param_groups {
303 group.lr = lr;
304 }
305 }
306
307 fn set_lrs(&mut self, lrs: &[f32]) {
308 if let Some(&lr) = lrs.first() {
309 self.lr = lr as f64;
310 }
311 for (group, &lr) in self.param_groups.iter_mut().zip(lrs.iter()) {
312 group.lr = lr;
313 }
314 }
315
316 fn add_param_group(&mut self, params: Vec<Arc<RwLock<Tensor>>>, options: HashMap<String, f32>) {
317 let mut options = options;
318 let lr = options.remove("lr").unwrap_or(self.lr as f32);
319 let mut group = ParamGroup::new(params, lr);
320 group.options = options;
321 self.param_groups.push(group);
322 }
323
324 fn parameters(&self) -> Vec<Arc<RwLock<Tensor>>> {
325 self.param_groups
326 .iter()
327 .flat_map(|group| group.params.iter().cloned())
328 .collect()
329 }
330
331 fn state_dict(&self) -> OptimizerResult<OptimizerState> {
332 let mut global_state = HashMap::new();
333 global_state.insert("lr".to_string(), self.lr as f32);
334 global_state.insert("beta1".to_string(), self.beta1 as f32);
335 global_state.insert("beta2".to_string(), self.beta2 as f32);
336 global_state.insert("eps".to_string(), self.eps as f32);
337 global_state.insert("weight_decay".to_string(), self.weight_decay as f32);
338 global_state.insert("step_count".to_string(), self.step_count as f32);
339 global_state.insert("amsgrad".to_string(), if self.amsgrad { 1.0 } else { 0.0 });
340
341 let param_groups = self
342 .param_groups
343 .iter()
344 .map(|group| ParamGroupState {
345 lr: group.lr,
346 options: group.options.clone(),
347 param_count: group.params.len(),
348 })
349 .collect();
350
351 let mut state = HashMap::new();
353 for (key, adam_state) in &self.state {
354 let mut entry = HashMap::new();
355 entry.insert("exp_avg".to_string(), adam_state.exp_avg.clone());
356 entry.insert("exp_avg_sq".to_string(), adam_state.exp_avg_sq.clone());
357 if let Some(max_exp_avg_sq) = &adam_state.max_exp_avg_sq {
358 entry.insert("max_exp_avg_sq".to_string(), max_exp_avg_sq.clone());
359 }
360 state.insert(key.clone(), entry);
361 }
362
363 Ok(OptimizerState {
364 optimizer_type: "AdvancedAdam".to_string(),
365 version: "1.0".to_string(),
366 param_groups,
367 state,
368 global_state,
369 })
370 }
371
372 fn load_state_dict(&mut self, state: OptimizerState) -> OptimizerResult<()> {
373 if state.optimizer_type != "AdvancedAdam" {
374 return Err(OptimizerError::InvalidParameter(format!(
375 "Expected AdvancedAdam, got {}",
376 state.optimizer_type
377 )));
378 }
379
380 if let Some(lr) = state.global_state.get("lr") {
381 self.lr = *lr as f64;
382 }
383 if let Some(beta1) = state.global_state.get("beta1") {
384 self.beta1 = *beta1 as f64;
385 }
386 if let Some(beta2) = state.global_state.get("beta2") {
387 self.beta2 = *beta2 as f64;
388 }
389 if let Some(eps) = state.global_state.get("eps") {
390 self.eps = *eps as f64;
391 }
392 if let Some(weight_decay) = state.global_state.get("weight_decay") {
393 self.weight_decay = *weight_decay as f64;
394 }
395 if let Some(step_count) = state.global_state.get("step_count") {
396 self.step_count = *step_count as u64;
397 }
398 if let Some(amsgrad) = state.global_state.get("amsgrad") {
399 self.amsgrad = *amsgrad != 0.0;
400 }
401
402 for (group, saved) in self.param_groups.iter_mut().zip(state.param_groups.iter()) {
404 group.lr = saved.lr;
405 group.options = saved.options.clone();
406 }
407
408 self.state.clear();
410 for (key, entry) in state.state {
411 let exp_avg = entry.get("exp_avg").ok_or_else(|| {
412 OptimizerError::StateError(format!("AdvancedAdam state for {key} has no exp_avg"))
413 })?;
414 let exp_avg_sq = entry.get("exp_avg_sq").ok_or_else(|| {
415 OptimizerError::StateError(format!(
416 "AdvancedAdam state for {key} has no exp_avg_sq"
417 ))
418 })?;
419 let max_exp_avg_sq = entry.get("max_exp_avg_sq").map(deep_copy).transpose()?;
420 self.state.insert(
421 key,
422 AdamState {
423 exp_avg: deep_copy(exp_avg)?,
424 exp_avg_sq: deep_copy(exp_avg_sq)?,
425 max_exp_avg_sq,
426 },
427 );
428 }
429
430 Ok(())
431 }
432}
433
434pub struct LAMB {
437 pub lr: f64,
438 pub beta1: f64,
439 pub beta2: f64,
440 pub eps: f64,
441 pub weight_decay: f64,
442 pub bias_correction: bool,
443
444 pub param_groups: Vec<ParamGroup>,
446
447 pub state: HashMap<String, LambState>,
448 pub step_count: u64,
449}
450
451#[derive(Debug, Clone)]
452pub struct LambState {
453 pub exp_avg: Tensor,
454 pub exp_avg_sq: Tensor,
455}
456
457impl LAMB {
458 pub fn new(lr: f64) -> Self {
460 Self {
461 lr,
462 beta1: 0.9,
463 beta2: 0.999,
464 eps: 1e-6,
465 weight_decay: 0.01,
466 bias_correction: true,
467 param_groups: Vec::new(),
468 state: HashMap::new(),
469 step_count: 0,
470 }
471 }
472
473 pub fn with_params(lr: f64, params: Vec<Arc<RwLock<Tensor>>>) -> Self {
475 let mut optimizer = Self::new(lr);
476 optimizer
477 .param_groups
478 .push(ParamGroup::new(params, lr as f32));
479 optimizer
480 }
481}
482
483impl Optimizer for LAMB {
484 fn step(&mut self) -> OptimizerResult<()> {
485 self.step_count += 1;
486 let step = self.step_count as i32;
487 let (bias_correction1, bias_correction2) = if self.bias_correction {
488 (1.0 - self.beta1.powi(step), 1.0 - self.beta2.powi(step))
489 } else {
490 (1.0, 1.0)
491 };
492
493 let groups: Vec<(f32, Vec<Arc<RwLock<Tensor>>>)> = self
494 .param_groups
495 .iter()
496 .map(|group| (group.lr, group.params.clone()))
497 .collect();
498
499 for (group_lr, params) in groups {
500 for param_arc in params {
501 let mut param = param_arc.write();
502 let Some(grad) = param.grad() else {
503 continue;
504 };
505
506 let key = param_key(¶m_arc);
507 if !self.state.contains_key(&key) {
508 let zeros = torsh_tensor::creation::zeros_like(¶m)
509 .map_err(OptimizerError::TensorError)?;
510 self.state.insert(
511 key.clone(),
512 LambState {
513 exp_avg: zeros.clone(),
514 exp_avg_sq: zeros,
515 },
516 );
517 }
518 let state = self
519 .state
520 .get_mut(&key)
521 .expect("state was just inserted for this key");
522
523 let grad_term = grad
524 .mul_scalar(1.0 - self.beta1 as f32)
525 .map_err(OptimizerError::TensorError)?;
526 state
527 .exp_avg
528 .mul_scalar_(self.beta1 as f32)
529 .map_err(OptimizerError::TensorError)?;
530 state
531 .exp_avg
532 .add_(&grad_term)
533 .map_err(OptimizerError::TensorError)?;
534
535 let grad_sq = grad.mul_op(&grad).map_err(OptimizerError::TensorError)?;
536 let grad_sq_term = grad_sq
537 .mul_scalar(1.0 - self.beta2 as f32)
538 .map_err(OptimizerError::TensorError)?;
539 state
540 .exp_avg_sq
541 .mul_scalar_(self.beta2 as f32)
542 .map_err(OptimizerError::TensorError)?;
543 state
544 .exp_avg_sq
545 .add_(&grad_sq_term)
546 .map_err(OptimizerError::TensorError)?;
547
548 let corrected_exp_avg = state
549 .exp_avg
550 .div_scalar(bias_correction1 as f32)
551 .map_err(OptimizerError::TensorError)?;
552 let corrected_exp_avg_sq = state
553 .exp_avg_sq
554 .div_scalar(bias_correction2 as f32)
555 .map_err(OptimizerError::TensorError)?;
556
557 let denom = corrected_exp_avg_sq
558 .sqrt()
559 .map_err(OptimizerError::TensorError)?
560 .add_scalar(self.eps as f32)
561 .map_err(OptimizerError::TensorError)?;
562
563 let mut direction = corrected_exp_avg
565 .div(&denom)
566 .map_err(OptimizerError::TensorError)?;
567 if self.weight_decay != 0.0 {
568 let decay = param
569 .mul_scalar(self.weight_decay as f32)
570 .map_err(OptimizerError::TensorError)?;
571 direction = direction.add(&decay).map_err(OptimizerError::TensorError)?;
572 }
573
574 let param_norm = param
577 .norm()
578 .map_err(OptimizerError::TensorError)?
579 .item()
580 .map_err(OptimizerError::TensorError)?;
581 let direction_norm = direction
582 .norm()
583 .map_err(OptimizerError::TensorError)?
584 .item()
585 .map_err(OptimizerError::TensorError)?;
586 let trust_ratio = if param_norm > 0.0 && direction_norm > 0.0 {
587 param_norm / direction_norm
588 } else {
589 1.0
590 };
591
592 let update = direction
593 .mul_scalar(group_lr * trust_ratio)
594 .map_err(OptimizerError::TensorError)?;
595 crate::param_update::sub_assign(&mut param, &update)
596 .map_err(OptimizerError::TensorError)?;
597 }
598 }
599
600 Ok(())
601 }
602
603 fn zero_grad(&mut self) {
604 for group in &self.param_groups {
605 for param in &group.params {
606 param.write().zero_grad();
607 }
608 }
609 }
610
611 fn get_lr(&self) -> Vec<f32> {
612 if self.param_groups.is_empty() {
613 vec![self.lr as f32]
614 } else {
615 self.param_groups.iter().map(|group| group.lr).collect()
616 }
617 }
618
619 fn set_lr(&mut self, lr: f32) {
620 self.lr = lr as f64;
621 for group in &mut self.param_groups {
622 group.lr = lr;
623 }
624 }
625
626 fn set_lrs(&mut self, lrs: &[f32]) {
627 if let Some(&lr) = lrs.first() {
628 self.lr = lr as f64;
629 }
630 for (group, &lr) in self.param_groups.iter_mut().zip(lrs.iter()) {
631 group.lr = lr;
632 }
633 }
634
635 fn add_param_group(&mut self, params: Vec<Arc<RwLock<Tensor>>>, options: HashMap<String, f32>) {
636 let mut options = options;
637 let lr = options.remove("lr").unwrap_or(self.lr as f32);
638 let mut group = ParamGroup::new(params, lr);
639 group.options = options;
640 self.param_groups.push(group);
641 }
642
643 fn parameters(&self) -> Vec<Arc<RwLock<Tensor>>> {
644 self.param_groups
645 .iter()
646 .flat_map(|group| group.params.iter().cloned())
647 .collect()
648 }
649
650 fn state_dict(&self) -> OptimizerResult<OptimizerState> {
651 let mut global_state = HashMap::new();
652 global_state.insert("lr".to_string(), self.lr as f32);
653 global_state.insert("beta1".to_string(), self.beta1 as f32);
654 global_state.insert("beta2".to_string(), self.beta2 as f32);
655 global_state.insert("eps".to_string(), self.eps as f32);
656 global_state.insert("weight_decay".to_string(), self.weight_decay as f32);
657 global_state.insert("step_count".to_string(), self.step_count as f32);
658 global_state.insert(
659 "bias_correction".to_string(),
660 if self.bias_correction { 1.0 } else { 0.0 },
661 );
662
663 let param_groups = self
664 .param_groups
665 .iter()
666 .map(|group| ParamGroupState {
667 lr: group.lr,
668 options: group.options.clone(),
669 param_count: group.params.len(),
670 })
671 .collect();
672
673 let mut state = HashMap::new();
674 for (key, lamb_state) in &self.state {
675 let mut entry = HashMap::new();
676 entry.insert("exp_avg".to_string(), lamb_state.exp_avg.clone());
677 entry.insert("exp_avg_sq".to_string(), lamb_state.exp_avg_sq.clone());
678 state.insert(key.clone(), entry);
679 }
680
681 Ok(OptimizerState {
682 optimizer_type: "LAMB".to_string(),
683 version: "1.0".to_string(),
684 param_groups,
685 state,
686 global_state,
687 })
688 }
689
690 fn load_state_dict(&mut self, state: OptimizerState) -> OptimizerResult<()> {
691 if state.optimizer_type != "LAMB" {
692 return Err(OptimizerError::InvalidParameter(format!(
693 "Expected LAMB, got {}",
694 state.optimizer_type
695 )));
696 }
697
698 if let Some(lr) = state.global_state.get("lr") {
699 self.lr = *lr as f64;
700 }
701 if let Some(beta1) = state.global_state.get("beta1") {
702 self.beta1 = *beta1 as f64;
703 }
704 if let Some(beta2) = state.global_state.get("beta2") {
705 self.beta2 = *beta2 as f64;
706 }
707 if let Some(eps) = state.global_state.get("eps") {
708 self.eps = *eps as f64;
709 }
710 if let Some(weight_decay) = state.global_state.get("weight_decay") {
711 self.weight_decay = *weight_decay as f64;
712 }
713 if let Some(step_count) = state.global_state.get("step_count") {
714 self.step_count = *step_count as u64;
715 }
716 if let Some(bias_correction) = state.global_state.get("bias_correction") {
717 self.bias_correction = *bias_correction != 0.0;
718 }
719
720 for (group, saved) in self.param_groups.iter_mut().zip(state.param_groups.iter()) {
721 group.lr = saved.lr;
722 group.options = saved.options.clone();
723 }
724
725 self.state.clear();
726 for (key, entry) in state.state {
727 let exp_avg = entry.get("exp_avg").ok_or_else(|| {
728 OptimizerError::StateError(format!("LAMB state for {key} has no exp_avg"))
729 })?;
730 let exp_avg_sq = entry.get("exp_avg_sq").ok_or_else(|| {
731 OptimizerError::StateError(format!("LAMB state for {key} has no exp_avg_sq"))
732 })?;
733 self.state.insert(
734 key,
735 LambState {
736 exp_avg: deep_copy(exp_avg)?,
737 exp_avg_sq: deep_copy(exp_avg_sq)?,
738 },
739 );
740 }
741
742 Ok(())
743 }
744}
745
746pub struct Lookahead<T: Optimizer> {
753 pub base_optimizer: T,
754 pub alpha: f64,
755 pub k: u64,
756
757 pub slow_weights: HashMap<String, Tensor>,
758 pub step_count: u64,
759}
760
761impl<T: Optimizer> Lookahead<T> {
762 pub fn new(base_optimizer: T, alpha: f64, k: u64) -> Self {
764 Self {
765 base_optimizer,
766 alpha,
767 k,
768 slow_weights: HashMap::new(),
769 step_count: 0,
770 }
771 }
772
773 fn initialize_slow_weights(&mut self) -> OptimizerResult<()> {
776 for param in self.base_optimizer.parameters() {
777 let key = param_key(¶m);
778 if !self.slow_weights.contains_key(&key) {
779 let snapshot = deep_copy(¶m.read())?;
780 self.slow_weights.insert(key, snapshot);
781 }
782 }
783 Ok(())
784 }
785
786 fn synchronize(&mut self) -> OptimizerResult<()> {
788 for param in self.base_optimizer.parameters() {
789 let key = param_key(¶m);
790 let Some(slow) = self.slow_weights.get_mut(&key) else {
791 continue;
792 };
793
794 let mut fast = param.write();
795 let diff = fast
797 .detach()
798 .sub(slow)
799 .map_err(OptimizerError::TensorError)?
800 .mul_scalar(self.alpha as f32)
801 .map_err(OptimizerError::TensorError)?;
802 slow.add_(&diff).map_err(OptimizerError::TensorError)?;
803
804 crate::param_update::assign(&mut fast, slow).map_err(OptimizerError::TensorError)?;
806 }
807 Ok(())
808 }
809}
810
811impl<T: Optimizer> Optimizer for Lookahead<T> {
812 fn step(&mut self) -> OptimizerResult<()> {
813 self.initialize_slow_weights()?;
814
815 self.base_optimizer.step()?;
817 self.step_count += 1;
818
819 if self.k > 0 && self.step_count % self.k == 0 {
820 self.synchronize()?;
821 }
822
823 Ok(())
824 }
825
826 fn zero_grad(&mut self) {
827 self.base_optimizer.zero_grad();
828 }
829
830 fn get_lr(&self) -> Vec<f32> {
831 self.base_optimizer.get_lr()
832 }
833
834 fn set_lr(&mut self, lr: f32) {
835 self.base_optimizer.set_lr(lr);
836 }
837
838 fn set_lrs(&mut self, lrs: &[f32]) {
839 self.base_optimizer.set_lrs(lrs);
840 }
841
842 fn add_param_group(&mut self, params: Vec<Arc<RwLock<Tensor>>>, options: HashMap<String, f32>) {
843 self.base_optimizer.add_param_group(params, options);
844 }
845
846 fn parameters(&self) -> Vec<Arc<RwLock<Tensor>>> {
847 self.base_optimizer.parameters()
848 }
849
850 fn state_dict(&self) -> OptimizerResult<OptimizerState> {
851 let mut base_state = self.base_optimizer.state_dict()?;
852
853 base_state
855 .global_state
856 .insert("alpha".to_string(), self.alpha as f32);
857 base_state
858 .global_state
859 .insert("k".to_string(), self.k as f32);
860 base_state
861 .global_state
862 .insert("step_count".to_string(), self.step_count as f32);
863
864 for (key, slow) in &self.slow_weights {
867 base_state
868 .state
869 .entry(key.clone())
870 .or_default()
871 .insert("lookahead_slow_weight".to_string(), slow.clone());
872 }
873
874 base_state.optimizer_type = format!("Lookahead<{}>", base_state.optimizer_type);
875
876 Ok(base_state)
877 }
878
879 fn load_state_dict(&mut self, mut state: OptimizerState) -> OptimizerResult<()> {
880 if let Some(alpha) = state.global_state.remove("alpha") {
882 self.alpha = alpha as f64;
883 }
884 if let Some(k) = state.global_state.remove("k") {
885 self.k = k as u64;
886 }
887 if let Some(step_count) = state.global_state.remove("step_count") {
888 self.step_count = step_count as u64;
889 }
890
891 self.slow_weights.clear();
902 for param in self.base_optimizer.parameters() {
903 let key = param_key(¶m);
904 if let Some(entry) = state.state.get_mut(&key) {
905 if let Some(slow) = entry.remove("lookahead_slow_weight") {
906 self.slow_weights.insert(key, slow);
907 }
908 }
909 }
910 for entry in state.state.values_mut() {
913 entry.remove("lookahead_slow_weight");
914 }
915
916 if state.optimizer_type.starts_with("Lookahead<") && state.optimizer_type.ends_with(">") {
918 let base_type = &state.optimizer_type[10..state.optimizer_type.len() - 1];
919 state.optimizer_type = base_type.to_string();
920 }
921
922 self.base_optimizer.load_state_dict(state)
924 }
925}
926
927#[cfg(test)]
928mod tests {
929 use super::*;
930 use torsh_core::device::DeviceType;
931
932 fn make_param(data: Vec<f32>) -> Arc<RwLock<Tensor>> {
933 let len = data.len();
934 let tensor = Tensor::from_data(data, vec![len], DeviceType::Cpu)
935 .expect("parameter creation")
936 .requires_grad_(true);
937 Arc::new(RwLock::new(tensor))
938 }
939
940 fn set_grad(param: &Arc<RwLock<Tensor>>, data: Vec<f32>) {
941 let len = data.len();
942 let grad = Tensor::from_data(data, vec![len], DeviceType::Cpu).expect("gradient creation");
943 param.read().set_grad(Some(grad));
944 }
945
946 #[test]
947 fn test_advanced_adam() {
948 let mut optimizer = AdvancedAdam::new(0.001)
949 .with_amsgrad()
950 .with_weight_decay(0.01)
951 .with_gradient_clipping(1.0);
952
953 assert_eq!(optimizer.get_lr(), vec![0.001]);
955 assert!(optimizer.step().is_ok());
956 }
957
958 #[test]
959 fn test_advanced_adam_updates_parameters() {
960 let param = make_param(vec![1.0, 1.0]);
961 set_grad(¶m, vec![1.0, 1.0]);
962
963 let mut optimizer = AdvancedAdam::with_params(0.1, vec![Arc::clone(¶m)]);
964 optimizer.step().expect("step");
965 let after_first = param.read().to_vec().expect("to_vec");
966 assert!(
967 after_first[0] < 1.0,
968 "parameter must move, got {after_first:?}"
969 );
970
971 optimizer.step().expect("step");
972 let after_second = param.read().to_vec().expect("to_vec");
973 assert!(
974 after_second[0] < after_first[0],
975 "second step must move further: {after_first:?} -> {after_second:?}"
976 );
977 }
978
979 #[test]
980 fn test_advanced_adam_add_param_group() {
981 let first = make_param(vec![1.0]);
982 let second = make_param(vec![1.0]);
983 set_grad(&first, vec![1.0]);
984 set_grad(&second, vec![1.0]);
985
986 let mut optimizer = AdvancedAdam::with_params(0.1, vec![Arc::clone(&first)]);
987 let mut options = HashMap::new();
988 options.insert("lr".to_string(), 0.5);
989 optimizer.add_param_group(vec![Arc::clone(&second)], options);
990
991 assert_eq!(optimizer.parameters().len(), 2);
992 assert_eq!(optimizer.get_lr(), vec![0.1, 0.5]);
993
994 optimizer.step().expect("step");
995 assert!(first.read().to_vec().expect("to_vec")[0] < 1.0);
996 assert!(second.read().to_vec().expect("to_vec")[0] < 1.0);
997 }
998
999 #[test]
1000 fn test_advanced_adam_state_dict_round_trip() {
1001 let param = make_param(vec![1.0, 2.0]);
1002 set_grad(¶m, vec![0.5, 0.5]);
1003
1004 let mut optimizer = AdvancedAdam::with_params(0.1, vec![Arc::clone(¶m)]);
1005 optimizer.step().expect("step");
1006
1007 let dict = optimizer.state_dict().expect("state_dict");
1008 assert_eq!(dict.param_groups.len(), 1);
1009 assert_eq!(dict.param_groups[0].param_count, 1);
1010 assert_eq!(dict.state.len(), 1);
1011 assert!(dict
1012 .state
1013 .values()
1014 .next()
1015 .expect("state entry")
1016 .contains_key("exp_avg"));
1017
1018 let mut restored = AdvancedAdam::with_params(0.0, vec![Arc::clone(¶m)]);
1019 restored.load_state_dict(dict).expect("load_state_dict");
1020 assert_eq!(restored.step_count, 1);
1021 assert_eq!(restored.state.len(), 1);
1022 }
1023
1024 #[test]
1025 fn test_lamb_optimizer() {
1026 let mut optimizer = LAMB::new(0.001);
1027
1028 assert_eq!(optimizer.get_lr(), vec![0.001]);
1030 assert!(optimizer.step().is_ok());
1031 }
1032
1033 #[test]
1034 fn test_lamb_updates_parameters_with_trust_ratio() {
1035 let param = make_param(vec![1.0, 1.0]);
1036 set_grad(¶m, vec![1.0, 1.0]);
1037
1038 let mut optimizer = LAMB::with_params(0.01, vec![Arc::clone(¶m)]);
1039 optimizer.step().expect("step");
1040
1041 let after = param.read().to_vec().expect("to_vec");
1042 assert!(
1043 after[0] < 1.0,
1044 "LAMB must move the parameter, got {after:?}"
1045 );
1046 assert!(after[0] > 0.0, "trust ratio must keep the step bounded");
1047 }
1048
1049 #[test]
1050 fn test_lookahead_wrapper() {
1051 let base_optimizer = AdvancedAdam::new(0.001);
1052 let mut lookahead = Lookahead::new(base_optimizer, 0.5, 5);
1053
1054 assert_eq!(lookahead.get_lr(), vec![0.001]);
1056 assert!(lookahead.step().is_ok());
1057 }
1058
1059 #[test]
1060 fn test_lookahead_state_dict_round_trip_keeps_slow_weights_per_parameter() {
1061 let first = make_param(vec![10.0]);
1062 let second = make_param(vec![-10.0]);
1063 set_grad(&first, vec![1.0]);
1064 set_grad(&second, vec![1.0]);
1065
1066 let base = AdvancedAdam::with_params(0.1, vec![Arc::clone(&first), Arc::clone(&second)]);
1067 let mut lookahead = Lookahead::new(base, 0.5, 1);
1068 lookahead.step().expect("step");
1069
1070 let expected: HashMap<String, f32> = lookahead
1071 .slow_weights
1072 .iter()
1073 .map(|(key, tensor)| (key.clone(), tensor.to_vec().expect("to_vec")[0]))
1074 .collect();
1075 assert_eq!(expected.len(), 2);
1076
1077 let dict = lookahead.state_dict().expect("state_dict");
1078 let restored_base =
1079 AdvancedAdam::with_params(0.1, vec![Arc::clone(&first), Arc::clone(&second)]);
1080 let mut restored = Lookahead::new(restored_base, 0.0, 1);
1081 restored.load_state_dict(dict).expect("load_state_dict");
1082
1083 assert_eq!(restored.slow_weights.len(), 2);
1084 for (key, value) in expected {
1085 let got = restored
1086 .slow_weights
1087 .get(&key)
1088 .unwrap_or_else(|| panic!("slow weight for {key} must be restored"))
1089 .to_vec()
1090 .expect("to_vec")[0];
1091 assert!(
1092 (got - value).abs() < 1e-6,
1093 "slow weight for {key} must land on the same parameter: {value} vs {got}"
1094 );
1095 }
1096 }
1097
1098 #[test]
1099 fn test_lookahead_updates_slow_weights_every_k_steps() {
1100 let param = make_param(vec![0.0]);
1101 set_grad(¶m, vec![1.0]);
1102
1103 let base = AdvancedAdam::with_params(0.1, vec![Arc::clone(¶m)]);
1104 let mut lookahead = Lookahead::new(base, 0.5, 2);
1105
1106 lookahead.step().expect("step 1");
1107 let fast_after_one = param.read().to_vec().expect("to_vec")[0];
1108 let slow = lookahead
1110 .slow_weights
1111 .values()
1112 .next()
1113 .expect("slow weight")
1114 .to_vec()
1115 .expect("to_vec")[0];
1116 assert!((slow - 0.0).abs() < 1e-6, "slow weight must not move yet");
1117
1118 lookahead.step().expect("step 2");
1119 let slow = lookahead
1120 .slow_weights
1121 .values()
1122 .next()
1123 .expect("slow weight")
1124 .to_vec()
1125 .expect("to_vec")[0];
1126 let fast = param.read().to_vec().expect("to_vec")[0];
1127 assert!(
1128 slow < 0.0,
1129 "slow weight must be pulled towards the fast weights at step k"
1130 );
1131 assert!(
1132 (fast - slow).abs() < 1e-6,
1133 "fast weights must be reset onto the slow weights: {fast} vs {slow}"
1134 );
1135 assert!(
1139 slow > 2.0 * fast_after_one && slow < 0.0,
1140 "the interpolated slow weight must lag the fast trajectory, got {slow}"
1141 );
1142 }
1143}