1use crate::{
4 Optimizer, OptimizerError, OptimizerResult, OptimizerState, ParamGroup, ParamGroupState,
5};
6use torsh_core::error::{Result, TorshError};
7use torsh_tensor::Tensor;
8use parking_lot::RwLock;
11use std::collections::HashMap;
12use std::sync::Arc;
13
14#[derive(Clone)]
16pub struct BaseOptimizer {
17 pub(crate) param_groups: Vec<ParamGroup>,
18 pub(crate) state: HashMap<String, HashMap<String, Tensor>>,
19 #[allow(dead_code)]
21 pub(crate) optimizer_type: String,
22 pub(crate) defaults: HashMap<String, f32>,
23}
24
25impl BaseOptimizer {
26 #[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 #[allow(dead_code)]
45 pub(crate) fn param_id(param: &Arc<RwLock<Tensor>>) -> String {
46 format!("{:p}", param.as_ref())
47 }
48
49 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 #[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 *exp_avg = exp_avg
203 .add(&grad_term)
204 .map_err(OptimizerError::TensorError)?;
205 Ok(())
206 }
207
208 #[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 *exp_avg_sq = exp_avg_sq
225 .add(&grad_sq_term)
226 .map_err(OptimizerError::TensorError)?;
227 Ok(())
228 }
229
230 #[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 #[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 pub(crate) fn parameters(&self) -> Vec<Arc<RwLock<Tensor>>> {
260 collect_parameters(&self.param_groups)
261 }
262}
263
264pub(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 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 fn load_state_dict(&mut self, state: OptimizerState) -> OptimizerResult<()> {
330 state
332 .validate()
333 .map_err(|e| OptimizerError::StateError(e.to_string()))?;
334
335 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 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 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 self.state = state.state;
367
368 for (key, value) in state.global_state {
370 self.defaults.insert(key, value);
371 }
372
373 Ok(())
374 }
375}
376
377impl BaseOptimizer {
378 #[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 for (key, value) in &self.defaults {
396 optimizer_state.global_state.insert(key.clone(), *value);
397 }
398
399 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
410pub mod functional {
412 use super::*;
413
414 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 0.0
425 } else {
426 0.0
427 }
428 }
429}