Skip to main content

optirs_core/optimizers/
sgd_simd.rs

1//! SIMD-accelerated SGD optimizer
2//!
3//! This module provides a SIMD-optimized implementation of Stochastic Gradient Descent
4//! for 1D parameter arrays using scirs2_core's SimdUnifiedOps.
5
6use 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/// SIMD-accelerated Stochastic Gradient Descent optimizer
15///
16/// This is a specialized version of SGD optimized for 1D arrays using SIMD operations.
17/// For maximum performance, use this when working with flattened parameter vectors.
18///
19/// Formula:
20/// v_t = momentum * v_{t-1} + learning_rate * (gradient + weight_decay * param)
21/// param_t = param_{t-1} - v_t
22///
23/// # Performance
24///
25/// This implementation uses SIMD instructions (AVX2/SSE/NEON) for:
26/// - Parameter updates
27/// - Momentum computation
28/// - Weight decay application
29///
30/// Expected speedup: 2-4x over scalar implementation for large parameter arrays
31///
32/// # Examples
33///
34/// ```
35/// use scirs2_core::ndarray::Array1;
36/// use optirs_core::optimizers::{SimdSGD, Optimizer};
37///
38/// // Initialize parameters and gradients
39/// let params = Array1::zeros(1000);
40/// let gradients = Array1::from_elem(1000, 0.1);
41///
42/// // Create SIMD-accelerated SGD optimizer
43/// let mut optimizer = SimdSGD::new(0.01);
44/// optimizer.set_momentum(0.9);
45///
46/// // Update parameters with SIMD acceleration
47/// let new_params = optimizer.step(&params, &gradients).expect("optimizer.step succeeds");
48/// ```
49#[derive(Debug, Clone)]
50pub struct SimdSGD<A: Float> {
51    /// Learning rate
52    learning_rate: A,
53    /// Momentum factor (0.0 means no momentum)
54    momentum: A,
55    /// Weight decay factor (L2 regularization)
56    weight_decay: A,
57    /// Velocity (momentum state)
58    velocity: Option<Array1<A>>,
59}
60
61impl<A: Float> SimdSGD<A> {
62    /// Creates a new SIMD-accelerated SGD optimizer
63    ///
64    /// # Arguments
65    ///
66    /// * `learning_rate` - The learning rate for parameter updates
67    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    /// Creates a new SIMD SGD optimizer with full configuration
77    ///
78    /// # Arguments
79    ///
80    /// * `learning_rate` - The learning rate for parameter updates
81    /// * `momentum` - The momentum factor (0.0 means no momentum)
82    /// * `weight_decay` - The weight decay factor (L2 regularization)
83    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    /// Sets the momentum factor
93    pub fn set_momentum(&mut self, momentum: A) -> &mut Self {
94        self.momentum = momentum;
95        self
96    }
97
98    /// Builder method to set momentum and return self
99    pub fn with_momentum(mut self, momentum: A) -> Self {
100        self.momentum = momentum;
101        self
102    }
103
104    /// Gets the current momentum factor
105    pub fn get_momentum(&self) -> A {
106        self.momentum
107    }
108
109    /// Gets the current learning rate
110    pub fn learning_rate(&self) -> A {
111        self.learning_rate
112    }
113
114    /// Sets the weight decay factor
115    pub fn set_weight_decay(&mut self, weight_decay: A) -> &mut Self {
116        self.weight_decay = weight_decay;
117        self
118    }
119
120    /// Builder method to set weight decay and return self
121    pub fn with_weight_decay(mut self, weight_decay: A) -> Self {
122        self.weight_decay = weight_decay;
123        self
124    }
125
126    /// Gets the current weight decay factor
127    pub fn get_weight_decay(&self) -> A {
128        self.weight_decay
129    }
130
131    /// Resets the optimizer state
132    pub fn reset(&mut self) {
133        self.velocity = None;
134    }
135}
136
137// Specialized SIMD implementation for f32
138impl Optimizer<f32, scirs2_core::ndarray::Ix1> for SimdSGD<f32> {
139    fn step(&mut self, params: &Array1<f32>, gradients: &Array1<f32>) -> Result<Array1<f32>> {
140        // Validate shapes
141        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        // Apply weight decay if needed
153        let adjusted_gradients = if self.weight_decay > 0.0 {
154            f32::simd_weight_decay(&gradients_view, &params_view, self.weight_decay)
155        } else {
156            gradients.to_owned()
157        };
158
159        // Initialize velocity if this is the first step
160        let velocity = self
161            .velocity
162            .get_or_insert_with(|| Array1::zeros(params.len()));
163
164        // Ensure velocity has correct dimensions
165        if velocity.len() != params.len() {
166            *velocity = Array1::zeros(params.len());
167        }
168
169        // Compute update using SIMD operations
170        let new_params = if self.momentum > 0.0 {
171            // SIMD-accelerated momentum update
172            let (updated_params, updated_velocity) = f32::simd_momentum_update(
173                &params_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            // SIMD-accelerated vanilla SGD
183            f32::simd_sgd_update(&params_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
198// Specialized SIMD implementation for f64
199impl Optimizer<f64, scirs2_core::ndarray::Ix1> for SimdSGD<f64> {
200    fn step(&mut self, params: &Array1<f64>, gradients: &Array1<f64>) -> Result<Array1<f64>> {
201        // Validate shapes
202        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        // Apply weight decay if needed
214        let adjusted_gradients = if self.weight_decay > 0.0 {
215            f64::simd_weight_decay(&gradients_view, &params_view, self.weight_decay)
216        } else {
217            gradients.to_owned()
218        };
219
220        // Initialize velocity if this is the first step
221        let velocity = self
222            .velocity
223            .get_or_insert_with(|| Array1::zeros(params.len()));
224
225        // Ensure velocity has correct dimensions
226        if velocity.len() != params.len() {
227            *velocity = Array1::zeros(params.len());
228        }
229
230        // Compute update using SIMD operations
231        let new_params = if self.momentum > 0.0 {
232            // SIMD-accelerated momentum update
233            let (updated_params, updated_velocity) = f64::simd_momentum_update(
234                &params_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            // SIMD-accelerated vanilla SGD
244            f64::simd_sgd_update(&params_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(&params, &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        // First step
288        let result1 = optimizer
289            .step(&params, &gradients)
290            .expect("optimizer.step succeeds in test_simd_sgd_momentum");
291
292        // Second step - should show momentum effect
293        let result2 = optimizer
294            .step(&result1, &gradients)
295            .expect("optimizer.step succeeds in test_simd_sgd_momentum");
296
297        // With momentum, the second step should move further
298        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(&params, &gradients)
309            .expect("optimizer.step succeeds in test_simd_sgd_weight_decay");
310
311        // Weight decay should reduce parameters more than vanilla SGD
312        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        // Test with large array to ensure SIMD path is taken
319        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(&params, &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(&params, &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        // Take a step to initialize velocity
357        let _ = optimizer
358            .step(&params, &gradients)
359            .expect("optimizer.step succeeds in test_simd_sgd_reset");
360        assert!(optimizer.velocity.is_some());
361
362        // Reset should clear velocity
363        optimizer.reset();
364        assert!(optimizer.velocity.is_none());
365    }
366}