Skip to main content

torsh_optim/
prodigy.rs

1//! Prodigy (An Adaptive Learning Rate Method) optimizer
2//!
3//! Prodigy is a state-of-the-art adaptive learning rate optimizer that automatically tunes the
4//! learning rate without manual tuning. It combines ideas from D-Adaptation with improved
5//! stability and convergence properties.
6//!
7//! ## Key Innovation
8//!
9//! Traditional optimizers require careful learning rate tuning. Prodigy automatically estimates
10//! the optimal learning rate by tracking gradient statistics and adapting based on:
11//! - Distance traveled in parameter space
12//! - Gradient variance
13//! - Convergence progress
14//!
15//! This makes it extremely user-friendly - you can use `lr=1.0` for almost any problem!
16//!
17//! ## Key Features
18//!
19//! - **Zero LR Tuning**: Use lr=1.0 for most problems
20//! - **Automatic Adaptation**: Learns optimal learning rate during training
21//! - **Robust**: Works across different domains (vision, NLP, RL)
22//! - **Memory Efficient**: Similar memory overhead to Adam
23//! - **Fast Convergence**: Often matches or beats carefully tuned Adam
24//!
25//! ## Algorithm
26//!
27//! ```text
28//! // Initialize
29//! d_0 = initial_d  // learning rate estimate
30//! s_0 = 0          // distance estimate
31//!
32//! // Update momentum and variance (like Adam)
33//! m_t = β₁ * m_{t-1} + (1 - β₁) * g_t
34//! v_t = β₂ * v_{t-1} + (1 - β₂) * g_t²
35//!
36//! // Estimate distance traveled
37//! s_t = s_{t-1} + ||θ_t - θ_{t-1}||
38//!
39//! // Adapt learning rate estimate
40//! if t > k:  // after warmup
41//!     d_t = d_{t-1} * (s_t / s_{t-1})^growth_rate
42//!
43//! // Compute adaptive step size
44//! α_t = lr / (d_t * √t)
45//!
46//! // Update parameters
47//! θ_t = θ_{t-1} - α_t * m̂_t / (√v̂_t + ε)
48//! ```
49//!
50//! Where:
51//! - `d_t` is the adaptive learning rate scale
52//! - `s_t` is the cumulative distance estimate
53//! - `m_t, v_t` are first and second moments
54//! - `β₁, β₂` are exponential decay rates
55//! - `growth_rate` controls adaptation speed (typically 1.0)
56//!
57//! ## Typical Hyperparameters
58//!
59//! - Learning rate: **1.0** (yes, really!)
60//! - β₁ (beta1): 0.9
61//! - β₂ (beta2): 0.999
62//! - Growth rate: 1.0
63//! - Weight decay: 0.0 (or small value like 0.01)
64//! - Initial d: 1e-6
65//! - Warmup steps: 0 (can use 100-1000 for stability)
66//!
67//! ## Usage Example
68//!
69//! ```rust
70//! # use torsh_tensor::creation::randn;
71//! # use torsh_core::error::Result;
72//! # fn main() -> Result<()> {
73//! use torsh_optim::prelude::{Prodigy, Optimizer};
74//! use parking_lot::RwLock;
75//! use std::sync::Arc;
76//!
77//! let param = Arc::new(RwLock::new(randn::<f32>(&[768, 768])?));
78//! let params = vec![param];
79//!
80//! // No learning rate tuning needed!
81//! let mut optimizer = Prodigy::new(
82//!     params,
83//!     1.0,    // lr (use 1.0 for almost everything!)
84//!     0.9,    // beta1
85//!     0.999,  // beta2
86//!     0.0,    // weight_decay
87//! );
88//!
89//! // Training loop - it just works!
90//! for _step in 0..10000 {
91//!     // ... compute gradients ...
92//!     optimizer.step()?;
93//!     optimizer.zero_grad();
94//! }
95//! # Ok(())
96//! # }
97//! ```
98//!
99//! ## Reference
100//!
101//! Mishchenko, K., & Defazio, A. (2024).
102//! "Prodigy: An Adaptive Learning Rate Method".
103//! arXiv preprint arXiv:2401.04536.
104
105use crate::{
106    Optimizer, OptimizerError, OptimizerResult, OptimizerState, ParamGroup, ParamGroupState,
107};
108use parking_lot::RwLock;
109use std::collections::HashMap;
110use std::sync::Arc;
111use torsh_tensor::Tensor;
112
113/// Prodigy optimizer configuration
114#[derive(Debug, Clone)]
115pub struct ProdigyConfig {
116    /// Learning rate (use 1.0 for most problems!)
117    pub lr: f32,
118    /// Beta1 for momentum (typically 0.9)
119    pub beta1: f32,
120    /// Beta2 for variance (typically 0.999)
121    pub beta2: f32,
122    /// Growth rate for d adaptation (typically 1.0)
123    pub growth_rate: f32,
124    /// Initial d estimate (typically 1e-6)
125    pub initial_d: f32,
126    /// Weight decay coefficient
127    pub weight_decay: f32,
128    /// Epsilon for numerical stability
129    pub eps: f32,
130    /// Warmup steps (0 for no warmup)
131    pub warmup_steps: usize,
132}
133
134impl Default for ProdigyConfig {
135    fn default() -> Self {
136        Self {
137            lr: 1.0,
138            beta1: 0.9,
139            beta2: 0.999,
140            growth_rate: 1.0,
141            initial_d: 1e-6,
142            weight_decay: 0.0,
143            eps: 1e-8,
144            warmup_steps: 0,
145        }
146    }
147}
148
149/// Prodigy optimizer (Adaptive Learning Rate Method)
150///
151/// Prodigy automatically tunes the learning rate without manual tuning.
152/// Just use lr=1.0 for almost any problem!
153pub struct Prodigy {
154    /// Parameter groups
155    param_groups: Vec<ParamGroup>,
156    /// Base learning rate (typically 1.0)
157    lr: f32,
158    /// Beta1 (momentum coefficient)
159    beta1: f32,
160    /// Beta2 (variance coefficient)
161    beta2: f32,
162    /// Growth rate for d adaptation
163    growth_rate: f32,
164    /// Current d estimate (adaptive lr scale)
165    d: f32,
166    /// Previous d estimate
167    d_prev: f32,
168    /// Cumulative distance estimate
169    s: f32,
170    /// Previous cumulative distance
171    s_prev: f32,
172    /// Weight decay coefficient
173    weight_decay: f32,
174    /// Epsilon for numerical stability
175    eps: f32,
176    /// Warmup steps
177    warmup_steps: usize,
178    /// Momentum buffers
179    momentum: HashMap<String, Tensor>,
180    /// Variance buffers
181    variance: HashMap<String, Tensor>,
182    /// Previous parameters (for distance computation)
183    prev_params: HashMap<String, Tensor>,
184    /// Current step count
185    step_count: usize,
186}
187
188impl Prodigy {
189    /// Create a new Prodigy optimizer
190    ///
191    /// # Arguments
192    ///
193    /// * `params` - Parameters to optimize
194    /// * `lr` - Learning rate (use 1.0 for most problems!)
195    /// * `beta1` - Momentum coefficient (default: 0.9)
196    /// * `beta2` - Variance coefficient (default: 0.999)
197    /// * `weight_decay` - Weight decay coefficient (default: 0.0)
198    ///
199    /// # Example
200    ///
201    /// ```rust
202    /// # use torsh_tensor::creation::randn;
203    /// # use torsh_core::error::Result;
204    /// # fn main() -> Result<()> {
205    /// use torsh_optim::prelude::Prodigy;
206    /// use parking_lot::RwLock;
207    /// use std::sync::Arc;
208    ///
209    /// let param = Arc::new(RwLock::new(randn::<f32>(&[100, 100])?));
210    /// let params = vec![param];
211    ///
212    /// // Use lr=1.0 - it adapts automatically!
213    /// let optimizer = Prodigy::new(params, 1.0, 0.9, 0.999, 0.0);
214    /// # Ok(())
215    /// # }
216    /// ```
217    pub fn new(
218        params: Vec<Arc<RwLock<Tensor>>>,
219        lr: f32,
220        beta1: f32,
221        beta2: f32,
222        weight_decay: f32,
223    ) -> Self {
224        let param_group = ParamGroup::new(params, lr);
225        let config = ProdigyConfig::default();
226        Self {
227            param_groups: vec![param_group],
228            lr,
229            beta1,
230            beta2,
231            growth_rate: config.growth_rate,
232            d: config.initial_d,
233            d_prev: config.initial_d,
234            s: 0.0,
235            s_prev: 0.0,
236            weight_decay,
237            eps: config.eps,
238            warmup_steps: config.warmup_steps,
239            momentum: HashMap::new(),
240            variance: HashMap::new(),
241            prev_params: HashMap::new(),
242            step_count: 0,
243        }
244    }
245
246    /// Create from configuration
247    pub fn from_config(params: Vec<Arc<RwLock<Tensor>>>, config: ProdigyConfig) -> Self {
248        let mut optimizer = Self::new(
249            params,
250            config.lr,
251            config.beta1,
252            config.beta2,
253            config.weight_decay,
254        );
255        optimizer.growth_rate = config.growth_rate;
256        optimizer.d = config.initial_d;
257        optimizer.d_prev = config.initial_d;
258        optimizer.eps = config.eps;
259        optimizer.warmup_steps = config.warmup_steps;
260        optimizer
261    }
262
263    /// Builder for Prodigy optimizer
264    pub fn builder() -> ProdigyBuilder {
265        ProdigyBuilder::default()
266    }
267
268    /// Get current learning rate scale (d)
269    pub fn get_d(&self) -> f32 {
270        self.d
271    }
272
273    /// Get effective learning rate
274    pub fn get_effective_lr(&self) -> f32 {
275        if self.step_count == 0 {
276            return 0.0;
277        }
278        self.lr / (self.d * (self.step_count as f32).sqrt())
279    }
280}
281
282impl Optimizer for Prodigy {
283    fn step(&mut self) -> OptimizerResult<()> {
284        self.step_count += 1;
285
286        // Compute warmup factor
287        let warmup_factor = if self.warmup_steps > 0 && self.step_count <= self.warmup_steps {
288            (self.step_count as f32) / (self.warmup_steps as f32)
289        } else {
290            1.0
291        };
292
293        // Adaptive step size: α = lr / (d * √t)
294        let base_step_size = self.lr / (self.d * (self.step_count as f32).sqrt());
295        let step_size = base_step_size * warmup_factor;
296
297        let mut distance_sum = 0.0f32;
298
299        for group in &self.param_groups {
300            let beta1 = self.beta1;
301            let beta2 = self.beta2;
302            let weight_decay = self.weight_decay;
303            let eps = self.eps;
304
305            for (idx, param) in group.params.iter().enumerate() {
306                let mut param_guard = param.write();
307
308                // Skip parameters without gradients
309                if !param_guard.has_grad() {
310                    continue;
311                }
312
313                let grad = param_guard
314                    .grad()
315                    .ok_or_else(|| OptimizerError::InvalidInput("No gradient found".to_string()))?;
316
317                let param_key = format!("param_{}", idx);
318
319                // Get or initialize momentum and variance
320                let m_entry = self.momentum.entry(param_key.clone()).or_insert_with(|| {
321                    grad.zeros_like().expect("Failed to create momentum buffer")
322                });
323                let v_entry = self.variance.entry(param_key.clone()).or_insert_with(|| {
324                    grad.zeros_like().expect("Failed to create variance buffer")
325                });
326
327                // Update momentum: m = β₁ * m + (1 - β₁) * g
328                let new_m = m_entry
329                    .mul_scalar(beta1)
330                    .map_err(|e| OptimizerError::TensorError(e))?
331                    .add(
332                        &grad
333                            .mul_scalar(1.0 - beta1)
334                            .map_err(|e| OptimizerError::TensorError(e))?,
335                    )
336                    .map_err(|e| OptimizerError::TensorError(e))?;
337
338                // Update variance: v = β₂ * v + (1 - β₂) * g²
339                let grad_squared = grad
340                    .mul(&grad)
341                    .map_err(|e| OptimizerError::TensorError(e))?;
342                let new_v = v_entry
343                    .mul_scalar(beta2)
344                    .map_err(|e| OptimizerError::TensorError(e))?
345                    .add(
346                        &grad_squared
347                            .mul_scalar(1.0 - beta2)
348                            .map_err(|e| OptimizerError::TensorError(e))?,
349                    )
350                    .map_err(|e| OptimizerError::TensorError(e))?;
351
352                // Bias correction
353                let bias_correction1 = 1.0 - beta1.powi(self.step_count as i32);
354                let bias_correction2 = 1.0 - beta2.powi(self.step_count as i32);
355
356                let m_hat = new_m
357                    .mul_scalar(1.0 / bias_correction1)
358                    .map_err(|e| OptimizerError::TensorError(e))?;
359                let v_hat = new_v
360                    .mul_scalar(1.0 / bias_correction2)
361                    .map_err(|e| OptimizerError::TensorError(e))?;
362
363                // Compute update: m̂ / (√v̂ + ε)
364                let v_sqrt = v_hat
365                    .sqrt()
366                    .map_err(|e| OptimizerError::TensorError(e))?
367                    .add_scalar(eps)
368                    .map_err(|e| OptimizerError::TensorError(e))?;
369
370                let update_direction = m_hat
371                    .div(&v_sqrt)
372                    .map_err(|e| OptimizerError::TensorError(e))?;
373
374                // Apply weight decay if specified (AdamW-style)
375                let param_data = param_guard.clone();
376                let update = if weight_decay > 0.0 {
377                    let decay_term = param_data
378                        .mul_scalar(weight_decay * step_size)
379                        .map_err(|e| OptimizerError::TensorError(e))?;
380                    update_direction
381                        .mul_scalar(step_size)
382                        .map_err(|e| OptimizerError::TensorError(e))?
383                        .add(&decay_term)
384                        .map_err(|e| OptimizerError::TensorError(e))?
385                } else {
386                    update_direction
387                        .mul_scalar(step_size)
388                        .map_err(|e| OptimizerError::TensorError(e))?
389                };
390
391                // Update parameters
392                let new_param = param_data
393                    .sub(&update)
394                    .map_err(|e| OptimizerError::TensorError(e))?;
395
396                // Compute distance for d adaptation
397                if let Some(prev_param) = self.prev_params.get(&param_key) {
398                    let param_diff = new_param
399                        .sub(prev_param)
400                        .map_err(|e| OptimizerError::TensorError(e))?;
401                    let diff_norm = param_diff
402                        .norm()
403                        .map_err(|e| OptimizerError::TensorError(e))?;
404                    distance_sum += diff_norm
405                        .to_vec()
406                        .map_err(|e| OptimizerError::TensorError(e))?[0];
407                }
408
409                // Store previous parameter for next iteration
410                self.prev_params.insert(param_key.clone(), param_data);
411
412                // Update momentum and variance
413                *m_entry = new_m;
414                *v_entry = new_v;
415                *param_guard = new_param;
416            }
417        }
418
419        // Update distance estimate
420        self.s_prev = self.s;
421        self.s += distance_sum;
422
423        // Adapt d after warmup
424        if self.step_count > self.warmup_steps.max(1) && self.s_prev > 0.0 {
425            let ratio = self.s / self.s_prev;
426            self.d_prev = self.d;
427            self.d = self.d * ratio.powf(self.growth_rate);
428
429            // Clamp d to prevent instability
430            self.d = self.d.max(1e-12).min(1e12);
431        }
432
433        Ok(())
434    }
435
436    fn zero_grad(&mut self) {
437        for group in &self.param_groups {
438            group.zero_grad();
439        }
440    }
441
442    fn get_lr(&self) -> Vec<f32> {
443        self.param_groups.iter().map(|g| g.lr).collect()
444    }
445
446    fn set_lr(&mut self, lr: f32) {
447        self.lr = lr;
448        for group in &mut self.param_groups {
449            group.lr = lr;
450        }
451    }
452
453    fn add_param_group(&mut self, params: Vec<Arc<RwLock<Tensor>>>, options: HashMap<String, f32>) {
454        let lr = options.get("lr").copied().unwrap_or(self.lr);
455        let group = ParamGroup::new(params, lr).with_options(options);
456        self.param_groups.push(group);
457    }
458
459    fn parameters(&self) -> Vec<Arc<RwLock<Tensor>>> {
460        crate::optimizer::collect_parameters(&self.param_groups)
461    }
462
463    fn state_dict(&self) -> OptimizerResult<OptimizerState> {
464        let param_group_states = self
465            .param_groups
466            .iter()
467            .map(|g| ParamGroupState::from_param_group(g))
468            .collect();
469
470        let mut state = HashMap::new();
471        for (key, _) in &self.momentum {
472            let mut param_state = HashMap::new();
473            if let Some(m) = self.momentum.get(key) {
474                param_state.insert("momentum".to_string(), m.clone());
475            }
476            if let Some(v) = self.variance.get(key) {
477                param_state.insert("variance".to_string(), v.clone());
478            }
479            if let Some(prev) = self.prev_params.get(key) {
480                param_state.insert("prev_param".to_string(), prev.clone());
481            }
482            state.insert(key.clone(), param_state);
483        }
484
485        let mut global_state = HashMap::new();
486        global_state.insert("beta1".to_string(), self.beta1);
487        global_state.insert("beta2".to_string(), self.beta2);
488        global_state.insert("growth_rate".to_string(), self.growth_rate);
489        global_state.insert("d".to_string(), self.d);
490        global_state.insert("d_prev".to_string(), self.d_prev);
491        global_state.insert("s".to_string(), self.s);
492        global_state.insert("s_prev".to_string(), self.s_prev);
493        global_state.insert("weight_decay".to_string(), self.weight_decay);
494        global_state.insert("step_count".to_string(), self.step_count as f32);
495
496        Ok(OptimizerState {
497            optimizer_type: "Prodigy".to_string(),
498            version: "1.0".to_string(),
499            param_groups: param_group_states,
500            state,
501            global_state,
502        })
503    }
504
505    fn load_state_dict(&mut self, state: OptimizerState) -> OptimizerResult<()> {
506        if state.optimizer_type != "Prodigy" {
507            return Err(OptimizerError::InvalidInput(format!(
508                "Expected Prodigy state dict, got {}",
509                state.optimizer_type
510            )));
511        }
512
513        // Restore hyperparameters
514        if let Some(&beta1) = state.global_state.get("beta1") {
515            self.beta1 = beta1;
516        }
517        if let Some(&beta2) = state.global_state.get("beta2") {
518            self.beta2 = beta2;
519        }
520        if let Some(&growth_rate) = state.global_state.get("growth_rate") {
521            self.growth_rate = growth_rate;
522        }
523        if let Some(&d) = state.global_state.get("d") {
524            self.d = d;
525        }
526        if let Some(&d_prev) = state.global_state.get("d_prev") {
527            self.d_prev = d_prev;
528        }
529        if let Some(&s) = state.global_state.get("s") {
530            self.s = s;
531        }
532        if let Some(&s_prev) = state.global_state.get("s_prev") {
533            self.s_prev = s_prev;
534        }
535        if let Some(&weight_decay) = state.global_state.get("weight_decay") {
536            self.weight_decay = weight_decay;
537        }
538        if let Some(&step_count) = state.global_state.get("step_count") {
539            self.step_count = step_count as usize;
540        }
541
542        // Restore optimizer state
543        self.momentum.clear();
544        self.variance.clear();
545        self.prev_params.clear();
546
547        for (key, param_state) in state.state {
548            if let Some(m) = param_state.get("momentum") {
549                self.momentum.insert(key.clone(), m.clone());
550            }
551            if let Some(v) = param_state.get("variance") {
552                self.variance.insert(key.clone(), v.clone());
553            }
554            if let Some(prev) = param_state.get("prev_param") {
555                self.prev_params.insert(key.clone(), prev.clone());
556            }
557        }
558
559        Ok(())
560    }
561}
562
563/// Builder for Prodigy optimizer
564#[derive(Debug, Clone)]
565pub struct ProdigyBuilder {
566    params: Vec<Arc<RwLock<Tensor>>>,
567    lr: f32,
568    beta1: f32,
569    beta2: f32,
570    growth_rate: f32,
571    initial_d: f32,
572    weight_decay: f32,
573    eps: f32,
574    warmup_steps: usize,
575}
576
577impl Default for ProdigyBuilder {
578    fn default() -> Self {
579        let config = ProdigyConfig::default();
580        Self {
581            params: Vec::new(),
582            lr: config.lr,
583            beta1: config.beta1,
584            beta2: config.beta2,
585            growth_rate: config.growth_rate,
586            initial_d: config.initial_d,
587            weight_decay: config.weight_decay,
588            eps: config.eps,
589            warmup_steps: config.warmup_steps,
590        }
591    }
592}
593
594impl ProdigyBuilder {
595    /// Create a new builder
596    pub fn new() -> Self {
597        Self::default()
598    }
599
600    /// Set parameters
601    pub fn params(mut self, params: Vec<Arc<RwLock<Tensor>>>) -> Self {
602        self.params = params;
603        self
604    }
605
606    /// Set learning rate (use 1.0 for most problems!)
607    pub fn lr(mut self, lr: f32) -> Self {
608        self.lr = lr;
609        self
610    }
611
612    /// Set beta1
613    pub fn beta1(mut self, beta1: f32) -> Self {
614        self.beta1 = beta1;
615        self
616    }
617
618    /// Set beta2
619    pub fn beta2(mut self, beta2: f32) -> Self {
620        self.beta2 = beta2;
621        self
622    }
623
624    /// Set growth rate
625    pub fn growth_rate(mut self, growth_rate: f32) -> Self {
626        self.growth_rate = growth_rate;
627        self
628    }
629
630    /// Set initial d
631    pub fn initial_d(mut self, initial_d: f32) -> Self {
632        self.initial_d = initial_d;
633        self
634    }
635
636    /// Set weight decay
637    pub fn weight_decay(mut self, weight_decay: f32) -> Self {
638        self.weight_decay = weight_decay;
639        self
640    }
641
642    /// Set epsilon
643    pub fn eps(mut self, eps: f32) -> Self {
644        self.eps = eps;
645        self
646    }
647
648    /// Set warmup steps
649    pub fn warmup_steps(mut self, warmup_steps: usize) -> Self {
650        self.warmup_steps = warmup_steps;
651        self
652    }
653
654    /// Build the optimizer
655    pub fn build(self) -> Prodigy {
656        let config = ProdigyConfig {
657            lr: self.lr,
658            beta1: self.beta1,
659            beta2: self.beta2,
660            growth_rate: self.growth_rate,
661            initial_d: self.initial_d,
662            weight_decay: self.weight_decay,
663            eps: self.eps,
664            warmup_steps: self.warmup_steps,
665        };
666        Prodigy::from_config(self.params, config)
667    }
668}
669
670#[cfg(test)]
671mod tests {
672    use super::*;
673    use torsh_tensor::creation::randn;
674
675    #[test]
676    fn test_prodigy_creation() -> OptimizerResult<()> {
677        let param = Arc::new(RwLock::new(randn::<f32>(&[64, 64])?));
678        let params = vec![param];
679
680        let optimizer = Prodigy::new(params, 1.0, 0.9, 0.999, 0.0);
681        assert_eq!(optimizer.lr, 1.0);
682        assert_eq!(optimizer.beta1, 0.9);
683        assert_eq!(optimizer.beta2, 0.999);
684
685        Ok(())
686    }
687
688    #[test]
689    fn test_prodigy_adaptive_lr() -> OptimizerResult<()> {
690        let param = Arc::new(RwLock::new(randn::<f32>(&[32, 32])?));
691        let params = vec![param.clone()];
692
693        let mut optimizer = Prodigy::new(params, 1.0, 0.9, 0.999, 0.0);
694
695        // Perform several steps and check that d adapts
696        let initial_d = optimizer.get_d();
697
698        for _ in 0..20 {
699            let grad = randn::<f32>(&[32, 32])?;
700            param.write().set_grad(Some(grad));
701            optimizer.step()?;
702            optimizer.zero_grad();
703        }
704
705        let final_d = optimizer.get_d();
706
707        // d should have changed (adapted)
708        assert_ne!(initial_d, final_d, "d should adapt during training");
709
710        Ok(())
711    }
712
713    #[test]
714    fn test_prodigy_step() -> OptimizerResult<()> {
715        let param = Arc::new(RwLock::new(randn::<f32>(&[16, 16])?));
716        let params = vec![param.clone()];
717
718        let mut optimizer = Prodigy::new(params, 1.0, 0.9, 0.999, 0.0);
719
720        let grad = randn::<f32>(&[16, 16])?;
721        param.write().set_grad(Some(grad));
722
723        let param_before = param.read().clone();
724        optimizer.step()?;
725        let param_after = param.read().clone();
726
727        // Parameters should change
728        let diff = param_before.sub(&param_after)?;
729        let diff_norm = diff.norm()?.to_vec()?[0];
730        assert!(diff_norm > 0.0, "Parameters should have changed");
731
732        Ok(())
733    }
734
735    #[test]
736    fn test_prodigy_effective_lr() -> OptimizerResult<()> {
737        let param = Arc::new(RwLock::new(randn::<f32>(&[8, 8])?));
738        let params = vec![param.clone()];
739
740        let mut optimizer = Prodigy::new(params, 1.0, 0.9, 0.999, 0.0);
741
742        assert_eq!(optimizer.get_effective_lr(), 0.0); // Before any steps
743
744        for _ in 0..5 {
745            let grad = randn::<f32>(&[8, 8])?;
746            param.write().set_grad(Some(grad));
747            optimizer.step()?;
748            optimizer.zero_grad();
749
750            // Effective LR should be > 0 after steps
751            assert!(optimizer.get_effective_lr() > 0.0);
752        }
753
754        Ok(())
755    }
756
757    #[test]
758    fn test_prodigy_state_dict() -> OptimizerResult<()> {
759        let param = Arc::new(RwLock::new(randn::<f32>(&[16, 16])?));
760        let params = vec![param.clone()];
761
762        let mut optimizer = Prodigy::new(params, 1.0, 0.9, 0.999, 0.0);
763
764        // Perform steps
765        for _ in 0..10 {
766            let grad = randn::<f32>(&[16, 16])?;
767            param.write().set_grad(Some(grad));
768            optimizer.step()?;
769            optimizer.zero_grad();
770        }
771
772        let state = optimizer.state_dict()?;
773        assert_eq!(state.optimizer_type, "Prodigy");
774        assert!(state.global_state.contains_key("d"));
775        assert!(state.global_state.contains_key("s"));
776
777        Ok(())
778    }
779}