Skip to main content

torsh_optim/
optimizer.rs

1//! Base optimizer implementation utilities
2
3use crate::{
4    Optimizer, OptimizerError, OptimizerResult, OptimizerState, ParamGroup, ParamGroupState,
5};
6use torsh_core::error::{Result, TorshError};
7use torsh_tensor::Tensor;
8// Temporarily disable scirs2 integration
9// use scirs2::optim::{Optimizer as SciOptimizer, OptimizerConfig};
10use parking_lot::RwLock;
11use std::collections::HashMap;
12use std::sync::Arc;
13
14/// Base optimizer struct (simplified without scirs2 integration)
15#[derive(Clone)]
16pub struct BaseOptimizer {
17    pub(crate) param_groups: Vec<ParamGroup>,
18    pub(crate) state: HashMap<String, HashMap<String, Tensor>>,
19    // Placeholder for optimizer-specific data
20    #[allow(dead_code)]
21    pub(crate) optimizer_type: String,
22    pub(crate) defaults: HashMap<String, f32>,
23}
24
25impl BaseOptimizer {
26    /// Apply weight decay if specified
27    #[allow(dead_code)]
28    pub(crate) fn apply_weight_decay(
29        &self,
30        param: &mut Tensor,
31        weight_decay: f32,
32    ) -> OptimizerResult<()> {
33        if weight_decay != 0.0 {
34            let decay = param
35                .mul_scalar(weight_decay)
36                .map_err(OptimizerError::TensorError)?;
37            crate::param_update::sub_assign(&mut *param, &decay)
38                .map_err(OptimizerError::TensorError)?;
39        }
40        Ok(())
41    }
42
43    /// Get parameter ID for state tracking
44    #[allow(dead_code)]
45    pub(crate) fn param_id(param: &Arc<RwLock<Tensor>>) -> String {
46        format!("{:p}", param.as_ref())
47    }
48
49    /// Initialize state for a parameter if not exists
50    #[allow(dead_code)]
51    pub(crate) fn init_state(&mut self, param_id: String) {
52        self.state.entry(param_id).or_default();
53    }
54
55    /// Get or create state tensor
56    #[allow(dead_code)]
57    pub(crate) fn get_or_create_state(
58        &mut self,
59        param_id: &str,
60        state_name: &str,
61        init_fn: impl FnOnce() -> Tensor,
62    ) -> Tensor {
63        self.state
64            .get_mut(param_id)
65            .expect("state should exist for param_id")
66            .entry(state_name.to_string())
67            .or_insert_with(init_fn)
68            .clone()
69    }
70
71    /// Update state tensor
72    #[allow(dead_code)]
73    pub(crate) fn update_state(&mut self, param_id: &str, state_name: &str, value: Tensor) {
74        self.state
75            .get_mut(param_id)
76            .expect("state should exist for param_id")
77            .insert(state_name.to_string(), value);
78    }
79
80    /// Initialize state with zeros_like for common optimizer states
81    #[allow(dead_code)]
82    pub(crate) fn init_state_with_zeros(
83        &mut self,
84        param_id: String,
85        param: &Tensor,
86        state_names: &[&str],
87    ) -> OptimizerResult<()> {
88        let state = self.state.entry(param_id).or_default();
89        for &name in state_names {
90            if !state.contains_key(name) {
91                let zeros = torsh_tensor::creation::zeros_like(param)
92                    .map_err(OptimizerError::TensorError)?;
93                state.insert(name.to_string(), zeros);
94            }
95        }
96        Ok(())
97    }
98
99    /// Initialize common Adam-like optimizer state
100    #[allow(dead_code)]
101    pub(crate) fn init_adam_state(
102        &mut self,
103        param_id: String,
104        param: &Tensor,
105        amsgrad: bool,
106    ) -> OptimizerResult<()> {
107        let state_names = if amsgrad {
108            vec!["step", "exp_avg", "exp_avg_sq", "max_exp_avg_sq"]
109        } else {
110            vec!["step", "exp_avg", "exp_avg_sq"]
111        };
112        self.init_state_with_zeros(param_id, param, &state_names)
113    }
114
115    /// Initialize common SGD-like optimizer state
116    #[allow(dead_code)]
117    pub(crate) fn init_sgd_state(
118        &mut self,
119        param_id: String,
120        param: &Tensor,
121        momentum: bool,
122    ) -> OptimizerResult<()> {
123        let state_names = if momentum {
124            vec!["momentum_buffer"]
125        } else {
126            vec![]
127        };
128        if !state_names.is_empty() {
129            self.init_state_with_zeros(param_id, param, &state_names)
130        } else {
131            self.init_state(param_id);
132            Ok(())
133        }
134    }
135
136    /// Apply weight decay to gradients
137    #[allow(dead_code)]
138    pub(crate) fn apply_weight_decay_to_grad(
139        &self,
140        grad: &mut Tensor,
141        param: &Tensor,
142        weight_decay: f32,
143    ) -> OptimizerResult<()> {
144        if weight_decay != 0.0 {
145            let weight_decay_term = param
146                .mul_scalar(weight_decay)
147                .map_err(OptimizerError::TensorError)?;
148            *grad = grad
149                .add_op(&weight_decay_term)
150                .map_err(OptimizerError::TensorError)?;
151        }
152        Ok(())
153    }
154
155    /// Get step count from state, incrementing if requested
156    #[allow(dead_code)]
157    pub(crate) fn get_step_count(
158        &mut self,
159        param_id: &str,
160        increment: bool,
161    ) -> OptimizerResult<i32> {
162        let state = self
163            .state
164            .get_mut(param_id)
165            .expect("state should exist for param_id");
166        let step_tensor = state.get_mut("step").expect("step state should exist");
167
168        if increment {
169            step_tensor
170                .add_scalar_(1.0)
171                .map_err(OptimizerError::TensorError)?;
172        }
173
174        let step = step_tensor.to_vec().map_err(OptimizerError::TensorError)?[0] as i32;
175        Ok(step)
176    }
177
178    /// Compute bias correction terms for Adam-like optimizers
179    #[allow(dead_code)]
180    pub(crate) fn compute_bias_correction(&self, betas: (f32, f32), step: i32) -> (f32, f32) {
181        let bias_correction1 = 1.0 - betas.0.powi(step);
182        let bias_correction2 = 1.0 - betas.1.powi(step);
183        (bias_correction1, bias_correction2)
184    }
185
186    /// Update exponential moving average
187    #[allow(dead_code)]
188    pub(crate) fn update_exp_avg(
189        &self,
190        exp_avg: &mut Tensor,
191        grad: &Tensor,
192        beta: f32,
193    ) -> OptimizerResult<()> {
194        exp_avg
195            .mul_scalar_(beta)
196            .map_err(OptimizerError::TensorError)?;
197        let grad_term = grad
198            .mul_scalar(1.0 - beta)
199            .map_err(OptimizerError::TensorError)?;
200        // `add` is non-mutating and returns a new tensor; write it back through
201        // the `&mut` reference so the moving average actually accumulates.
202        *exp_avg = exp_avg
203            .add(&grad_term)
204            .map_err(OptimizerError::TensorError)?;
205        Ok(())
206    }
207
208    /// Update exponential moving average of squared gradients
209    #[allow(dead_code)]
210    pub(crate) fn update_exp_avg_sq(
211        &self,
212        exp_avg_sq: &mut Tensor,
213        grad: &Tensor,
214        beta: f32,
215    ) -> OptimizerResult<()> {
216        exp_avg_sq
217            .mul_scalar_(beta)
218            .map_err(OptimizerError::TensorError)?;
219        let grad_squared = grad.mul_op(grad).map_err(OptimizerError::TensorError)?;
220        let grad_sq_term = grad_squared
221            .mul_scalar(1.0 - beta)
222            .map_err(OptimizerError::TensorError)?;
223        // `add` is non-mutating; write it back through the `&mut` reference.
224        *exp_avg_sq = exp_avg_sq
225            .add(&grad_sq_term)
226            .map_err(OptimizerError::TensorError)?;
227        Ok(())
228    }
229
230    /// Apply gradient clipping to a gradient tensor
231    #[allow(dead_code)]
232    pub(crate) fn clip_gradient(&self, grad: &mut Tensor, max_norm: f32) -> OptimizerResult<f32> {
233        let norm = grad.norm().map_err(OptimizerError::TensorError)?;
234        let norm_value = norm.to_vec().map_err(OptimizerError::TensorError)?[0];
235
236        if norm_value > max_norm {
237            let scale = max_norm / norm_value;
238            *grad = grad
239                .mul_scalar(scale)
240                .map_err(OptimizerError::TensorError)?;
241        }
242
243        Ok(norm_value)
244    }
245
246    /// Check if all parameters have gradients
247    #[allow(dead_code)]
248    pub(crate) fn validate_gradients(&self) -> bool {
249        self.param_groups
250            .iter()
251            .all(|group| group.params.iter().all(|param| param.read().has_grad()))
252    }
253
254    /// Collect all parameter tensor handles across every parameter group.
255    ///
256    /// Returns cheap `Arc` clones that share the underlying tensors, so callers
257    /// can read parameters and access their gradients. Used by [`Optimizer::parameters`]
258    /// implementations of optimizers built on top of [`BaseOptimizer`].
259    pub(crate) fn parameters(&self) -> Vec<Arc<RwLock<Tensor>>> {
260        collect_parameters(&self.param_groups)
261    }
262}
263
264/// Collect all parameter tensor handles from a slice of parameter groups.
265///
266/// Shared by the [`Optimizer::parameters`] implementations of every optimizer
267/// that stores its parameters as a `Vec<ParamGroup>` (either directly or via
268/// [`BaseOptimizer`]).
269pub(crate) fn collect_parameters(param_groups: &[ParamGroup]) -> Vec<Arc<RwLock<Tensor>>> {
270    param_groups
271        .iter()
272        .flat_map(|group| group.params.iter().cloned())
273        .collect()
274}
275
276impl Optimizer for BaseOptimizer {
277    fn step(&mut self) -> OptimizerResult<()> {
278        // Temporarily disabled - would use scirs2's optimizer when integrated
279        // For now, return a placeholder error
280        Err(OptimizerError::TensorError(TorshError::Other(
281            "Optimizer step not yet implemented - scirs2 integration pending".to_string(),
282        )))
283    }
284
285    fn zero_grad(&mut self) {
286        for group in &self.param_groups {
287            for param in &group.params {
288                param.write().zero_grad();
289            }
290        }
291    }
292
293    fn get_lr(&self) -> Vec<f32> {
294        self.param_groups.iter().map(|g| g.lr).collect()
295    }
296
297    fn set_lr(&mut self, lr: f32) {
298        for group in &mut self.param_groups {
299            group.lr = lr;
300        }
301    }
302
303    fn set_lrs(&mut self, lrs: &[f32]) {
304        for (group, &lr) in self.param_groups.iter_mut().zip(lrs.iter()) {
305            group.lr = lr;
306        }
307    }
308
309    fn add_param_group(
310        &mut self,
311        params: Vec<Arc<RwLock<Tensor>>>,
312        mut options: HashMap<String, f32>,
313    ) {
314        let lr = options
315            .remove("lr")
316            .unwrap_or_else(|| self.defaults.get("lr").copied().unwrap_or(1e-3));
317
318        let mut group = ParamGroup::new(params, lr);
319        group.options = options;
320        self.param_groups.push(group);
321    }
322
323    fn state_dict(&self) -> OptimizerResult<OptimizerState> {
324        self.create_state_dict(None)
325    }
326
327    // This method is moved outside of the trait implementation block below
328
329    fn load_state_dict(&mut self, state: OptimizerState) -> OptimizerResult<()> {
330        // Validate the incoming state
331        state
332            .validate()
333            .map_err(|e| OptimizerError::StateError(e.to_string()))?;
334
335        // Check compatibility
336        if state.param_groups.len() != self.param_groups.len() {
337            return Err(OptimizerError::StateError(
338                "Loaded state dict has different number of parameter groups".to_string(),
339            ));
340        }
341
342        // Check parameter counts match
343        for (i, (group, state_group)) in self
344            .param_groups
345            .iter()
346            .zip(state.param_groups.iter())
347            .enumerate()
348        {
349            if group.params.len() != state_group.param_count {
350                return Err(OptimizerError::StateError(format!(
351                    "Parameter count mismatch in group {}: expected {}, got {}",
352                    i,
353                    group.params.len(),
354                    state_group.param_count
355                )));
356            }
357        }
358
359        // Update parameter groups
360        for (group, state_group) in self.param_groups.iter_mut().zip(state.param_groups.iter()) {
361            group.lr = state_group.lr;
362            group.options = state_group.options.clone();
363        }
364
365        // Update optimizer state
366        self.state = state.state;
367
368        // Update defaults from global state
369        for (key, value) in state.global_state {
370            self.defaults.insert(key, value);
371        }
372
373        Ok(())
374    }
375}
376
377impl BaseOptimizer {
378    /// Create a standardized state dict with optional additional global state
379    #[allow(dead_code)]
380    pub(crate) fn create_state_dict(
381        &self,
382        additional_global_state: Option<HashMap<String, f32>>,
383    ) -> OptimizerResult<OptimizerState> {
384        let param_groups = self
385            .param_groups
386            .iter()
387            .map(|g| ParamGroupState::from_param_group(g))
388            .collect();
389
390        let mut optimizer_state = OptimizerState::new(self.optimizer_type.clone());
391        optimizer_state.param_groups = param_groups;
392        optimizer_state.state = self.state.clone();
393
394        // Add any global state from defaults
395        for (key, value) in &self.defaults {
396            optimizer_state.global_state.insert(key.clone(), *value);
397        }
398
399        // Add additional global state if provided
400        if let Some(additional) = additional_global_state {
401            for (key, value) in additional {
402                optimizer_state.global_state.insert(key, value);
403            }
404        }
405
406        Ok(optimizer_state)
407    }
408}
409
410/// Functional utilities for optimizers
411pub mod functional {
412    use super::*;
413
414    /// Apply gradient clipping before optimizer step
415    pub fn clip_grad_before_step<O: Optimizer>(
416        _optimizer: &O,
417        max_norm: Option<f32>,
418        _norm_type: f32,
419    ) -> f32 {
420        if let Some(_max_norm) = max_norm {
421            // Collect all parameters from optimizer
422            // This would need access to parameters through the optimizer trait
423            // For now, return 0.0 as placeholder
424            0.0
425        } else {
426            0.0
427        }
428    }
429}