1use crate::{
4 Optimizer, OptimizerError, OptimizerResult, OptimizerState, ParamGroup, ParamGroupState,
5};
6use parking_lot::RwLock;
7use std::collections::HashMap;
8use std::ops::{Add, Mul};
9use std::sync::Arc;
10use torsh_core::error::{Result, TorshError};
11use torsh_tensor::Tensor;
12
13pub struct SparseAdam {
18 param_groups: Vec<ParamGroup>,
19 state: HashMap<String, HashMap<String, Tensor>>,
20 step_count: usize,
21 beta1: f32,
22 beta2: f32,
23 eps: f32,
24}
25
26impl SparseAdam {
27 pub fn new(
37 params: Vec<Arc<RwLock<Tensor>>>,
38 lr: Option<f32>,
39 beta1: Option<f32>,
40 beta2: Option<f32>,
41 eps: Option<f32>,
42 weight_decay: Option<f32>,
43 ) -> Self {
44 let lr = lr.unwrap_or(1e-3);
45 let beta1 = beta1.unwrap_or(0.9);
46 let beta2 = beta2.unwrap_or(0.999);
47 let eps = eps.unwrap_or(1e-8);
48 let weight_decay = weight_decay.unwrap_or(0.0);
49
50 let mut options = HashMap::new();
51 options.insert("weight_decay".to_string(), weight_decay);
52
53 let param_group = ParamGroup::new(params, lr).with_options(options);
54
55 Self {
56 param_groups: vec![param_group],
57 state: HashMap::new(),
58 step_count: 0,
59 beta1,
60 beta2,
61 eps,
62 }
63 }
64
65 fn get_param_id(param: &Arc<RwLock<Tensor>>) -> String {
66 format!("{:p}", Arc::as_ptr(param))
67 }
68
69 fn sparse_update(
71 &self,
72 param: &mut Tensor,
73 grad: &Tensor,
74 exp_avg: &mut Tensor,
75 exp_avg_sq: &mut Tensor,
76 lr: f32,
77 step: usize,
78 ) -> Result<()> {
79 let exp_avg_update = grad.mul_scalar(1.0 - self.beta1)?;
85 *exp_avg = exp_avg.mul_scalar(self.beta1)?.add(&exp_avg_update)?;
86
87 let grad_sq = grad.mul(grad)?;
89 let exp_avg_sq_update = grad_sq.mul_scalar(1.0 - self.beta2)?;
90 *exp_avg_sq = exp_avg_sq.mul_scalar(self.beta2)?.add(&exp_avg_sq_update)?;
91
92 let bias_correction1 = 1.0 - self.beta1.powi(step as i32);
94 let bias_correction2 = 1.0 - self.beta2.powi(step as i32);
95
96 let exp_avg_corrected = exp_avg.div_scalar(bias_correction1)?;
98 let exp_avg_sq_corrected = exp_avg_sq.div_scalar(bias_correction2)?;
99
100 let denom = exp_avg_sq_corrected.sqrt()?.add_scalar(self.eps)?;
102 let update = exp_avg_corrected.div(&denom)?;
103
104 crate::param_update::sub_assign(&mut *param, &update.mul_scalar(lr)?)?;
106
107 Ok(())
108 }
109}
110
111impl Optimizer for SparseAdam {
112 fn step(&mut self) -> OptimizerResult<()> {
113 self.step_count += 1;
114
115 for group in &self.param_groups {
116 let lr = group.lr;
117 let weight_decay = group.options.get("weight_decay").copied().unwrap_or(0.0);
118
119 for param in &group.params {
120 let param_id = Self::get_param_id(param);
121 let param_read = param.read();
122
123 let grad = param_read.grad().ok_or_else(|| {
124 TorshError::invalid_argument_with_context(
125 "Parameter has no gradient",
126 "sparse_adam_step",
127 )
128 })?;
129
130 let mut exp_avg = if let Some(state) = self.state.get(¶m_id) {
139 if let Some(exp_avg) = state.get("exp_avg") {
140 exp_avg.clone()
141 } else {
142 Tensor::zeros_like(¶m_read)?
143 }
144 } else {
145 Tensor::zeros_like(¶m_read)?
146 };
147
148 let mut exp_avg_sq = if let Some(state) = self.state.get(¶m_id) {
149 if let Some(exp_avg_sq) = state.get("exp_avg_sq") {
150 exp_avg_sq.clone()
151 } else {
152 Tensor::zeros_like(¶m_read)?
153 }
154 } else {
155 Tensor::zeros_like(¶m_read)?
156 };
157
158 let mut grad_to_use = grad.clone();
160 if weight_decay != 0.0 {
161 grad_to_use = grad_to_use.add(¶m_read.mul_scalar(weight_decay)?)?;
162 }
163
164 drop(param_read);
165 let mut param_write = param.write();
166
167 self.sparse_update(
169 &mut param_write,
170 &grad_to_use,
171 &mut exp_avg,
172 &mut exp_avg_sq,
173 lr,
174 self.step_count,
175 )?;
176
177 let param_state = self.state.entry(param_id.clone()).or_default();
179 param_state.insert("exp_avg".to_string(), exp_avg);
180 param_state.insert("exp_avg_sq".to_string(), exp_avg_sq);
181 }
182 }
183
184 Ok(())
185 }
186
187 fn zero_grad(&mut self) {
188 for group in &self.param_groups {
189 for param in &group.params {
190 param.write().zero_grad();
191 }
192 }
193 }
194
195 fn get_lr(&self) -> Vec<f32> {
196 self.param_groups.iter().map(|g| g.lr).collect()
197 }
198
199 fn set_lr(&mut self, lr: f32) {
200 for group in &mut self.param_groups {
201 group.lr = lr;
202 }
203 }
204
205 fn add_param_group(&mut self, params: Vec<Arc<RwLock<Tensor>>>, options: HashMap<String, f32>) {
206 let lr = options.get("lr").copied().unwrap_or(1e-3);
207 let group = ParamGroup::new(params, lr).with_options(options);
208 self.param_groups.push(group);
209 }
210
211 fn parameters(&self) -> Vec<Arc<RwLock<Tensor>>> {
212 crate::optimizer::collect_parameters(&self.param_groups)
213 }
214
215 fn state_dict(&self) -> OptimizerResult<OptimizerState> {
216 let param_groups = self
217 .param_groups
218 .iter()
219 .map(|g| ParamGroupState {
220 lr: g.lr,
221 options: g.options.clone(),
222 param_count: g.params.len(),
223 })
224 .collect();
225
226 Ok(OptimizerState {
227 optimizer_type: "SparseAdam".to_string(),
228 version: "0.1.0".to_string(),
229 param_groups,
230 state: self.state.clone(),
231 global_state: HashMap::new(),
232 })
233 }
234
235 fn load_state_dict(&mut self, state: OptimizerState) -> OptimizerResult<()> {
236 if state.param_groups.len() != self.param_groups.len() {
237 return Err(OptimizerError::InvalidParameter(
238 "Parameter group count mismatch".to_string(),
239 ));
240 }
241
242 for (i, group_state) in state.param_groups.iter().enumerate() {
243 self.param_groups[i].lr = group_state.lr;
244 self.param_groups[i].options = group_state.options.clone();
245 }
246
247 self.state = state.state;
248 Ok(())
249 }
250}
251
252#[cfg(test)]
253mod tests {
254 use super::*;
255 use torsh_tensor::creation::randn;
256
257 #[test]
258 fn test_sparse_adam_creation() {
259 let params = vec![Arc::new(RwLock::new(randn::<f32>(&[2, 2]).unwrap()))];
260 let optimizer = SparseAdam::new(params, None, None, None, None, None);
261 assert_eq!(optimizer.get_lr()[0], 1e-3);
262 }
263
264 #[test]
265 #[ignore = "Temporarily disabled due to potential deadlock"]
266 fn test_sparse_adam_step() -> OptimizerResult<()> {
267 let param = Arc::new(RwLock::new(randn::<f32>(&[2, 2]).unwrap()));
268 let mut param_write = param.write();
269 param_write.set_grad(Some(randn::<f32>(&[2, 2]).unwrap()));
270 drop(param_write);
271
272 let params = vec![param];
273 let mut optimizer = SparseAdam::new(params, Some(0.1), None, None, None, None);
274
275 optimizer.step()?;
276 Ok(())
277 }
278
279 #[test]
280 fn test_sparse_adam_basic() -> OptimizerResult<()> {
281 let params = vec![Arc::new(RwLock::new(randn::<f32>(&[2, 2])?))];
283 let mut optimizer = SparseAdam::new(
284 params,
285 Some(0.01),
286 Some(0.9),
287 Some(0.999),
288 Some(1e-10),
289 Some(0.01),
290 );
291
292 assert_eq!(optimizer.get_lr()[0], 0.01);
293 optimizer.set_lr(0.001);
294 assert_eq!(optimizer.get_lr()[0], 0.001);
295
296 let state = optimizer.state_dict()?;
298 assert_eq!(state.param_groups.len(), 1);
299 assert_eq!(state.param_groups[0].lr, 0.001);
300 Ok(())
301 }
302}