Skip to main content

torsh_optim/
advanced.rs

1//! Advanced optimizers using SciRS2 optimization algorithms
2//!
3//! The optimizers here own their parameter groups, so `step()` performs a real
4//! update, `add_param_group` extends the set of optimised tensors, and
5//! `state_dict` / `load_state_dict` round-trip both the parameter groups and the
6//! per-parameter moment buffers.
7
8use 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
17/// Stable per-parameter key derived from the handle's identity.
18fn param_key(param: &Arc<RwLock<Tensor>>) -> String {
19    format!("param_{:p}", Arc::as_ptr(param))
20}
21
22/// Deep-copy a tensor's data into a fresh, gradient-free tensor.
23///
24/// Used for snapshots (slow weights, restored state) that must not observe later
25/// in-place updates to the source tensor.
26fn 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
32/// Advanced Adam optimizer with SciRS2 enhancements
33pub 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    /// Parameter groups optimised by this instance
42    pub param_groups: Vec<ParamGroup>,
43
44    // State variables
45    pub state: HashMap<String, AdamState>,
46    pub step_count: u64,
47
48    // SciRS2 enhancements
49    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    /// Create a new advanced Adam optimizer
63    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    /// Create an optimizer that already owns `params`
81    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    /// Enable AMSGrad variant
90    pub fn with_amsgrad(mut self) -> Self {
91        self.amsgrad = true;
92        self
93    }
94
95    /// Add weight decay (L2 regularization)
96    pub fn with_weight_decay(mut self, weight_decay: f64) -> Self {
97        self.weight_decay = weight_decay;
98        self
99    }
100
101    /// Enable the adaptive (inverse square-root) learning rate schedule
102    ///
103    /// With this enabled the learning rate decays as `sqrt(t_ref / t)` once the
104    /// step count passes `t_ref` (the warmup length, or 1 when no warmup is
105    /// configured) — the "Noam" schedule used for transformer training.
106    pub fn with_adaptive_lr(mut self) -> Self {
107        self.adaptive_lr = true;
108        self
109    }
110
111    /// Add gradient clipping
112    pub fn with_gradient_clipping(mut self, max_norm: f64) -> Self {
113        self.gradient_clipping = Some(max_norm);
114        self
115    }
116
117    /// Add learning rate warmup
118    pub fn with_warmup(mut self, warmup_steps: u64) -> Self {
119        self.warmup_steps = Some(warmup_steps);
120        self
121    }
122
123    /// Multiplier applied to every group's learning rate at the current step.
124    ///
125    /// Combines linear warmup (`step / warmup_steps`, capped at 1) with the
126    /// optional inverse square-root decay.
127    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        // Snapshot the (Arc) handles so the per-parameter state map can be
161        // borrowed mutably while iterating.
162        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                // Gradient clipping by global norm of this parameter's gradient.
178                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                // L2 regularisation folded into the gradient.
192                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(&param_arc);
200                if !self.state.contains_key(&key) {
201                    let zeros = torsh_tensor::creation::zeros_like(&param)
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                // m_t = b1 * m_{t-1} + (1 - b1) * g
218                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                // v_t = b2 * v_{t-1} + (1 - b2) * g^2
231                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                // AMSGrad maxes over the bias-corrected second moments.
254                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        // Per-parameter moment buffers, keyed the same way `step` keys them.
352        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        // Restore per-group learning rates for the groups this optimizer owns.
403        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        // Restore per-parameter moment buffers.
409        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
434/// LAMB (Layer-wise Adaptive Moments optimizer for Batch training)
435/// Particularly effective for large batch training
436pub 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    /// Parameter groups optimised by this instance
445    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    /// Create a new LAMB optimizer
459    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    /// Create an optimizer that already owns `params`
474    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(&param_arc);
507                if !self.state.contains_key(&key) {
508                    let zeros = torsh_tensor::creation::zeros_like(&param)
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                // Adam direction plus decoupled weight decay.
564                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                // Layer-wise trust ratio: ||w|| / ||r||, defaulting to 1 when
575                // either norm vanishes (as in the LAMB paper).
576                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
746/// Lookahead optimizer wrapper
747///
748/// Wraps any optimizer that exposes its parameters via
749/// [`Optimizer::parameters`]. Every `k` fast-weight steps the slow weights are
750/// pulled a fraction `alpha` towards the fast weights and the fast weights are
751/// reset onto them: `phi <- phi + alpha * (theta - phi)`, `theta <- phi`.
752pub 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    /// Create a new Lookahead optimizer
763    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    /// Snapshot the current fast weights as slow weights, for parameters that do
774    /// not have a snapshot yet.
775    fn initialize_slow_weights(&mut self) -> OptimizerResult<()> {
776        for param in self.base_optimizer.parameters() {
777            let key = param_key(&param);
778            if !self.slow_weights.contains_key(&key) {
779                let snapshot = deep_copy(&param.read())?;
780                self.slow_weights.insert(key, snapshot);
781            }
782        }
783        Ok(())
784    }
785
786    /// Perform the slow/fast synchronisation.
787    fn synchronize(&mut self) -> OptimizerResult<()> {
788        for param in self.base_optimizer.parameters() {
789            let key = param_key(&param);
790            let Some(slow) = self.slow_weights.get_mut(&key) else {
791                continue;
792            };
793
794            let mut fast = param.write();
795            // phi <- phi + alpha * (theta - phi)
796            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            // theta <- phi
805            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        // Perform base optimizer step
816        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        // Add Lookahead-specific state
854        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        // Slow weights live alongside the base optimizer's per-parameter state,
865        // under a reserved key so they round-trip through the same state dict.
866        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        // Extract Lookahead-specific state
881        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        // Pull the slow weights back out before handing the rest to the base.
892        //
893        // The lookup is driven by the *parameters*, never by iteration order over
894        // the state map: `HashMap` iteration is unordered, so pairing entries
895        // positionally would attach each slow weight to an arbitrary parameter.
896        // A parameter whose key is absent from the checkpoint simply gets no slow
897        // weight and is re-snapshotted by `initialize_slow_weights` on the next
898        // step. (Keys are derived from handle addresses, so — as for every
899        // optimizer in this crate — they identify parameters only within the
900        // process that produced them.)
901        self.slow_weights.clear();
902        for param in self.base_optimizer.parameters() {
903            let key = param_key(&param);
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        // Drop any slow weights belonging to parameters this optimizer no longer
911        // holds, so they cannot leak into the base optimizer's state.
912        for entry in state.state.values_mut() {
913            entry.remove("lookahead_slow_weight");
914        }
915
916        // Restore base optimizer type
917        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        // Load base optimizer state
923        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        // Test interface
954        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(&param, vec![1.0, 1.0]);
962
963        let mut optimizer = AdvancedAdam::with_params(0.1, vec![Arc::clone(&param)]);
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(&param, vec![0.5, 0.5]);
1003
1004        let mut optimizer = AdvancedAdam::with_params(0.1, vec![Arc::clone(&param)]);
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(&param)]);
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        // Test interface
1029        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(&param, vec![1.0, 1.0]);
1037
1038        let mut optimizer = LAMB::with_params(0.01, vec![Arc::clone(&param)]);
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        // Test interface
1055        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(&param, vec![1.0]);
1102
1103        let base = AdvancedAdam::with_params(0.1, vec![Arc::clone(&param)]);
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        // Slow weights still hold the initial value after a non-sync step.
1109        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        // Two Adam steps at lr = 0.1 move the fast weight to about -0.2, and
1136        // alpha = 0.5 places the slow weight halfway there: strictly between the
1137        // starting point and the fast trajectory, never beyond it.
1138        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}