1use 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
13pub 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 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 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 let mut grad_to_use = grad.clone();
128 if weight_decay != 0.0 {
129 grad_to_use = grad_to_use.add(¶m_read.mul_scalar(weight_decay)?)?;
130 }
131
132 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 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 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(¶m_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(¶m_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}