Skip to main content

torsh_optim/
asgd.rs

1//! Averaged Stochastic Gradient Descent optimizer
2
3use crate::{
4    Optimizer, OptimizerError, OptimizerResult, OptimizerState, ParamGroup, ParamGroupState,
5};
6use parking_lot::RwLock;
7use std::collections::HashMap;
8use std::ops::Add;
9use std::sync::Arc;
10use torsh_core::error::{Result, TorshError};
11use torsh_tensor::Tensor;
12
13/// Averaged Stochastic Gradient Descent (ASGD) optimizer
14///
15/// This implements the Averaged SGD algorithm which maintains a running average
16/// of parameters during training. It can achieve better convergence properties
17/// than standard SGD in certain scenarios.
18pub struct ASGD {
19    param_groups: Vec<ParamGroup>,
20    state: HashMap<String, HashMap<String, Tensor>>,
21    step_count: usize,
22    alpha: f32,
23    t0: f32,
24    lambd: f32,
25}
26
27impl ASGD {
28    /// Create a new ASGD optimizer
29    ///
30    /// # Arguments
31    /// * `params` - Parameters to optimize
32    /// * `lr` - Learning rate (default: 1e-2)
33    /// * `alpha` - Power for computing average (default: 0.75)
34    /// * `t0` - Point at which to start averaging (default: 1e6)
35    /// * `lambd` - Decay term (default: 1e-4)
36    /// * `weight_decay` - Weight decay (L2 penalty) (default: 0.0)
37    pub fn new(
38        params: Vec<Arc<RwLock<Tensor>>>,
39        lr: Option<f32>,
40        alpha: Option<f32>,
41        t0: Option<f32>,
42        lambd: Option<f32>,
43        weight_decay: Option<f32>,
44    ) -> Self {
45        let lr = lr.unwrap_or(1e-2);
46        let alpha = alpha.unwrap_or(0.75);
47        let t0 = t0.unwrap_or(1e6);
48        let lambd = lambd.unwrap_or(1e-4);
49        let weight_decay = weight_decay.unwrap_or(0.0);
50
51        let mut options = HashMap::new();
52        options.insert("weight_decay".to_string(), weight_decay);
53
54        let param_group = ParamGroup::new(params, lr).with_options(options);
55
56        Self {
57            param_groups: vec![param_group],
58            state: HashMap::new(),
59            step_count: 0,
60            alpha,
61            t0,
62            lambd,
63        }
64    }
65
66    fn get_param_id(param: &Arc<RwLock<Tensor>>) -> String {
67        format!("{:p}", Arc::as_ptr(param))
68    }
69}
70
71impl Optimizer for ASGD {
72    fn step(&mut self) -> OptimizerResult<()> {
73        self.step_count += 1;
74
75        for group in &self.param_groups {
76            let lr = group.lr;
77            let weight_decay = group.options.get("weight_decay").copied().unwrap_or(0.0);
78
79            for param in &group.params {
80                let param_id = Self::get_param_id(param);
81                let param_read = param.read();
82
83                let grad = param_read.grad().ok_or_else(|| {
84                    TorshError::invalid_argument_with_context(
85                        "Parameter has no gradient",
86                        "asgd_step",
87                    )
88                })?;
89
90                // Get or initialize state
91                let param_state = self.state.entry(param_id.clone()).or_default();
92
93                let mut eta = if !param_state.contains_key("eta") {
94                    let eta = lr;
95                    param_state.insert("eta".to_string(), Tensor::scalar(eta)?);
96                    eta
97                } else {
98                    param_state
99                        .get("eta")
100                        .expect("eta state should exist")
101                        .item()?
102                };
103
104                let ax = if !param_state.contains_key("ax") {
105                    let ax = param_read.clone();
106                    param_state.insert("ax".to_string(), ax.clone());
107                    ax
108                } else {
109                    param_state
110                        .get("ax")
111                        .expect("ax state should exist")
112                        .clone()
113                };
114
115                let mu = if !param_state.contains_key("mu") {
116                    let mu = 1.0;
117                    param_state.insert("mu".to_string(), Tensor::scalar(mu)?);
118                    mu
119                } else {
120                    param_state
121                        .get("mu")
122                        .expect("mu state should exist")
123                        .item()?
124                };
125
126                // Apply weight decay
127                let mut grad_to_use = grad.clone();
128                if weight_decay != 0.0 {
129                    grad_to_use = grad_to_use.add(&param_read.mul_scalar(weight_decay)?)?;
130                }
131
132                // Update eta
133                if self.step_count > 1 {
134                    eta = lr / (1.0 + (self.step_count as f32 - 1.0) * self.lambd).powf(self.alpha);
135                    param_state.insert("eta".to_string(), Tensor::scalar(eta)?);
136                }
137
138                // Update parameter
139                drop(param_read);
140                let mut param_write = param.write();
141                crate::param_update::sub_assign(&mut param_write, &grad_to_use.mul_scalar(eta)?)?;
142
143                // Update averaged parameter
144                if self.step_count as f32 >= self.t0 {
145                    let new_mu = mu / (mu + 1.0);
146                    param_state.insert("mu".to_string(), Tensor::scalar(new_mu)?);
147
148                    let new_ax = ax
149                        .mul_scalar(new_mu)?
150                        .add(&param_write.mul_scalar(1.0 - new_mu)?)?;
151                    param_state.insert("ax".to_string(), new_ax);
152                } else {
153                    let new_mu = 1.0 / self.step_count as f32;
154                    param_state.insert("mu".to_string(), Tensor::scalar(new_mu)?);
155
156                    let new_ax = ax
157                        .mul_scalar(1.0 - new_mu)?
158                        .add(&param_write.mul_scalar(new_mu)?)?;
159                    param_state.insert("ax".to_string(), new_ax);
160                }
161            }
162        }
163
164        Ok(())
165    }
166
167    fn zero_grad(&mut self) {
168        for group in &self.param_groups {
169            for param in &group.params {
170                param.write().zero_grad();
171            }
172        }
173    }
174
175    fn get_lr(&self) -> Vec<f32> {
176        self.param_groups.iter().map(|g| g.lr).collect()
177    }
178
179    fn set_lr(&mut self, lr: f32) {
180        for group in &mut self.param_groups {
181            group.lr = lr;
182        }
183    }
184
185    fn add_param_group(&mut self, params: Vec<Arc<RwLock<Tensor>>>, options: HashMap<String, f32>) {
186        let lr = options.get("lr").copied().unwrap_or(1e-2);
187        let group = ParamGroup::new(params, lr).with_options(options);
188        self.param_groups.push(group);
189    }
190
191    fn parameters(&self) -> Vec<Arc<RwLock<Tensor>>> {
192        crate::optimizer::collect_parameters(&self.param_groups)
193    }
194
195    fn state_dict(&self) -> OptimizerResult<OptimizerState> {
196        let param_groups = self
197            .param_groups
198            .iter()
199            .map(|g| ParamGroupState {
200                lr: g.lr,
201                options: g.options.clone(),
202                param_count: g.params.len(),
203            })
204            .collect();
205
206        Ok(OptimizerState {
207            optimizer_type: "ASGD".to_string(),
208            version: "1.0".to_string(),
209            param_groups,
210            state: self.state.clone(),
211            global_state: HashMap::new(),
212        })
213    }
214
215    fn load_state_dict(&mut self, state: OptimizerState) -> OptimizerResult<()> {
216        if state.param_groups.len() != self.param_groups.len() {
217            return Err(OptimizerError::InvalidParameter(
218                "Parameter group count mismatch".to_string(),
219            ));
220        }
221
222        for (i, group_state) in state.param_groups.iter().enumerate() {
223            self.param_groups[i].lr = group_state.lr;
224            self.param_groups[i].options = group_state.options.clone();
225        }
226
227        self.state = state.state;
228        Ok(())
229    }
230}
231
232#[cfg(test)]
233mod tests {
234    use super::*;
235    use torsh_tensor::creation::randn;
236
237    #[test]
238    fn test_asgd_creation() {
239        let params = vec![Arc::new(RwLock::new(randn::<f32>(&[2, 2]).unwrap()))];
240        let optimizer = ASGD::new(params, None, None, None, None, None);
241        assert_eq!(optimizer.get_lr()[0], 1e-2);
242    }
243
244    #[test]
245    fn test_asgd_step() -> OptimizerResult<()> {
246        let param = Arc::new(RwLock::new(randn::<f32>(&[2, 2])?));
247        let mut param_write = param.write();
248        param_write.set_grad(Some(randn::<f32>(&[2, 2])?));
249        drop(param_write);
250
251        let params = vec![param];
252        let mut optimizer = ASGD::new(params, Some(0.1), None, None, None, None);
253
254        optimizer.step()?;
255        Ok(())
256    }
257}