1use crate::traits::StatefulOptimizer;
8use crate::{Adam, AdamW, LRScheduler, SGD};
9use serde::{Deserialize, Serialize};
10use std::collections::HashMap;
11use std::sync::{Arc, Mutex};
12use trustformers_core::errors::{Result, TrustformersError};
13use trustformers_core::traits::Optimizer;
14use trustformers_core::Tensor;
15
16#[derive(Debug, Clone, Serialize, Deserialize)]
18pub struct PyTorchParamGroup {
19 pub params: Vec<String>, pub lr: f64,
21 pub weight_decay: f64,
22 pub momentum: Option<f64>,
23 pub dampening: Option<f64>,
24 pub eps: Option<f64>,
25 pub betas: Option<(f64, f64)>,
26 pub alpha: Option<f64>,
27 pub amsgrad: Option<bool>,
28 pub maximize: Option<bool>,
29 pub foreach: Option<bool>,
30 pub differentiable: Option<bool>,
31}
32
33impl Default for PyTorchParamGroup {
34 fn default() -> Self {
35 Self {
36 params: Vec::new(),
37 lr: 0.001,
38 weight_decay: 0.0,
39 momentum: None,
40 dampening: None,
41 eps: Some(1e-8),
42 betas: Some((0.9, 0.999)),
43 alpha: None,
44 amsgrad: Some(false),
45 maximize: Some(false),
46 foreach: None,
47 differentiable: Some(false),
48 }
49 }
50}
51
52#[derive(Debug, Clone, Serialize, Deserialize)]
54pub struct PyTorchOptimizerState {
55 pub state: HashMap<String, serde_json::Value>,
56 pub param_groups: Vec<PyTorchParamGroup>,
57}
58
59#[derive(Debug, Clone, Serialize, Deserialize)]
61pub struct PyTorchOptimizerConfig {
62 pub optimizer_type: String,
63 pub learning_rate: f64,
64 pub betas: (f64, f64),
65 pub epsilon: f64,
66 pub weight_decay: f64,
67 pub amsgrad: bool,
68 pub maximize: bool,
69 pub parameters: HashMap<String, serde_json::Value>,
70}
71
72impl Default for PyTorchOptimizerConfig {
73 fn default() -> Self {
74 Self {
75 optimizer_type: "Adam".to_string(),
76 learning_rate: 1e-3,
77 betas: (0.9, 0.999),
78 epsilon: 1e-8,
79 weight_decay: 0.0,
80 amsgrad: false,
81 maximize: false,
82 parameters: HashMap::new(),
83 }
84 }
85}
86
87pub trait PyTorchOptimizer: Send + Sync {
89 fn param_groups(&self) -> &[PyTorchParamGroup];
91
92 fn param_groups_mut(&mut self) -> &mut [PyTorchParamGroup];
94
95 fn state_dict(&self) -> Result<PyTorchOptimizerState>;
103
104 fn load_state_dict(&mut self, state: PyTorchOptimizerState) -> Result<()>;
106
107 fn step(&mut self, closure: Option<Box<dyn Fn() -> f64>>) -> Result<Option<f64>>;
109
110 fn zero_grad(&mut self, set_to_none: bool) -> Result<()>;
112
113 fn add_param_group(&mut self, param_group: PyTorchParamGroup) -> Result<()>;
115
116 fn defaults(&self) -> PyTorchParamGroup;
118}
119
120const TENSOR_SHAPE_KEY: &str = "shape";
122const TENSOR_DATA_KEY: &str = "data";
124
125fn encode_tensor_state(state: &HashMap<String, Tensor>) -> Result<serde_json::Value> {
136 let mut encoded = serde_json::Map::new();
137 for (name, tensor) in state {
138 let data = tensor.data_f32()?;
139 let mut numbers = Vec::with_capacity(data.len());
140 for value in &data {
141 let number = serde_json::Number::from_f64(*value as f64).ok_or_else(|| {
142 TrustformersError::invalid_input(format!(
143 "optimizer state entry '{name}' contains {value}, which JSON cannot represent"
144 ))
145 })?;
146 numbers.push(serde_json::Value::Number(number));
147 }
148
149 let mut entry = serde_json::Map::new();
150 entry.insert(
151 TENSOR_SHAPE_KEY.to_string(),
152 serde_json::json!(tensor.shape()),
153 );
154 entry.insert(
155 TENSOR_DATA_KEY.to_string(),
156 serde_json::Value::Array(numbers),
157 );
158 encoded.insert(name.clone(), serde_json::Value::Object(entry));
159 }
160 Ok(serde_json::Value::Object(encoded))
161}
162
163fn decode_tensor_state(value: &serde_json::Value) -> Result<HashMap<String, Tensor>> {
171 let object = value.as_object().ok_or_else(|| {
172 TrustformersError::invalid_input(
173 "optimizer state must be a JSON object of tensor entries".to_string(),
174 )
175 })?;
176
177 let mut decoded = HashMap::new();
178 for (name, entry) in object {
179 let entry = entry.as_object().ok_or_else(|| {
180 TrustformersError::invalid_input(format!(
181 "optimizer state entry '{name}' is not an object"
182 ))
183 })?;
184
185 let shape: Vec<usize> = entry
186 .get(TENSOR_SHAPE_KEY)
187 .and_then(|v| v.as_array())
188 .ok_or_else(|| {
189 TrustformersError::invalid_input(format!(
190 "optimizer state entry '{name}' has no '{TENSOR_SHAPE_KEY}' array"
191 ))
192 })?
193 .iter()
194 .map(|v| {
195 v.as_u64().map(|n| n as usize).ok_or_else(|| {
196 TrustformersError::invalid_input(format!(
197 "optimizer state entry '{name}' has a non-integer dimension"
198 ))
199 })
200 })
201 .collect::<Result<Vec<usize>>>()?;
202
203 let data: Vec<f32> = entry
204 .get(TENSOR_DATA_KEY)
205 .and_then(|v| v.as_array())
206 .ok_or_else(|| {
207 TrustformersError::invalid_input(format!(
208 "optimizer state entry '{name}' has no '{TENSOR_DATA_KEY}' array"
209 ))
210 })?
211 .iter()
212 .map(|v| {
213 v.as_f64().map(|n| n as f32).ok_or_else(|| {
214 TrustformersError::invalid_input(format!(
215 "optimizer state entry '{name}' has a non-numeric element"
216 ))
217 })
218 })
219 .collect::<Result<Vec<f32>>>()?;
220
221 let expected: usize = shape.iter().product();
222 if data.len() != expected {
223 return Err(TrustformersError::invalid_input(format!(
224 "optimizer state entry '{name}' has {} elements but shape {shape:?} needs {expected}",
225 data.len()
226 )));
227 }
228
229 decoded.insert(name.clone(), Tensor::from_vec(data, &shape)?);
230 }
231
232 Ok(decoded)
233}
234
235#[derive(Debug)]
237pub struct PyTorchAdam {
238 inner: Adam,
239 param_groups: Vec<PyTorchParamGroup>,
240 parameters: Arc<Mutex<HashMap<String, Tensor>>>,
241 gradients: Arc<Mutex<HashMap<String, Tensor>>>,
242}
243
244impl PyTorchAdam {
245 pub fn new(
247 params: Vec<PyTorchParamGroup>,
248 lr: f64,
249 betas: (f64, f64),
250 eps: f64,
251 weight_decay: f64,
252 _amsgrad: bool,
253 ) -> Result<Self> {
254 let inner = Adam::new(
255 lr as f32,
256 (betas.0 as f32, betas.1 as f32),
257 eps as f32,
258 weight_decay as f32,
259 );
260
261 Ok(Self {
262 inner,
263 param_groups: params,
264 parameters: Arc::new(Mutex::new(HashMap::new())),
265 gradients: Arc::new(Mutex::new(HashMap::new())),
266 })
267 }
268
269 pub fn from_params(params: impl IntoIterator<Item = (String, Tensor)>) -> Result<Self> {
271 let param_group = PyTorchParamGroup {
272 params: params.into_iter().map(|(name, _)| name).collect(),
273 ..Default::default()
274 };
275
276 Self::new(vec![param_group], 0.001, (0.9, 0.999), 1e-8, 0.0, false)
277 }
278
279 pub fn from_config(config: PyTorchOptimizerConfig) -> Result<Self> {
281 let param_group = PyTorchParamGroup {
283 params: config.parameters.keys().cloned().collect(),
284 lr: config.learning_rate,
285 weight_decay: config.weight_decay,
286 eps: Some(config.epsilon),
287 betas: Some(config.betas),
288 amsgrad: Some(config.amsgrad),
289 maximize: Some(config.maximize),
290 ..Default::default()
291 };
292
293 Self::new(
294 vec![param_group],
295 config.learning_rate,
296 config.betas,
297 config.epsilon,
298 config.weight_decay,
299 config.amsgrad,
300 )
301 }
302
303 pub fn from_cross_framework_config(
305 config: crate::cross_framework::PyTorchOptimizerConfig,
306 ) -> Result<Self> {
307 let betas = if let Some(betas_val) = config.parameters.get("betas") {
309 if let Some(arr) = betas_val.as_array() {
310 (
311 arr[0].as_f64().unwrap_or(0.9),
312 arr[1].as_f64().unwrap_or(0.999),
313 )
314 } else {
315 (0.9, 0.999)
316 }
317 } else {
318 (0.9, 0.999)
319 };
320
321 let epsilon = config.parameters.get("epsilon").and_then(|v| v.as_f64()).unwrap_or(1e-8);
322
323 let weight_decay =
324 config.parameters.get("weight_decay").and_then(|v| v.as_f64()).unwrap_or(0.0);
325
326 let amsgrad = config.parameters.get("amsgrad").and_then(|v| v.as_bool()).unwrap_or(false);
327
328 let param_group = PyTorchParamGroup {
330 params: Vec::new(),
331 lr: config.learning_rate as f64,
332 weight_decay,
333 eps: Some(epsilon),
334 betas: Some(betas),
335 amsgrad: Some(amsgrad),
336 maximize: Some(false),
337 ..Default::default()
338 };
339
340 Self::new(
341 vec![param_group],
342 config.learning_rate as f64,
343 betas,
344 epsilon,
345 weight_decay,
346 amsgrad,
347 )
348 }
349
350 pub fn register_param(&mut self, name: String, param: Tensor) -> Result<()> {
352 let mut params = self
353 .parameters
354 .lock()
355 .map_err(|_| TrustformersError::runtime_error("Mutex lock poisoned".into()))?;
356 params.insert(name, param);
357 Ok(())
358 }
359
360 pub fn set_grad(&mut self, name: String, grad: Tensor) -> Result<()> {
362 let mut grads = self
363 .gradients
364 .lock()
365 .map_err(|_| TrustformersError::runtime_error("Mutex lock poisoned".into()))?;
366 grads.insert(name, grad);
367 Ok(())
368 }
369}
370
371impl PyTorchOptimizer for PyTorchAdam {
372 fn param_groups(&self) -> &[PyTorchParamGroup] {
373 &self.param_groups
374 }
375
376 fn param_groups_mut(&mut self) -> &mut [PyTorchParamGroup] {
377 &mut self.param_groups
378 }
379
380 fn state_dict(&self) -> Result<PyTorchOptimizerState> {
381 let inner_state = StatefulOptimizer::state_dict(&self.inner)?;
384 Ok(PyTorchOptimizerState {
385 state: [(
386 String::from("adam_state"),
387 encode_tensor_state(&inner_state)?,
388 )]
389 .into(),
390 param_groups: self.param_groups.clone(),
391 })
392 }
393
394 fn load_state_dict(&mut self, state: PyTorchOptimizerState) -> Result<()> {
395 self.param_groups = state.param_groups;
396
397 let raw = state.state.get("adam_state").ok_or_else(|| {
398 TrustformersError::invalid_input("checkpoint has no 'adam_state' entry".to_string())
399 })?;
400 let decoded = decode_tensor_state(raw)?;
401 StatefulOptimizer::load_state_dict(&mut self.inner, decoded)?;
402 Ok(())
403 }
404
405 fn step(&mut self, closure: Option<Box<dyn Fn() -> f64>>) -> Result<Option<f64>> {
406 let loss = closure.map(|closure_fn| closure_fn());
407
408 for group in &self.param_groups {
410 for param_name in &group.params {
411 let param_copy = {
413 let params = self.parameters.lock().map_err(|_| {
414 TrustformersError::runtime_error("Mutex lock poisoned".into())
415 })?;
416 params.get(param_name).cloned()
417 };
418 let grad_copy = {
419 let grads = self.gradients.lock().map_err(|_| {
420 TrustformersError::runtime_error("Mutex lock poisoned".into())
421 })?;
422 grads.get(param_name).cloned()
423 };
424
425 if let (Some(mut param), Some(grad)) = (param_copy, grad_copy) {
426 self.inner.update_named(param_name, &mut param, &grad)?;
431
432 let mut params = self.parameters.lock().map_err(|_| {
434 TrustformersError::runtime_error("Mutex lock poisoned".into())
435 })?;
436 params.insert(param_name.clone(), param);
437 }
438 }
439 }
440
441 Optimizer::step(&mut self.inner);
443
444 Ok(loss)
445 }
446
447 fn zero_grad(&mut self, _set_to_none: bool) -> Result<()> {
448 let mut grads = self
449 .gradients
450 .lock()
451 .map_err(|_| TrustformersError::runtime_error("Mutex lock poisoned".into()))?;
452 grads.clear();
453 Ok(())
454 }
455
456 fn add_param_group(&mut self, param_group: PyTorchParamGroup) -> Result<()> {
457 self.param_groups.push(param_group);
458 Ok(())
459 }
460
461 fn defaults(&self) -> PyTorchParamGroup {
462 PyTorchParamGroup {
463 lr: 0.001,
464 betas: Some((0.9, 0.999)),
465 eps: Some(1e-8),
466 weight_decay: 0.0,
467 amsgrad: Some(false),
468 ..Default::default()
469 }
470 }
471}
472
473#[derive(Debug)]
475pub struct PyTorchAdamW {
476 inner: AdamW,
477 param_groups: Vec<PyTorchParamGroup>,
478 parameters: Arc<Mutex<HashMap<String, Tensor>>>,
479 gradients: Arc<Mutex<HashMap<String, Tensor>>>,
480}
481
482impl PyTorchAdamW {
483 pub fn new(
485 params: Vec<PyTorchParamGroup>,
486 lr: f64,
487 betas: (f64, f64),
488 eps: f64,
489 weight_decay: f64,
490 _amsgrad: bool,
491 ) -> Result<Self> {
492 let inner = AdamW::new(
493 lr as f32,
494 (betas.0 as f32, betas.1 as f32),
495 eps as f32,
496 weight_decay as f32,
497 );
498
499 Ok(Self {
500 inner,
501 param_groups: params,
502 parameters: Arc::new(Mutex::new(HashMap::new())),
503 gradients: Arc::new(Mutex::new(HashMap::new())),
504 })
505 }
506
507 pub fn from_params(params: impl IntoIterator<Item = (String, Tensor)>) -> Result<Self> {
509 let param_group = PyTorchParamGroup {
510 params: params.into_iter().map(|(name, _)| name).collect(),
511 ..Default::default()
512 };
513
514 Self::new(vec![param_group], 0.001, (0.9, 0.999), 1e-8, 0.01, false)
515 }
516
517 pub fn register_param(&mut self, name: String, param: Tensor) -> Result<()> {
519 let mut params = self
520 .parameters
521 .lock()
522 .map_err(|_| TrustformersError::runtime_error("Mutex lock poisoned".into()))?;
523 params.insert(name, param);
524 Ok(())
525 }
526
527 pub fn set_grad(&mut self, name: String, grad: Tensor) -> Result<()> {
529 let mut grads = self
530 .gradients
531 .lock()
532 .map_err(|_| TrustformersError::runtime_error("Mutex lock poisoned".into()))?;
533 grads.insert(name, grad);
534 Ok(())
535 }
536}
537
538impl PyTorchOptimizer for PyTorchAdamW {
539 fn param_groups(&self) -> &[PyTorchParamGroup] {
540 &self.param_groups
541 }
542
543 fn param_groups_mut(&mut self) -> &mut [PyTorchParamGroup] {
544 &mut self.param_groups
545 }
546
547 fn state_dict(&self) -> Result<PyTorchOptimizerState> {
548 let inner_state = StatefulOptimizer::state_dict(&self.inner)?;
551 Ok(PyTorchOptimizerState {
552 state: [(
553 String::from("adamw_state"),
554 encode_tensor_state(&inner_state)?,
555 )]
556 .into(),
557 param_groups: self.param_groups.clone(),
558 })
559 }
560
561 fn load_state_dict(&mut self, state: PyTorchOptimizerState) -> Result<()> {
562 self.param_groups = state.param_groups;
563
564 let raw = state.state.get("adamw_state").ok_or_else(|| {
565 TrustformersError::invalid_input("checkpoint has no 'adamw_state' entry".to_string())
566 })?;
567 let decoded = decode_tensor_state(raw)?;
568 StatefulOptimizer::load_state_dict(&mut self.inner, decoded)?;
569 Ok(())
570 }
571
572 fn step(&mut self, closure: Option<Box<dyn Fn() -> f64>>) -> Result<Option<f64>> {
573 let loss = closure.map(|closure_fn| closure_fn());
574
575 for group in &self.param_groups {
576 for param_name in &group.params {
577 let param_copy = {
579 let params = self.parameters.lock().map_err(|_| {
580 TrustformersError::runtime_error("Mutex lock poisoned".into())
581 })?;
582 params.get(param_name).cloned()
583 };
584 let grad_copy = {
585 let grads = self.gradients.lock().map_err(|_| {
586 TrustformersError::runtime_error("Mutex lock poisoned".into())
587 })?;
588 grads.get(param_name).cloned()
589 };
590
591 if let (Some(mut param), Some(grad)) = (param_copy, grad_copy) {
592 self.inner.update_named(param_name, &mut param, &grad)?;
597
598 let mut params = self.parameters.lock().map_err(|_| {
600 TrustformersError::runtime_error("Mutex lock poisoned".into())
601 })?;
602 params.insert(param_name.clone(), param);
603 }
604 }
605 }
606
607 Optimizer::step(&mut self.inner);
609
610 Ok(loss)
611 }
612
613 fn zero_grad(&mut self, _set_to_none: bool) -> Result<()> {
614 let mut grads = self
615 .gradients
616 .lock()
617 .map_err(|_| TrustformersError::runtime_error("Mutex lock poisoned".into()))?;
618 grads.clear();
619 Ok(())
620 }
621
622 fn add_param_group(&mut self, param_group: PyTorchParamGroup) -> Result<()> {
623 self.param_groups.push(param_group);
624 Ok(())
625 }
626
627 fn defaults(&self) -> PyTorchParamGroup {
628 PyTorchParamGroup {
629 lr: 0.001,
630 betas: Some((0.9, 0.999)),
631 eps: Some(1e-8),
632 weight_decay: 0.01,
633 amsgrad: Some(false),
634 ..Default::default()
635 }
636 }
637}
638
639#[derive(Debug)]
641pub struct PyTorchSGD {
642 inner: SGD,
643 param_groups: Vec<PyTorchParamGroup>,
644 parameters: Arc<Mutex<HashMap<String, Tensor>>>,
645 gradients: Arc<Mutex<HashMap<String, Tensor>>>,
646}
647
648impl PyTorchSGD {
649 pub fn new(
651 params: Vec<PyTorchParamGroup>,
652 lr: f64,
653 momentum: f64,
654 dampening: f64,
655 weight_decay: f64,
656 nesterov: bool,
657 ) -> Result<Self> {
658 let config = crate::sgd::SGDConfig {
659 lr: lr as f32,
660 momentum: momentum as f32,
661 dampening: dampening as f32,
662 weight_decay: weight_decay as f32,
663 nesterov,
664 };
665
666 let inner = SGD::from_config(config);
667
668 Ok(Self {
669 inner,
670 param_groups: params,
671 parameters: Arc::new(Mutex::new(HashMap::new())),
672 gradients: Arc::new(Mutex::new(HashMap::new())),
673 })
674 }
675
676 pub fn from_params(params: impl IntoIterator<Item = (String, Tensor)>) -> Result<Self> {
678 let param_group = PyTorchParamGroup {
679 params: params.into_iter().map(|(name, _)| name).collect(),
680 lr: 0.01,
681 momentum: Some(0.0),
682 dampening: Some(0.0),
683 weight_decay: 0.0,
684 ..Default::default()
685 };
686
687 Self::new(vec![param_group], 0.01, 0.0, 0.0, 0.0, false)
688 }
689
690 pub fn register_param(&mut self, name: String, param: Tensor) -> Result<()> {
692 let mut params = self
693 .parameters
694 .lock()
695 .map_err(|_| TrustformersError::runtime_error("Mutex lock poisoned".into()))?;
696 params.insert(name, param);
697 Ok(())
698 }
699
700 pub fn set_grad(&mut self, name: String, grad: Tensor) -> Result<()> {
702 let mut grads = self
703 .gradients
704 .lock()
705 .map_err(|_| TrustformersError::runtime_error("Mutex lock poisoned".into()))?;
706 grads.insert(name, grad);
707 Ok(())
708 }
709}
710
711impl PyTorchOptimizer for PyTorchSGD {
712 fn param_groups(&self) -> &[PyTorchParamGroup] {
713 &self.param_groups
714 }
715
716 fn param_groups_mut(&mut self) -> &mut [PyTorchParamGroup] {
717 &mut self.param_groups
718 }
719
720 fn state_dict(&self) -> Result<PyTorchOptimizerState> {
721 let inner_state = StatefulOptimizer::state_dict(&self.inner)?;
724 Ok(PyTorchOptimizerState {
725 state: [(
726 String::from("sgd_state"),
727 encode_tensor_state(&inner_state)?,
728 )]
729 .into(),
730 param_groups: self.param_groups.clone(),
731 })
732 }
733
734 fn load_state_dict(&mut self, state: PyTorchOptimizerState) -> Result<()> {
735 self.param_groups = state.param_groups;
736
737 let raw = state.state.get("sgd_state").ok_or_else(|| {
738 TrustformersError::invalid_input("checkpoint has no 'sgd_state' entry".to_string())
739 })?;
740 let decoded = decode_tensor_state(raw)?;
741 StatefulOptimizer::load_state_dict(&mut self.inner, decoded)?;
742 Ok(())
743 }
744
745 fn step(&mut self, closure: Option<Box<dyn Fn() -> f64>>) -> Result<Option<f64>> {
746 let loss = closure.map(|closure_fn| closure_fn());
747
748 for group in &self.param_groups {
749 for param_name in &group.params {
750 let param_copy = {
752 let params = self.parameters.lock().map_err(|_| {
753 TrustformersError::runtime_error("Mutex lock poisoned".into())
754 })?;
755 params.get(param_name).cloned()
756 };
757 let grad_copy = {
758 let grads = self.gradients.lock().map_err(|_| {
759 TrustformersError::runtime_error("Mutex lock poisoned".into())
760 })?;
761 grads.get(param_name).cloned()
762 };
763
764 if let (Some(mut param), Some(grad)) = (param_copy, grad_copy) {
765 self.inner.update_named(param_name, &mut param, &grad)?;
770
771 let mut params = self.parameters.lock().map_err(|_| {
773 TrustformersError::runtime_error("Mutex lock poisoned".into())
774 })?;
775 params.insert(param_name.clone(), param);
776 }
777 }
778 }
779
780 Optimizer::step(&mut self.inner);
782
783 Ok(loss)
784 }
785
786 fn zero_grad(&mut self, _set_to_none: bool) -> Result<()> {
787 let mut grads = self
788 .gradients
789 .lock()
790 .map_err(|_| TrustformersError::runtime_error("Mutex lock poisoned".into()))?;
791 grads.clear();
792 Ok(())
793 }
794
795 fn add_param_group(&mut self, param_group: PyTorchParamGroup) -> Result<()> {
796 self.param_groups.push(param_group);
797 Ok(())
798 }
799
800 fn defaults(&self) -> PyTorchParamGroup {
801 PyTorchParamGroup {
802 lr: 0.01,
803 momentum: Some(0.0),
804 dampening: Some(0.0),
805 weight_decay: 0.0,
806 ..Default::default()
807 }
808 }
809}
810
811pub struct PyTorchOptimizerFactory;
813
814impl PyTorchOptimizerFactory {
815 pub fn adam(
817 params: impl IntoIterator<Item = (String, Tensor)>,
818 lr: f64,
819 betas: (f64, f64),
820 eps: f64,
821 weight_decay: f64,
822 amsgrad: bool,
823 ) -> Result<PyTorchAdam> {
824 let param_group = PyTorchParamGroup {
825 params: params.into_iter().map(|(name, _)| name).collect(),
826 lr,
827 betas: Some(betas),
828 eps: Some(eps),
829 weight_decay,
830 amsgrad: Some(amsgrad),
831 ..Default::default()
832 };
833
834 PyTorchAdam::new(vec![param_group], lr, betas, eps, weight_decay, amsgrad)
835 }
836
837 pub fn adamw(
839 params: impl IntoIterator<Item = (String, Tensor)>,
840 lr: f64,
841 betas: (f64, f64),
842 eps: f64,
843 weight_decay: f64,
844 amsgrad: bool,
845 ) -> Result<PyTorchAdamW> {
846 let param_group = PyTorchParamGroup {
847 params: params.into_iter().map(|(name, _)| name).collect(),
848 lr,
849 betas: Some(betas),
850 eps: Some(eps),
851 weight_decay,
852 amsgrad: Some(amsgrad),
853 ..Default::default()
854 };
855
856 PyTorchAdamW::new(vec![param_group], lr, betas, eps, weight_decay, amsgrad)
857 }
858
859 pub fn sgd(
861 params: impl IntoIterator<Item = (String, Tensor)>,
862 lr: f64,
863 momentum: f64,
864 dampening: f64,
865 weight_decay: f64,
866 nesterov: bool,
867 ) -> Result<PyTorchSGD> {
868 let param_group = PyTorchParamGroup {
869 params: params.into_iter().map(|(name, _)| name).collect(),
870 lr,
871 momentum: Some(momentum),
872 dampening: Some(dampening),
873 weight_decay,
874 ..Default::default()
875 };
876
877 PyTorchSGD::new(
878 vec![param_group],
879 lr,
880 momentum,
881 dampening,
882 weight_decay,
883 nesterov,
884 )
885 }
886}
887
888pub struct PyTorchLRScheduler {
890 inner_scheduler: Box<dyn LRScheduler>,
891 optimizer: Box<dyn PyTorchOptimizer>,
892 last_epoch: i64,
893}
894
895impl PyTorchLRScheduler {
896 pub fn new(optimizer: Box<dyn PyTorchOptimizer>, scheduler: Box<dyn LRScheduler>) -> Self {
898 Self {
899 inner_scheduler: scheduler,
900 optimizer,
901 last_epoch: -1,
902 }
903 }
904
905 pub fn step(&mut self, epoch: Option<i64>) -> Result<()> {
907 let current_epoch = epoch.unwrap_or(self.last_epoch + 1);
908 self.last_epoch = current_epoch;
909
910 let new_lr = self.inner_scheduler.get_lr(current_epoch as usize);
911
912 for group in self.optimizer.param_groups_mut() {
914 group.lr = new_lr as f64;
915 }
916
917 Ok(())
918 }
919
920 pub fn get_last_lr(&self) -> f64 {
922 self.inner_scheduler.get_lr(self.last_epoch.max(0) as usize) as f64
923 }
924
925 pub fn state_dict(&self) -> serde_json::Value {
927 serde_json::json!({
928 "last_epoch": self.last_epoch,
929 "scheduler_state": "serialized_state" })
931 }
932
933 pub fn load_state_dict(&mut self, state: serde_json::Value) -> Result<()> {
935 if let Some(epoch) = state.get("last_epoch").and_then(|e| e.as_i64()) {
936 self.last_epoch = epoch;
937 }
938 Ok(())
939 }
940}
941
942#[cfg(test)]
943mod tests {
944 use super::*;
945 use trustformers_core::Tensor;
946
947 #[test]
948 fn test_pytorch_adam_creation() {
949 let params = vec![
950 (
951 "param1".to_string(),
952 Tensor::zeros(&[10, 10]).expect("Failed to create tensor"),
953 ),
954 (
955 "param2".to_string(),
956 Tensor::zeros(&[5, 5]).expect("Failed to create tensor"),
957 ),
958 ];
959
960 let optimizer =
961 PyTorchAdam::from_params(params).expect("Failed to create optimizer from params");
962 assert_eq!(optimizer.param_groups().len(), 1);
963 assert_eq!(optimizer.param_groups()[0].params.len(), 2);
964 }
965
966 #[test]
967 fn test_pytorch_adamw_creation() {
968 let params = vec![(
969 "param1".to_string(),
970 Tensor::zeros(&[10, 10]).expect("Failed to create tensor"),
971 )];
972
973 let optimizer =
974 PyTorchAdamW::from_params(params).expect("Failed to create optimizer from params");
975 assert_eq!(optimizer.param_groups().len(), 1);
976 assert_eq!(optimizer.defaults().weight_decay, 0.01);
977 }
978
979 #[test]
980 fn test_pytorch_sgd_creation() {
981 let params = vec![(
982 "param1".to_string(),
983 Tensor::zeros(&[10, 10]).expect("Failed to create tensor"),
984 )];
985
986 let optimizer =
987 PyTorchSGD::from_params(params).expect("Failed to create optimizer from params");
988 assert_eq!(optimizer.param_groups().len(), 1);
989 assert_eq!(optimizer.defaults().lr, 0.01);
990 }
991
992 #[test]
993 fn test_pytorch_optimizer_factory() {
994 let params = vec![(
995 "param1".to_string(),
996 Tensor::zeros(&[10, 10]).expect("Failed to create tensor"),
997 )];
998
999 let adam =
1000 PyTorchOptimizerFactory::adam(params.clone(), 0.001, (0.9, 0.999), 1e-8, 0.0, false)
1001 .expect("Operation failed in test");
1002 assert_eq!(adam.param_groups()[0].lr, 0.001);
1003
1004 let adamw =
1005 PyTorchOptimizerFactory::adamw(params.clone(), 0.001, (0.9, 0.999), 1e-8, 0.01, false)
1006 .expect("Operation failed in test");
1007 assert_eq!(adamw.param_groups()[0].weight_decay, 0.01);
1008
1009 let sgd = PyTorchOptimizerFactory::sgd(params, 0.01, 0.9, 0.0, 0.0, false)
1010 .expect("Operation failed in test");
1011 assert_eq!(sgd.param_groups()[0].momentum, Some(0.9));
1012 }
1013
1014 #[test]
1015 fn test_param_group_operations() {
1016 let params = vec![(
1017 "param1".to_string(),
1018 Tensor::zeros(&[10, 10]).expect("Failed to create tensor"),
1019 )];
1020
1021 let mut optimizer =
1022 PyTorchAdam::from_params(params).expect("Failed to create optimizer from params");
1023
1024 let new_group = PyTorchParamGroup {
1025 params: vec!["param2".to_string()],
1026 lr: 0.002,
1027 ..Default::default()
1028 };
1029
1030 optimizer.add_param_group(new_group).expect("Failed to add param group");
1031 assert_eq!(optimizer.param_groups().len(), 2);
1032 assert_eq!(optimizer.param_groups()[1].lr, 0.002);
1033 }
1034
1035 #[test]
1036 fn test_state_dict_operations() {
1037 let params = vec![(
1038 "param1".to_string(),
1039 Tensor::zeros(&[10, 10]).expect("Failed to create tensor"),
1040 )];
1041
1042 let optimizer =
1043 PyTorchAdam::from_params(params).expect("Failed to create optimizer from params");
1044 let state_dict = optimizer.state_dict().expect("state_dict");
1045
1046 assert_eq!(state_dict.param_groups.len(), 1);
1047 assert!(state_dict.state.contains_key("adam_state"));
1048 }
1049
1050 #[test]
1051 fn test_zero_grad() {
1052 let params = vec![(
1053 "param1".to_string(),
1054 Tensor::zeros(&[10, 10]).expect("Failed to create tensor"),
1055 )];
1056
1057 let mut optimizer =
1058 PyTorchAdam::from_params(params).expect("Failed to create optimizer from params");
1059 optimizer
1060 .set_grad(
1061 "param1".to_string(),
1062 Tensor::ones(&[10, 10]).expect("Failed to create tensor"),
1063 )
1064 .expect("Operation failed in test");
1065
1066 assert_eq!(
1068 optimizer.gradients.lock().expect("Mutex lock poisoned").len(),
1069 1
1070 );
1071
1072 optimizer.zero_grad(false).expect("Zero grad failed");
1074 assert_eq!(
1075 optimizer.gradients.lock().expect("Mutex lock poisoned").len(),
1076 0
1077 );
1078 }
1079
1080 #[test]
1086 fn state_dict_round_trip_reproduces_the_trajectory() {
1087 fn build() -> PyTorchAdam {
1088 let params = vec![(
1089 "w".to_string(),
1090 Tensor::from_vec(vec![1.0_f32, 2.0], &[2]).expect("tensor"),
1091 )];
1092 PyTorchAdam::from_params(params).expect("optimizer")
1093 }
1094
1095 fn drive(optimizer: &mut PyTorchAdam, steps: usize) {
1096 let grad = Tensor::from_vec(vec![0.5_f32, -0.5], &[2]).expect("grad");
1097 for _ in 0..steps {
1098 optimizer.set_grad("w".to_string(), grad.clone()).expect("grad");
1099 optimizer.step(None).expect("step");
1100 }
1101 }
1102
1103 fn value_of(optimizer: &PyTorchAdam) -> Vec<f32> {
1104 optimizer
1105 .parameters
1106 .lock()
1107 .expect("registry")
1108 .get("w")
1109 .expect("parameter")
1110 .data_f32()
1111 .expect("data")
1112 }
1113
1114 let mut original = build();
1115 original
1116 .register_param(
1117 "w".to_string(),
1118 Tensor::from_vec(vec![1.0_f32, 2.0], &[2]).expect("tensor"),
1119 )
1120 .expect("register");
1121 drive(&mut original, 3);
1122
1123 let checkpoint = original.state_dict().expect("state_dict");
1124 let entries = checkpoint
1125 .state
1126 .get("adam_state")
1127 .and_then(|v| v.as_object())
1128 .expect("adam_state object");
1129 assert!(
1130 entries.keys().any(|k| k.starts_with("exp_avg_")),
1131 "the moment buffers must be checkpointed, found {:?}",
1132 entries.keys().collect::<Vec<_>>()
1133 );
1134
1135 let resume_point = value_of(&original);
1136
1137 let mut resumed = build();
1138 resumed.load_state_dict(checkpoint).expect("load_state_dict");
1139 assert_eq!(
1140 resumed.parameters.lock().expect("registry").len(),
1141 0,
1142 "loading optimizer state must not inject buffers into the parameter map"
1143 );
1144 resumed
1145 .register_param(
1146 "w".to_string(),
1147 Tensor::from_vec(resume_point.clone(), &[2]).expect("tensor"),
1148 )
1149 .expect("register");
1150
1151 drive(&mut original, 1);
1152 drive(&mut resumed, 1);
1153
1154 let expected = value_of(&original);
1155 let actual = value_of(&resumed);
1156 for (a, b) in actual.iter().zip(expected.iter()) {
1157 assert!(
1158 (a - b).abs() < 1e-6,
1159 "resume diverged from the uninterrupted run: {a} vs {b}"
1160 );
1161 }
1162 assert!(
1163 actual.iter().zip(resume_point.iter()).any(|(a, b)| (a - b).abs() > 1e-9),
1164 "the post-resume step must actually move the parameter"
1165 );
1166 }
1167
1168 #[test]
1170 fn malformed_checkpoint_is_rejected() {
1171 let params = vec![("w".to_string(), Tensor::zeros(&[2]).expect("tensor"))];
1172 let mut optimizer = PyTorchAdam::from_params(params).expect("optimizer");
1173
1174 let bogus = PyTorchOptimizerState {
1175 state: [(
1176 "adam_state".to_string(),
1177 serde_json::json!({"exp_avg_p:0": {"shape": [4], "data": [1.0, 2.0]}}),
1178 )]
1179 .into(),
1180 param_groups: Vec::new(),
1181 };
1182 assert!(
1183 optimizer.load_state_dict(bogus).is_err(),
1184 "a shape/payload mismatch must be an error"
1185 );
1186
1187 let missing = PyTorchOptimizerState {
1188 state: HashMap::new(),
1189 param_groups: Vec::new(),
1190 };
1191 assert!(
1192 optimizer.load_state_dict(missing).is_err(),
1193 "a checkpoint with no optimizer state must be an error"
1194 );
1195 }
1196}