optirs_core/optimizers/
sgd.rs1use scirs2_core::ndarray::{Array, Dimension, IxDyn, ScalarOperand, Zip};
4use scirs2_core::numeric::Float;
5use std::fmt::Debug;
6
7use crate::error::{OptimError, Result};
8use crate::optimizers::Optimizer;
9
10#[derive(Debug, Clone)]
35pub struct SGD<A: Float + ScalarOperand + Debug> {
36 learning_rate: A,
38 momentum: A,
40 weight_decay: A,
42 velocity: Option<Vec<Array<A, IxDyn>>>,
44}
45
46impl<A: Float + ScalarOperand + Debug + Send + Sync> SGD<A> {
47 pub fn new(learning_rate: A) -> Self {
53 Self {
54 learning_rate,
55 momentum: A::zero(),
56 weight_decay: A::zero(),
57 velocity: None,
58 }
59 }
60
61 pub fn new_with_config(learning_rate: A, momentum: A, weight_decay: A) -> Self {
69 Self {
70 learning_rate,
71 momentum,
72 weight_decay,
73 velocity: None,
74 }
75 }
76
77 pub fn set_momentum(&mut self, momentum: A) -> &mut Self {
83 self.momentum = momentum;
84 self
85 }
86
87 pub fn with_momentum(mut self, momentum: A) -> Self {
93 self.momentum = momentum;
94 self
95 }
96
97 pub fn get_momentum(&self) -> A {
99 self.momentum
100 }
101
102 pub fn learning_rate(&self) -> A {
104 self.learning_rate
105 }
106
107 pub fn set_weight_decay(&mut self, weight_decay: A) -> &mut Self {
113 self.weight_decay = weight_decay;
114 self
115 }
116
117 pub fn with_weight_decay(mut self, weight_decay: A) -> Self {
123 self.weight_decay = weight_decay;
124 self
125 }
126
127 pub fn get_weight_decay(&self) -> A {
129 self.weight_decay
130 }
131
132 pub fn reset(&mut self) {
134 self.velocity = None;
135 }
136
137 fn ensure_state(&mut self, index: usize, dim: &IxDyn) {
139 let velocity = self.velocity.get_or_insert_with(Vec::new);
140 while velocity.len() <= index {
141 velocity.push(Array::zeros(dim.clone()));
142 }
143 if velocity[index].raw_dim() != *dim {
144 velocity[index] = Array::zeros(dim.clone());
145 }
146 }
147
148 pub fn step_inplace_indexed<D: Dimension>(
153 &mut self,
154 index: usize,
155 params: &mut Array<A, D>,
156 gradients: &Array<A, D>,
157 ) -> Result<()> {
158 if params.shape() != gradients.shape() {
159 return Err(OptimError::DimensionMismatch(format!(
160 "Incompatible shapes: parameters have shape {:?}, gradients have shape {:?}",
161 params.shape(),
162 gradients.shape()
163 )));
164 }
165
166 let dim = params.raw_dim().into_dyn();
167 self.ensure_state(index, &dim);
168
169 let momentum = self.momentum;
170 let lr = self.learning_rate;
171 let weight_decay = self.weight_decay;
172 let use_weight_decay = weight_decay > A::zero();
173 let use_momentum = momentum > A::zero();
174
175 let velocity = self
176 .velocity
177 .as_mut()
178 .ok_or_else(|| OptimError::InvalidConfig("SGD state not initialized".to_string()))?;
179
180 let mut params_view = params.view_mut().into_dyn();
181 let gradients_view = gradients.view().into_dyn();
182
183 Zip::from(&mut params_view)
184 .and(&gradients_view)
185 .and(&mut velocity[index])
186 .for_each(|p, &g, v| {
187 let grad = if use_weight_decay {
188 g + weight_decay * *p
189 } else {
190 g
191 };
192 *v = if use_momentum {
193 *v * momentum + grad * lr
194 } else {
195 grad * lr
196 };
197 *p = *p - *v;
198 });
199
200 Ok(())
201 }
202
203 pub fn step_inplace<D: Dimension>(
205 &mut self,
206 params: &mut Array<A, D>,
207 gradients: &Array<A, D>,
208 ) -> Result<()> {
209 self.step_inplace_indexed(0, params, gradients)
210 }
211
212 pub fn step_indexed<D: Dimension>(
217 &mut self,
218 index: usize,
219 params: &Array<A, D>,
220 gradients: &Array<A, D>,
221 ) -> Result<Array<A, D>> {
222 let mut updated = params.to_owned();
223 self.step_inplace_indexed(index, &mut updated, gradients)?;
224 Ok(updated)
225 }
226}
227
228impl<A, D> Optimizer<A, D> for SGD<A>
229where
230 A: Float + ScalarOperand + Debug + Send + Sync,
231 D: Dimension,
232{
233 fn step(&mut self, params: &Array<A, D>, gradients: &Array<A, D>) -> Result<Array<A, D>> {
234 self.step_indexed(0, params, gradients)
235 }
236
237 fn step_list(
238 &mut self,
239 params_list: &[&Array<A, D>],
240 gradients_list: &[&Array<A, D>],
241 ) -> Result<Vec<Array<A, D>>> {
242 if params_list.len() != gradients_list.len() {
243 return Err(OptimError::InvalidConfig(format!(
244 "Number of parameter arrays ({}) does not match number of gradient arrays ({})",
245 params_list.len(),
246 gradients_list.len()
247 )));
248 }
249
250 let mut results = Vec::with_capacity(params_list.len());
251 for (index, (params, grads)) in params_list.iter().zip(gradients_list.iter()).enumerate() {
252 results.push(self.step_indexed(index, params, grads)?);
253 }
254 Ok(results)
255 }
256
257 fn get_learning_rate(&self) -> A {
258 self.learning_rate
259 }
260
261 fn set_learning_rate(&mut self, learning_rate: A) {
262 self.learning_rate = learning_rate;
263 }
264}