Skip to main content

trustformers_optim/
pytorch_compat.rs

1//! PyTorch Optimizer API Compatibility Layer
2//!
3//! This module provides PyTorch-compatible optimizer interfaces for seamless
4//! integration with PyTorch-based training workflows. It wraps our native
5//! optimizers to provide the familiar PyTorch API while maintaining high performance.
6
7use 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/// PyTorch-compatible optimizer parameter group
17#[derive(Debug, Clone, Serialize, Deserialize)]
18pub struct PyTorchParamGroup {
19    pub params: Vec<String>, // Parameter names/IDs
20    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/// PyTorch-compatible optimizer state
53#[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/// PyTorch-compatible optimizer configuration
60#[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
87/// PyTorch-compatible optimizer interface
88pub trait PyTorchOptimizer: Send + Sync {
89    /// Get parameter groups
90    fn param_groups(&self) -> &[PyTorchParamGroup];
91
92    /// Get mutable parameter groups
93    fn param_groups_mut(&mut self) -> &mut [PyTorchParamGroup];
94
95    /// Get optimizer state
96    /// Serialises the optimizer state.
97    ///
98    /// # Errors
99    ///
100    /// Returns an error when the inner optimizer's state cannot be serialised;
101    /// silently emitting an empty state would make a checkpoint look valid.
102    fn state_dict(&self) -> Result<PyTorchOptimizerState>;
103
104    /// Load optimizer state
105    fn load_state_dict(&mut self, state: PyTorchOptimizerState) -> Result<()>;
106
107    /// Perform optimization step
108    fn step(&mut self, closure: Option<Box<dyn Fn() -> f64>>) -> Result<Option<f64>>;
109
110    /// Zero gradients
111    fn zero_grad(&mut self, set_to_none: bool) -> Result<()>;
112
113    /// Add parameter group
114    fn add_param_group(&mut self, param_group: PyTorchParamGroup) -> Result<()>;
115
116    /// Get defaults
117    fn defaults(&self) -> PyTorchParamGroup;
118}
119
120/// JSON key under which a tensor's logical shape is stored.
121const TENSOR_SHAPE_KEY: &str = "shape";
122/// JSON key under which a tensor's flattened `f32` payload is stored.
123const TENSOR_DATA_KEY: &str = "data";
124
125/// Encodes a [`StatefulOptimizer`] tensor state dict as JSON.
126///
127/// Each entry becomes `{"shape": [...], "data": [...]}` so the round trip is lossless
128/// for `f32` payloads and, crucially, carries the *real* optimizer moments rather than
129/// a summary of them.
130///
131/// # Errors
132///
133/// Returns an error when a tensor cannot be read as `f32` or a value is not
134/// representable in JSON (NaN / infinity).
135fn 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
163/// Decodes the JSON produced by [`encode_tensor_state`].
164///
165/// # Errors
166///
167/// Returns an error for any structural problem — a missing shape, a payload whose
168/// length disagrees with the shape, or a non-object value. A malformed checkpoint must
169/// never load "successfully" with nothing restored.
170fn 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/// PyTorch-compatible Adam optimizer
236#[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    /// Create new PyTorch-compatible Adam optimizer
246    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    /// Create with default parameters
270    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    /// Create PyTorch Adam optimizer from configuration
280    pub fn from_config(config: PyTorchOptimizerConfig) -> Result<Self> {
281        // Create parameter group from config
282        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    /// Create PyTorch Adam optimizer from cross-framework configuration
304    pub fn from_cross_framework_config(
305        config: crate::cross_framework::PyTorchOptimizerConfig,
306    ) -> Result<Self> {
307        // Extract parameters from the HashMap
308        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        // Create parameter group from config
329        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    /// Register parameter
351    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    /// Set gradient for parameter
361    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        // Route through the inner optimizer's own checkpoint format so the moment
382        // buffers really are written out.
383        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        // Apply gradients to parameters using the inner optimizer
409        for group in &self.param_groups {
410            for param_name in &group.params {
411                // Get copies of parameter and gradient to avoid borrow conflicts
412                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                    // Use the *named* identity: this API already carries parameter
427                    // names, and the tensor is cloned out of the registry on every
428                    // step, so an address-derived key would allocate a fresh state
429                    // slot each time.
430                    self.inner.update_named(param_name, &mut param, &grad)?;
431
432                    // Store updated parameter back
433                    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        // Advance the inner optimizer's global step so bias correction progresses.
442        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/// PyTorch-compatible AdamW optimizer
474#[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    /// Create new PyTorch-compatible AdamW optimizer
484    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    /// Create with default parameters
508    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    /// Register parameter
518    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    /// Set gradient for parameter
528    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        // Route through the inner optimizer's own checkpoint format so the moment
549        // buffers really are written out.
550        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                // Get copies of parameter and gradient to avoid borrow conflicts
578                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                    // Use the *named* identity: this API already carries parameter
593                    // names, and the tensor is cloned out of the registry on every
594                    // step, so an address-derived key would allocate a fresh state
595                    // slot each time.
596                    self.inner.update_named(param_name, &mut param, &grad)?;
597
598                    // Store updated parameter back
599                    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        // Advance the inner optimizer's global step so bias correction progresses.
608        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/// PyTorch-compatible SGD optimizer
640#[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    /// Create new PyTorch-compatible SGD optimizer
650    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    /// Create with default parameters
677    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    /// Register parameter
691    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    /// Set gradient for parameter
701    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        // Route through the inner optimizer's own checkpoint format so the moment
722        // buffers really are written out.
723        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                // Get copies of parameter and gradient to avoid borrow conflicts
751                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                    // Use the *named* identity: this API already carries parameter
766                    // names, and the tensor is cloned out of the registry on every
767                    // step, so an address-derived key would allocate a fresh state
768                    // slot each time.
769                    self.inner.update_named(param_name, &mut param, &grad)?;
770
771                    // Store updated parameter back
772                    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        // Advance the inner optimizer's global step so bias correction progresses.
781        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
811/// PyTorch optimizer factory for creating optimizers with PyTorch-compatible API
812pub struct PyTorchOptimizerFactory;
813
814impl PyTorchOptimizerFactory {
815    /// Create Adam optimizer with PyTorch API
816    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    /// Create AdamW optimizer with PyTorch API
838    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    /// Create SGD optimizer with PyTorch API
860    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
888/// PyTorch-compatible learning rate scheduler wrapper
889pub struct PyTorchLRScheduler {
890    inner_scheduler: Box<dyn LRScheduler>,
891    optimizer: Box<dyn PyTorchOptimizer>,
892    last_epoch: i64,
893}
894
895impl PyTorchLRScheduler {
896    /// Create new scheduler wrapper
897    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    /// Step the scheduler
906    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        // Update all parameter groups
913        for group in self.optimizer.param_groups_mut() {
914            group.lr = new_lr as f64;
915        }
916
917        Ok(())
918    }
919
920    /// Get current learning rate
921    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    /// Get current state dict
926    pub fn state_dict(&self) -> serde_json::Value {
927        serde_json::json!({
928            "last_epoch": self.last_epoch,
929            "scheduler_state": "serialized_state" // Would need scheduler serialization
930        })
931    }
932
933    /// Load state dict
934    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        // Check that gradient is set
1067        assert_eq!(
1068            optimizer.gradients.lock().expect("Mutex lock poisoned").len(),
1069            1
1070        );
1071
1072        // Zero gradients
1073        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    /// Regression: `load_state_dict` used to restore nothing, insert momentum buffers
1081    /// into the *parameter* registry, and report success on a malformed checkpoint.
1082    ///
1083    /// The check is a real resume: after save → new optimizer → load, the next step
1084    /// must land exactly where the uninterrupted run's next step lands.
1085    #[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    /// A malformed checkpoint must fail loudly instead of "loading" nothing.
1169    #[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}