1use scirs2_core::ndarray::Array1;
7use scirs2_core::numeric::Float;
8use std::fmt::Debug;
9
10use crate::error::Result;
11use crate::optimizers::Optimizer;
12use crate::simd_optimizer::SimdOptimizer;
13
14#[derive(Debug, Clone)]
50pub struct SimdSGD<A: Float> {
51 learning_rate: A,
53 momentum: A,
55 weight_decay: A,
57 velocity: Option<Array1<A>>,
59}
60
61impl<A: Float> SimdSGD<A> {
62 pub fn new(learning_rate: A) -> Self {
68 Self {
69 learning_rate,
70 momentum: A::zero(),
71 weight_decay: A::zero(),
72 velocity: None,
73 }
74 }
75
76 pub fn new_with_config(learning_rate: A, momentum: A, weight_decay: A) -> Self {
84 Self {
85 learning_rate,
86 momentum,
87 weight_decay,
88 velocity: None,
89 }
90 }
91
92 pub fn set_momentum(&mut self, momentum: A) -> &mut Self {
94 self.momentum = momentum;
95 self
96 }
97
98 pub fn with_momentum(mut self, momentum: A) -> Self {
100 self.momentum = momentum;
101 self
102 }
103
104 pub fn get_momentum(&self) -> A {
106 self.momentum
107 }
108
109 pub fn learning_rate(&self) -> A {
111 self.learning_rate
112 }
113
114 pub fn set_weight_decay(&mut self, weight_decay: A) -> &mut Self {
116 self.weight_decay = weight_decay;
117 self
118 }
119
120 pub fn with_weight_decay(mut self, weight_decay: A) -> Self {
122 self.weight_decay = weight_decay;
123 self
124 }
125
126 pub fn get_weight_decay(&self) -> A {
128 self.weight_decay
129 }
130
131 pub fn reset(&mut self) {
133 self.velocity = None;
134 }
135}
136
137impl Optimizer<f32, scirs2_core::ndarray::Ix1> for SimdSGD<f32> {
139 fn step(&mut self, params: &Array1<f32>, gradients: &Array1<f32>) -> Result<Array1<f32>> {
140 if params.shape() != gradients.shape() {
142 return Err(crate::error::OptimError::DimensionMismatch(format!(
143 "Incompatible shapes: parameters have shape {:?}, gradients have shape {:?}",
144 params.shape(),
145 gradients.shape()
146 )));
147 }
148
149 let params_view = params.view();
150 let gradients_view = gradients.view();
151
152 let adjusted_gradients = if self.weight_decay > 0.0 {
154 f32::simd_weight_decay(&gradients_view, ¶ms_view, self.weight_decay)
155 } else {
156 gradients.to_owned()
157 };
158
159 let velocity = self
161 .velocity
162 .get_or_insert_with(|| Array1::zeros(params.len()));
163
164 if velocity.len() != params.len() {
166 *velocity = Array1::zeros(params.len());
167 }
168
169 let new_params = if self.momentum > 0.0 {
171 let (updated_params, updated_velocity) = f32::simd_momentum_update(
173 ¶ms_view,
174 &adjusted_gradients.view(),
175 &velocity.view(),
176 self.learning_rate,
177 self.momentum,
178 );
179 *velocity = updated_velocity;
180 updated_params
181 } else {
182 f32::simd_sgd_update(¶ms_view, &adjusted_gradients.view(), self.learning_rate)
184 };
185
186 Ok(new_params)
187 }
188
189 fn get_learning_rate(&self) -> f32 {
190 self.learning_rate
191 }
192
193 fn set_learning_rate(&mut self, learning_rate: f32) {
194 self.learning_rate = learning_rate;
195 }
196}
197
198impl Optimizer<f64, scirs2_core::ndarray::Ix1> for SimdSGD<f64> {
200 fn step(&mut self, params: &Array1<f64>, gradients: &Array1<f64>) -> Result<Array1<f64>> {
201 if params.shape() != gradients.shape() {
203 return Err(crate::error::OptimError::DimensionMismatch(format!(
204 "Incompatible shapes: parameters have shape {:?}, gradients have shape {:?}",
205 params.shape(),
206 gradients.shape()
207 )));
208 }
209
210 let params_view = params.view();
211 let gradients_view = gradients.view();
212
213 let adjusted_gradients = if self.weight_decay > 0.0 {
215 f64::simd_weight_decay(&gradients_view, ¶ms_view, self.weight_decay)
216 } else {
217 gradients.to_owned()
218 };
219
220 let velocity = self
222 .velocity
223 .get_or_insert_with(|| Array1::zeros(params.len()));
224
225 if velocity.len() != params.len() {
227 *velocity = Array1::zeros(params.len());
228 }
229
230 let new_params = if self.momentum > 0.0 {
232 let (updated_params, updated_velocity) = f64::simd_momentum_update(
234 ¶ms_view,
235 &adjusted_gradients.view(),
236 &velocity.view(),
237 self.learning_rate,
238 self.momentum,
239 );
240 *velocity = updated_velocity;
241 updated_params
242 } else {
243 f64::simd_sgd_update(¶ms_view, &adjusted_gradients.view(), self.learning_rate)
245 };
246
247 Ok(new_params)
248 }
249
250 fn get_learning_rate(&self) -> f64 {
251 self.learning_rate
252 }
253
254 fn set_learning_rate(&mut self, learning_rate: f64) {
255 self.learning_rate = learning_rate;
256 }
257}
258
259#[cfg(test)]
260mod tests {
261 use super::*;
262 use approx::assert_relative_eq;
263
264 #[test]
265 fn test_simd_sgd_basic() {
266 let params = Array1::from_vec(vec![1.0f32, 2.0, 3.0, 4.0]);
267 let gradients = Array1::from_vec(vec![0.1, 0.2, 0.3, 0.4]);
268
269 let mut optimizer = SimdSGD::new(0.1);
270 let result = optimizer
271 .step(¶ms, &gradients)
272 .expect("optimizer.step succeeds in test_simd_sgd_basic");
273
274 assert_relative_eq!(result[0], 0.99, epsilon = 1e-6);
275 assert_relative_eq!(result[1], 1.98, epsilon = 1e-6);
276 assert_relative_eq!(result[2], 2.97, epsilon = 1e-6);
277 assert_relative_eq!(result[3], 3.96, epsilon = 1e-6);
278 }
279
280 #[test]
281 fn test_simd_sgd_momentum() {
282 let params = Array1::from_vec(vec![1.0f32, 2.0, 3.0, 4.0]);
283 let gradients = Array1::from_vec(vec![0.1, 0.2, 0.3, 0.4]);
284
285 let mut optimizer = SimdSGD::new_with_config(0.1, 0.9, 0.0);
286
287 let result1 = optimizer
289 .step(¶ms, &gradients)
290 .expect("optimizer.step succeeds in test_simd_sgd_momentum");
291
292 let result2 = optimizer
294 .step(&result1, &gradients)
295 .expect("optimizer.step succeeds in test_simd_sgd_momentum");
296
297 assert!(result2[0] < result1[0]);
299 }
300
301 #[test]
302 fn test_simd_sgd_weight_decay() {
303 let params = Array1::from_vec(vec![1.0f32, 2.0, 3.0, 4.0]);
304 let gradients = Array1::from_vec(vec![0.1, 0.2, 0.3, 0.4]);
305
306 let mut optimizer = SimdSGD::new_with_config(0.1, 0.0, 0.01);
307 let result = optimizer
308 .step(¶ms, &gradients)
309 .expect("optimizer.step succeeds in test_simd_sgd_weight_decay");
310
311 let expected_grad = 0.1 + 0.01 * 1.0;
313 assert_relative_eq!(result[0], 1.0 - 0.1 * expected_grad, epsilon = 1e-6);
314 }
315
316 #[test]
317 fn test_simd_sgd_large_array() {
318 let size = 1000;
320 let params: Array1<f32> = Array1::from_vec((0..size).map(|i| i as f32).collect());
321 let gradients: Array1<f32> = Array1::from_elem(size, 0.1);
322
323 let mut optimizer = SimdSGD::new(0.01);
324 let result = optimizer
325 .step(¶ms, &gradients)
326 .expect("optimizer.step succeeds in test_simd_sgd_large_array");
327
328 for i in 0..size {
329 assert_relative_eq!(result[i], (i as f32) - 0.01 * 0.1, epsilon = 1e-6);
330 }
331 }
332
333 #[test]
334 fn test_simd_sgd_f64() {
335 let params = Array1::from_vec(vec![1.0f64, 2.0, 3.0, 4.0]);
336 let gradients = Array1::from_vec(vec![0.1, 0.2, 0.3, 0.4]);
337
338 let mut optimizer = SimdSGD::new(0.1);
339 let result = optimizer
340 .step(¶ms, &gradients)
341 .expect("optimizer.step succeeds in test_simd_sgd_f64");
342
343 assert_relative_eq!(result[0], 0.99, epsilon = 1e-10);
344 assert_relative_eq!(result[1], 1.98, epsilon = 1e-10);
345 assert_relative_eq!(result[2], 2.97, epsilon = 1e-10);
346 assert_relative_eq!(result[3], 3.96, epsilon = 1e-10);
347 }
348
349 #[test]
350 fn test_simd_sgd_reset() {
351 let params = Array1::from_vec(vec![1.0f32, 2.0, 3.0, 4.0]);
352 let gradients = Array1::from_vec(vec![0.1, 0.2, 0.3, 0.4]);
353
354 let mut optimizer = SimdSGD::new_with_config(0.1, 0.9, 0.0);
355
356 let _ = optimizer
358 .step(¶ms, &gradients)
359 .expect("optimizer.step succeeds in test_simd_sgd_reset");
360 assert!(optimizer.velocity.is_some());
361
362 optimizer.reset();
364 assert!(optimizer.velocity.is_none());
365 }
366}