Skip to main content

torsh_optim/
lr_scheduler.rs

1//! Learning rate schedulers
2
3use crate::{Optimizer, OptimizerError, OptimizerResult};
4use torsh_core::error::{Result, TorshError};
5
6/// Base trait for learning rate schedulers
7pub trait LRScheduler {
8    /// Update learning rates based on current epoch/step
9    fn step(&mut self) -> OptimizerResult<()>;
10
11    /// Update learning rates with optional metrics (for ReduceLROnPlateau)
12    fn step_with_metric(&mut self, metric: Option<f32>) -> OptimizerResult<()> {
13        // Default implementation ignores the metric
14        self.step()
15    }
16
17    /// Get current learning rates
18    fn get_last_lr(&self) -> &[f32];
19
20    /// Get base learning rates
21    fn get_base_lrs(&self) -> &[f32];
22
23    /// Get current epoch/step count
24    fn get_last_epoch(&self) -> i32;
25
26    /// Reset the scheduler state
27    fn reset(&mut self);
28
29    /// Get scheduler state for serialization
30    fn state_dict(&self) -> SchedulerState;
31
32    /// Load scheduler state from serialization
33    fn load_state_dict(&mut self, state: SchedulerState) -> OptimizerResult<()>;
34}
35
36/// Macro to implement common LRScheduler methods for schedulers with a `base` field
37#[macro_export]
38macro_rules! impl_base_scheduler_methods {
39    ($scheduler_type:ty, $scheduler_name:expr) => {
40        impl<O: Optimizer> LRScheduler for $scheduler_type {
41            fn step(&mut self) -> OptimizerResult<()> {
42                // Default implementation - should be overridden
43                Ok(())
44            }
45
46            fn get_last_lr(&self) -> &[f32] {
47                &self.base.last_lr
48            }
49
50            fn get_base_lrs(&self) -> &[f32] {
51                &self.base.base_lrs
52            }
53
54            fn get_last_epoch(&self) -> i32 {
55                self.base.last_epoch
56            }
57
58            fn reset(&mut self) {
59                self.base.last_epoch = 0;
60                self.base.last_lr = self.base.base_lrs.clone();
61                // Push the base rates back onto the optimizer so a reset actually
62                // restores them (per group, not broadcast).
63                let base_lrs = self.base.base_lrs.clone();
64                self.base.optimizer.set_lrs(&base_lrs);
65            }
66
67            fn state_dict(&self) -> SchedulerState {
68                let mut state = SchedulerState::new($scheduler_name.to_string());
69                state.last_epoch = self.base.last_epoch;
70                state.base_lrs = self.base.base_lrs.clone();
71                state.last_lr = self.base.last_lr.clone();
72                state
73            }
74
75            fn load_state_dict(&mut self, state: SchedulerState) -> OptimizerResult<()> {
76                self.base.last_epoch = state.last_epoch;
77                self.base.base_lrs = state.base_lrs;
78                self.base.last_lr = state.last_lr;
79                Ok(())
80            }
81        }
82    };
83}
84
85/// Macro to implement common LRScheduler methods with custom state handling
86#[macro_export]
87macro_rules! impl_scheduler_with_state {
88    ($scheduler_type:ty, $scheduler_name:expr, $state_fields:expr, $load_state_fields:expr) => {
89        impl<O: Optimizer> LRScheduler for $scheduler_type {
90            fn step(&mut self) -> OptimizerResult<()> {
91                // Default implementation - should be overridden
92                Ok(())
93            }
94
95            fn get_last_lr(&self) -> &[f32] {
96                &self.base.last_lr
97            }
98
99            fn get_base_lrs(&self) -> &[f32] {
100                &self.base.base_lrs
101            }
102
103            fn get_last_epoch(&self) -> i32 {
104                self.base.last_epoch
105            }
106
107            fn reset(&mut self) {
108                self.base.last_epoch = 0;
109                self.base.last_lr = self.base.base_lrs.clone();
110                // Push the base rates back onto the optimizer so a reset actually
111                // restores them (per group, not broadcast).
112                let base_lrs = self.base.base_lrs.clone();
113                self.base.optimizer.set_lrs(&base_lrs);
114            }
115
116            fn state_dict(&self) -> SchedulerState {
117                let mut state = SchedulerState::new($scheduler_name.to_string());
118                state.last_epoch = self.base.last_epoch;
119                state.base_lrs = self.base.base_lrs.clone();
120                state.last_lr = self.base.last_lr.clone();
121
122                // Add custom state fields
123                $state_fields(&self, &mut state);
124
125                state
126            }
127
128            fn load_state_dict(&mut self, state: SchedulerState) -> OptimizerResult<()> {
129                self.base.last_epoch = state.last_epoch;
130                self.base.base_lrs = state.base_lrs;
131                self.base.last_lr = state.last_lr;
132
133                // Load custom state fields
134                $load_state_fields(self, &state)?;
135
136                Ok(())
137            }
138        }
139    };
140}
141
142/// Scheduler state for serialization
143#[derive(Debug, Clone)]
144pub struct SchedulerState {
145    pub scheduler_type: String,
146    pub last_epoch: i32,
147    pub base_lrs: Vec<f32>,
148    pub last_lr: Vec<f32>,
149    pub state: std::collections::HashMap<String, f32>,
150}
151
152impl SchedulerState {
153    pub fn new(scheduler_type: String) -> Self {
154        Self {
155            scheduler_type,
156            last_epoch: 0,
157            base_lrs: Vec::new(),
158            last_lr: Vec::new(),
159            state: std::collections::HashMap::new(),
160        }
161    }
162}
163
164/// Base scheduler implementation
165pub struct BaseScheduler<O: Optimizer> {
166    pub optimizer: O,
167    pub base_lrs: Vec<f32>,
168    pub last_lr: Vec<f32>,
169    pub last_epoch: i32,
170}
171
172impl<O: Optimizer> BaseScheduler<O> {
173    pub fn new(optimizer: O) -> Self {
174        let base_lrs = optimizer.get_lr();
175        let last_lr = base_lrs.clone();
176
177        Self {
178            optimizer,
179            base_lrs,
180            last_lr,
181            last_epoch: 0,
182        }
183    }
184
185    /// Set the learning rates for all parameter groups
186    ///
187    /// A single entry is broadcast to every group; several entries are applied
188    /// group-by-group so per-group rates configured through
189    /// [`Optimizer::add_param_group`] are preserved.
190    pub fn set_learning_rates(&mut self, lrs: &[f32]) {
191        if !lrs.is_empty() {
192            if lrs.len() == 1 {
193                // Single learning rate - apply to all groups
194                self.optimizer.set_lr(lrs[0]);
195            } else {
196                self.optimizer.set_lrs(lrs);
197            }
198        }
199        self.last_lr = lrs.to_vec();
200    }
201
202    /// Get a mutable reference to the optimizer
203    pub fn optimizer_mut(&mut self) -> &mut O {
204        &mut self.optimizer
205    }
206
207    /// Get a reference to the optimizer
208    pub fn optimizer(&self) -> &O {
209        &self.optimizer
210    }
211
212    /// Increment the epoch counter
213    pub fn increment_epoch(&mut self) {
214        self.last_epoch += 1;
215    }
216
217    /// Set the epoch counter
218    pub fn set_epoch(&mut self, epoch: i32) {
219        self.last_epoch = epoch;
220    }
221}
222
223impl<O: Optimizer> LRScheduler for BaseScheduler<O> {
224    fn step(&mut self) -> OptimizerResult<()> {
225        self.increment_epoch();
226        // Base implementation does nothing - schedulers override this
227        Ok(())
228    }
229
230    fn get_last_lr(&self) -> &[f32] {
231        &self.last_lr
232    }
233
234    fn get_base_lrs(&self) -> &[f32] {
235        &self.base_lrs
236    }
237
238    fn get_last_epoch(&self) -> i32 {
239        self.last_epoch
240    }
241
242    fn reset(&mut self) {
243        self.last_epoch = 0;
244        self.last_lr = self.base_lrs.clone();
245        self.set_learning_rates(&self.base_lrs.clone());
246    }
247
248    fn state_dict(&self) -> SchedulerState {
249        let mut state = SchedulerState::new("BaseScheduler".to_string());
250        state.last_epoch = self.last_epoch;
251        state.base_lrs = self.base_lrs.clone();
252        state.last_lr = self.last_lr.clone();
253        state
254    }
255
256    fn load_state_dict(&mut self, state: SchedulerState) -> OptimizerResult<()> {
257        self.last_epoch = state.last_epoch;
258        self.base_lrs = state.base_lrs;
259        self.last_lr = state.last_lr.clone();
260        self.set_learning_rates(&state.last_lr);
261        Ok(())
262    }
263}
264
265/// Step learning rate scheduler
266pub struct StepLR<O: Optimizer> {
267    base: BaseScheduler<O>,
268    step_size: i32,
269    gamma: f32,
270}
271
272impl<O: Optimizer> StepLR<O> {
273    pub fn new(optimizer: O, step_size: i32, gamma: f32) -> Self {
274        Self {
275            base: BaseScheduler::new(optimizer),
276            step_size,
277            gamma,
278        }
279    }
280}
281
282impl<O: Optimizer> LRScheduler for StepLR<O> {
283    fn step(&mut self) -> OptimizerResult<()> {
284        self.base.increment_epoch();
285
286        let new_lrs: Vec<f32> = self
287            .base
288            .base_lrs
289            .iter()
290            .map(|&base_lr| {
291                let num_steps = self.base.last_epoch / self.step_size;
292                base_lr * self.gamma.powi(num_steps)
293            })
294            .collect();
295
296        self.base.set_learning_rates(&new_lrs);
297        Ok(())
298    }
299
300    fn get_last_lr(&self) -> &[f32] {
301        self.base.get_last_lr()
302    }
303
304    fn get_base_lrs(&self) -> &[f32] {
305        self.base.get_base_lrs()
306    }
307
308    fn get_last_epoch(&self) -> i32 {
309        self.base.get_last_epoch()
310    }
311
312    fn reset(&mut self) {
313        self.base.reset()
314    }
315
316    fn state_dict(&self) -> SchedulerState {
317        let mut state = self.base.state_dict();
318        state.scheduler_type = "StepLR".to_string();
319        state
320            .state
321            .insert("step_size".to_string(), self.step_size as f32);
322        state.state.insert("gamma".to_string(), self.gamma);
323        state
324    }
325
326    fn load_state_dict(&mut self, state: SchedulerState) -> OptimizerResult<()> {
327        self.base.load_state_dict(state.clone())?;
328
329        if let Some(&step_size) = state.state.get("step_size") {
330            self.step_size = step_size as i32;
331        }
332        if let Some(&gamma) = state.state.get("gamma") {
333            self.gamma = gamma;
334        }
335
336        Ok(())
337    }
338}
339
340/// Exponential learning rate scheduler
341pub struct ExponentialLR<O: Optimizer> {
342    base: BaseScheduler<O>,
343    gamma: f32,
344}
345
346impl<O: Optimizer> ExponentialLR<O> {
347    pub fn new(optimizer: O, gamma: f32) -> Self {
348        Self {
349            base: BaseScheduler::new(optimizer),
350            gamma,
351        }
352    }
353
354    /// Get a reference to the wrapped optimizer
355    pub fn optimizer(&self) -> &O {
356        &self.base.optimizer
357    }
358
359    /// Get a mutable reference to the wrapped optimizer
360    pub fn optimizer_mut(&mut self) -> &mut O {
361        &mut self.base.optimizer
362    }
363}
364
365impl<O: Optimizer> LRScheduler for ExponentialLR<O> {
366    fn step(&mut self) -> OptimizerResult<()> {
367        self.base.last_epoch += 1;
368
369        let new_lrs: Vec<f32> = self
370            .base
371            .base_lrs
372            .iter()
373            .map(|&base_lr| base_lr * self.gamma.powi(self.base.last_epoch))
374            .collect();
375
376        self.base.optimizer.set_lrs(&new_lrs);
377
378        self.base.last_lr = new_lrs;
379        Ok(())
380    }
381
382    fn get_last_lr(&self) -> &[f32] {
383        &self.base.last_lr
384    }
385
386    fn get_base_lrs(&self) -> &[f32] {
387        &self.base.base_lrs
388    }
389
390    fn get_last_epoch(&self) -> i32 {
391        self.base.last_epoch
392    }
393
394    fn reset(&mut self) {
395        self.base.last_epoch = 0;
396        self.base.last_lr = self.base.base_lrs.clone();
397        // Push the base rates back onto the optimizer so a reset actually
398        // restores them (per group, not broadcast).
399        let base_lrs = self.base.base_lrs.clone();
400        self.base.optimizer.set_lrs(&base_lrs);
401    }
402
403    fn state_dict(&self) -> SchedulerState {
404        let mut state = SchedulerState::new("ExponentialLR".to_string());
405        state.last_epoch = self.base.last_epoch;
406        state.base_lrs = self.base.base_lrs.clone();
407        state.last_lr = self.base.last_lr.clone();
408        state.state.insert("gamma".to_string(), self.gamma);
409        state
410    }
411
412    fn load_state_dict(&mut self, state: SchedulerState) -> OptimizerResult<()> {
413        self.base.last_epoch = state.last_epoch;
414        self.base.base_lrs = state.base_lrs;
415        self.base.last_lr = state.last_lr;
416        if let Some(&gamma) = state.state.get("gamma") {
417            self.gamma = gamma;
418        }
419        Ok(())
420    }
421}
422
423/// Cosine annealing learning rate scheduler
424pub struct CosineAnnealingLR<O: Optimizer> {
425    base: BaseScheduler<O>,
426    t_max: i32,
427    eta_min: f32,
428}
429
430impl<O: Optimizer> CosineAnnealingLR<O> {
431    pub fn new(optimizer: O, t_max: i32, eta_min: f32) -> Self {
432        Self {
433            base: BaseScheduler::new(optimizer),
434            t_max,
435            eta_min,
436        }
437    }
438}
439
440impl<O: Optimizer> LRScheduler for CosineAnnealingLR<O> {
441    fn step(&mut self) -> OptimizerResult<()> {
442        self.base.last_epoch += 1;
443
444        let new_lrs: Vec<f32> = self
445            .base
446            .base_lrs
447            .iter()
448            .map(|&base_lr| {
449                self.eta_min
450                    + (base_lr - self.eta_min)
451                        * (1.0
452                            + (std::f32::consts::PI * self.base.last_epoch as f32
453                                / self.t_max as f32)
454                                .cos())
455                        / 2.0
456            })
457            .collect();
458
459        self.base.optimizer.set_lrs(&new_lrs);
460
461        self.base.last_lr = new_lrs;
462        Ok(())
463    }
464
465    fn get_last_lr(&self) -> &[f32] {
466        &self.base.last_lr
467    }
468
469    fn get_base_lrs(&self) -> &[f32] {
470        &self.base.base_lrs
471    }
472
473    fn get_last_epoch(&self) -> i32 {
474        self.base.last_epoch
475    }
476
477    fn reset(&mut self) {
478        self.base.last_epoch = 0;
479        self.base.last_lr = self.base.base_lrs.clone();
480        // Push the base rates back onto the optimizer so a reset actually
481        // restores them (per group, not broadcast).
482        let base_lrs = self.base.base_lrs.clone();
483        self.base.optimizer.set_lrs(&base_lrs);
484    }
485
486    fn state_dict(&self) -> SchedulerState {
487        let mut state = SchedulerState::new("CosineAnnealingLR".to_string());
488        state.last_epoch = self.base.last_epoch;
489        state.base_lrs = self.base.base_lrs.clone();
490        state.last_lr = self.base.last_lr.clone();
491        state.state.insert("t_max".to_string(), self.t_max as f32);
492        state.state.insert("eta_min".to_string(), self.eta_min);
493        state
494    }
495
496    fn load_state_dict(&mut self, state: SchedulerState) -> OptimizerResult<()> {
497        self.base.last_epoch = state.last_epoch;
498        self.base.base_lrs = state.base_lrs;
499        self.base.last_lr = state.last_lr;
500        if let Some(&t_max) = state.state.get("t_max") {
501            self.t_max = t_max as i32;
502        }
503        if let Some(&eta_min) = state.state.get("eta_min") {
504            self.eta_min = eta_min;
505        }
506        Ok(())
507    }
508}
509
510/// Reduce learning rate on plateau
511pub struct ReduceLROnPlateau<O: Optimizer> {
512    optimizer: O,
513    mode: String,
514    factor: f32,
515    patience: i32,
516    threshold: f32,
517    threshold_mode: String,
518    cooldown: i32,
519    min_lr: f32,
520    eps: f32,
521    best: Option<f32>,
522    num_bad_epochs: i32,
523    cooldown_counter: i32,
524}
525
526impl<O: Optimizer> ReduceLROnPlateau<O> {
527    #[allow(clippy::too_many_arguments)]
528    pub fn new(
529        optimizer: O,
530        mode: &str,
531        factor: f32,
532        patience: i32,
533        threshold: f32,
534        threshold_mode: &str,
535        cooldown: i32,
536        min_lr: f32,
537        eps: f32,
538    ) -> Result<Self> {
539        if factor >= 1.0 {
540            return Err(TorshError::Other("Factor should be < 1.0".to_string()));
541        }
542
543        Ok(Self {
544            optimizer,
545            mode: mode.to_string(),
546            factor,
547            patience,
548            threshold,
549            threshold_mode: threshold_mode.to_string(),
550            cooldown,
551            min_lr,
552            eps,
553            best: None,
554            num_bad_epochs: 0,
555            cooldown_counter: 0,
556        })
557    }
558
559    pub fn step(&mut self, metrics: f32) {
560        let current = metrics;
561
562        if self.best.is_none() {
563            self.best = Some(current);
564        } else {
565            let best_value = self.best.expect("best should exist after is_none check");
566            let is_better = match self.mode.as_str() {
567                "min" => match self.threshold_mode.as_str() {
568                    "rel" => current < best_value * (1.0 - self.threshold),
569                    "abs" => current < best_value - self.threshold,
570                    _ => false,
571                },
572                "max" => match self.threshold_mode.as_str() {
573                    "rel" => current > best_value * (1.0 + self.threshold),
574                    "abs" => current > best_value + self.threshold,
575                    _ => false,
576                },
577                _ => false,
578            };
579
580            if is_better {
581                self.best = Some(current);
582                self.num_bad_epochs = 0;
583            } else {
584                self.num_bad_epochs += 1;
585            }
586
587            if self.cooldown_counter > 0 {
588                self.cooldown_counter -= 1;
589                self.num_bad_epochs = 0;
590            }
591
592            if self.num_bad_epochs > self.patience {
593                self.reduce_lr();
594                self.cooldown_counter = self.cooldown;
595                self.num_bad_epochs = 0;
596            }
597        }
598    }
599
600    fn reduce_lr(&mut self) {
601        let old_lrs = self.optimizer.get_lr();
602        let new_lrs: Vec<f32> = old_lrs
603            .iter()
604            .map(|&lr| (lr * self.factor).max(self.min_lr))
605            .collect();
606
607        // Only reduce the groups whose change is significant; the others keep
608        // their current rate (a per-group assignment, not a broadcast).
609        let mut applied: Vec<f32> = old_lrs.clone();
610        let mut any_reduced = false;
611        for (idx, (old_lr, new_lr)) in old_lrs.iter().zip(new_lrs.iter()).enumerate() {
612            if old_lr - new_lr > self.eps {
613                applied[idx] = *new_lr;
614                any_reduced = true;
615                log::debug!("Reducing learning rate from {old_lr} to {new_lr}");
616            }
617        }
618        if any_reduced {
619            self.optimizer.set_lrs(&applied);
620        }
621    }
622}
623
624/// One cycle learning rate scheduler
625pub struct OneCycleLR<O: Optimizer> {
626    base: BaseScheduler<O>,
627    max_lr: Vec<f32>,
628    total_steps: i32,
629    pct_start: f32,
630    anneal_strategy: String,
631    #[allow(dead_code)]
632    cycle_momentum: bool,
633    #[allow(dead_code)]
634    base_momentum: f32,
635    #[allow(dead_code)]
636    max_momentum: f32,
637    #[allow(dead_code)]
638    div_factor: f32,
639    final_div_factor: f32,
640    step_count: i32,
641}
642
643impl<O: Optimizer> OneCycleLR<O> {
644    #[allow(clippy::too_many_arguments)]
645    pub fn new(
646        optimizer: O,
647        max_lr: Vec<f32>,
648        total_steps: i32,
649        pct_start: Option<f32>,
650        anneal_strategy: Option<&str>,
651        cycle_momentum: Option<bool>,
652        base_momentum: Option<f32>,
653        max_momentum: Option<f32>,
654        div_factor: Option<f32>,
655        final_div_factor: Option<f32>,
656    ) -> Self {
657        let pct_start = pct_start.unwrap_or(0.3);
658        let anneal_strategy = anneal_strategy.unwrap_or("cos").to_string();
659        let cycle_momentum = cycle_momentum.unwrap_or(true);
660        let base_momentum = base_momentum.unwrap_or(0.85);
661        let max_momentum = max_momentum.unwrap_or(0.95);
662        let div_factor = div_factor.unwrap_or(25.0);
663        let final_div_factor = final_div_factor.unwrap_or(10000.0);
664
665        let mut base = BaseScheduler::new(optimizer);
666
667        // Initialize base learning rates
668        base.base_lrs = max_lr.iter().map(|&lr| lr / div_factor).collect();
669
670        Self {
671            base,
672            max_lr,
673            total_steps,
674            pct_start,
675            anneal_strategy,
676            cycle_momentum,
677            base_momentum,
678            max_momentum,
679            div_factor,
680            final_div_factor,
681            step_count: 0,
682        }
683    }
684}
685
686impl<O: Optimizer> LRScheduler for OneCycleLR<O> {
687    fn step(&mut self) -> OptimizerResult<()> {
688        self.step_count += 1;
689
690        let step_num = self.step_count as f32;
691        let step_size_up = (self.pct_start * self.total_steps as f32).floor();
692        let step_size_down = self.total_steps as f32 - step_size_up;
693
694        let new_lrs: Vec<f32> = if step_num <= step_size_up {
695            // Increase phase
696            let computed_lr =
697                |base_lr: f32, max_lr: f32| base_lr + (max_lr - base_lr) * step_num / step_size_up;
698
699            self.base
700                .base_lrs
701                .iter()
702                .zip(self.max_lr.iter())
703                .map(|(&base, &max)| computed_lr(base, max))
704                .collect()
705        } else {
706            // Decrease phase
707            let down_step_num = step_num - step_size_up;
708            match self.anneal_strategy.as_str() {
709                "cos" => {
710                    let computed_lr = |max_lr: f32, base_lr: f32| {
711                        let min_lr = base_lr / self.final_div_factor;
712                        min_lr
713                            + (max_lr - min_lr)
714                                * (1.0
715                                    + (std::f32::consts::PI * down_step_num / step_size_down).cos())
716                                / 2.0
717                    };
718
719                    self.max_lr
720                        .iter()
721                        .zip(self.base.base_lrs.iter())
722                        .map(|(&max, &base)| computed_lr(max, base))
723                        .collect()
724                }
725                "linear" => {
726                    let computed_lr = |max_lr: f32, base_lr: f32| {
727                        let min_lr = base_lr / self.final_div_factor;
728                        max_lr - (max_lr - min_lr) * down_step_num / step_size_down
729                    };
730
731                    self.max_lr
732                        .iter()
733                        .zip(self.base.base_lrs.iter())
734                        .map(|(&max, &base)| computed_lr(max, base))
735                        .collect()
736                }
737                _ => {
738                    return Err(OptimizerError::InvalidParameter(format!(
739                        "Unknown anneal strategy: {}",
740                        self.anneal_strategy
741                    )))
742                }
743            }
744        };
745
746        // Update learning rates
747        self.base.optimizer.set_lrs(&new_lrs);
748
749        self.base.last_lr = new_lrs;
750        Ok(())
751    }
752
753    fn get_last_lr(&self) -> &[f32] {
754        &self.base.last_lr
755    }
756
757    fn get_base_lrs(&self) -> &[f32] {
758        &self.base.base_lrs
759    }
760
761    fn get_last_epoch(&self) -> i32 {
762        self.step_count
763    }
764
765    fn reset(&mut self) {
766        self.step_count = 0;
767        self.base.last_lr = self.base.base_lrs.clone();
768        // Push the base rates back onto the optimizer so a reset actually
769        // restores them (per group, not broadcast).
770        let base_lrs = self.base.base_lrs.clone();
771        self.base.optimizer.set_lrs(&base_lrs);
772    }
773
774    fn state_dict(&self) -> SchedulerState {
775        let mut state = SchedulerState::new("OneCycleLR".to_string());
776        state.last_epoch = self.step_count;
777        state.base_lrs = self.base.base_lrs.clone();
778        state.last_lr = self.base.last_lr.clone();
779        state
780            .state
781            .insert("total_steps".to_string(), self.total_steps as f32);
782        state.state.insert("pct_start".to_string(), self.pct_start);
783        state
784            .state
785            .insert("final_div_factor".to_string(), self.final_div_factor);
786        state
787    }
788
789    fn load_state_dict(&mut self, state: SchedulerState) -> OptimizerResult<()> {
790        self.step_count = state.last_epoch;
791        self.base.base_lrs = state.base_lrs;
792        self.base.last_lr = state.last_lr;
793        if let Some(&total_steps) = state.state.get("total_steps") {
794            self.total_steps = total_steps as i32;
795        }
796        if let Some(&pct_start) = state.state.get("pct_start") {
797            self.pct_start = pct_start;
798        }
799        if let Some(&final_div_factor) = state.state.get("final_div_factor") {
800            self.final_div_factor = final_div_factor;
801        }
802        Ok(())
803    }
804}