1use 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#[derive(Debug, Clone)]
115pub struct ProdigyConfig {
116 pub lr: f32,
118 pub beta1: f32,
120 pub beta2: f32,
122 pub growth_rate: f32,
124 pub initial_d: f32,
126 pub weight_decay: f32,
128 pub eps: f32,
130 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
149pub struct Prodigy {
154 param_groups: Vec<ParamGroup>,
156 lr: f32,
158 beta1: f32,
160 beta2: f32,
162 growth_rate: f32,
164 d: f32,
166 d_prev: f32,
168 s: f32,
170 s_prev: f32,
172 weight_decay: f32,
174 eps: f32,
176 warmup_steps: usize,
178 momentum: HashMap<String, Tensor>,
180 variance: HashMap<String, Tensor>,
182 prev_params: HashMap<String, Tensor>,
184 step_count: usize,
186}
187
188impl Prodigy {
189 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 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 pub fn builder() -> ProdigyBuilder {
265 ProdigyBuilder::default()
266 }
267
268 pub fn get_d(&self) -> f32 {
270 self.d
271 }
272
273 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 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 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 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 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 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 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 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 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 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 let new_param = param_data
393 .sub(&update)
394 .map_err(|e| OptimizerError::TensorError(e))?;
395
396 if let Some(prev_param) = self.prev_params.get(¶m_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 self.prev_params.insert(param_key.clone(), param_data);
411
412 *m_entry = new_m;
414 *v_entry = new_v;
415 *param_guard = new_param;
416 }
417 }
418
419 self.s_prev = self.s;
421 self.s += distance_sum;
422
423 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 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 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 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#[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 pub fn new() -> Self {
597 Self::default()
598 }
599
600 pub fn params(mut self, params: Vec<Arc<RwLock<Tensor>>>) -> Self {
602 self.params = params;
603 self
604 }
605
606 pub fn lr(mut self, lr: f32) -> Self {
608 self.lr = lr;
609 self
610 }
611
612 pub fn beta1(mut self, beta1: f32) -> Self {
614 self.beta1 = beta1;
615 self
616 }
617
618 pub fn beta2(mut self, beta2: f32) -> Self {
620 self.beta2 = beta2;
621 self
622 }
623
624 pub fn growth_rate(mut self, growth_rate: f32) -> Self {
626 self.growth_rate = growth_rate;
627 self
628 }
629
630 pub fn initial_d(mut self, initial_d: f32) -> Self {
632 self.initial_d = initial_d;
633 self
634 }
635
636 pub fn weight_decay(mut self, weight_decay: f32) -> Self {
638 self.weight_decay = weight_decay;
639 self
640 }
641
642 pub fn eps(mut self, eps: f32) -> Self {
644 self.eps = eps;
645 self
646 }
647
648 pub fn warmup_steps(mut self, warmup_steps: usize) -> Self {
650 self.warmup_steps = warmup_steps;
651 self
652 }
653
654 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 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 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 let diff = param_before.sub(¶m_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); 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 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 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}