Skip to main content

optirs_core/
simd_optimizer.rs

1//! SIMD-accelerated optimizer operations
2//!
3//! This module provides SIMD-optimized implementations of common optimizer
4//! operations using scirs2_core's SimdUnifiedOps infrastructure.
5//!
6//! The module automatically selects the best SIMD backend available on the
7//! target platform (AVX2, SSE, NEON, or scalar fallback).
8
9use scirs2_core::ndarray::{Array1, ArrayView1};
10use scirs2_core::numeric::Float;
11use scirs2_core::simd_ops::SimdUnifiedOps;
12
13/// Trait for SIMD-accelerated optimizer operations
14///
15/// This trait provides high-performance implementations of common
16/// operations found in optimization algorithms.
17pub trait SimdOptimizer<T: Float> {
18    /// SIMD-accelerated parameter update: params - learning_rate * gradient
19    ///
20    /// # Arguments
21    ///
22    /// * `params` - Parameter array
23    /// * `gradients` - Gradient array
24    /// * `learning_rate` - Learning rate scalar
25    ///
26    /// # Returns
27    ///
28    /// Updated parameters
29    fn simd_sgd_update(
30        params: &ArrayView1<T>,
31        gradients: &ArrayView1<T>,
32        learning_rate: T,
33    ) -> Array1<T>;
34
35    /// SIMD-accelerated momentum update
36    ///
37    /// velocity = momentum * velocity + learning_rate * gradient
38    /// params = params - velocity
39    ///
40    /// # Arguments
41    ///
42    /// * `params` - Parameter array
43    /// * `gradients` - Gradient array
44    /// * `velocity` - Velocity array (momentum state)
45    /// * `learning_rate` - Learning rate scalar
46    /// * `momentum` - Momentum coefficient
47    ///
48    /// # Returns
49    ///
50    /// Tuple of (updated_params, updated_velocity)
51    fn simd_momentum_update(
52        params: &ArrayView1<T>,
53        gradients: &ArrayView1<T>,
54        velocity: &ArrayView1<T>,
55        learning_rate: T,
56        momentum: T,
57    ) -> (Array1<T>, Array1<T>);
58
59    /// SIMD-accelerated Adam first moment update
60    ///
61    /// m = beta1 * m + (1 - beta1) * gradient
62    ///
63    /// # Arguments
64    ///
65    /// * `m` - First moment array
66    /// * `gradients` - Gradient array
67    /// * `beta1` - Exponential decay rate for first moment
68    ///
69    /// # Returns
70    ///
71    /// Updated first moment
72    fn simd_adam_first_moment(m: &ArrayView1<T>, gradients: &ArrayView1<T>, beta1: T) -> Array1<T>;
73
74    /// SIMD-accelerated Adam second moment update
75    ///
76    /// v = beta2 * v + (1 - beta2) * gradient^2
77    ///
78    /// # Arguments
79    ///
80    /// * `v` - Second moment array
81    /// * `gradients` - Gradient array
82    /// * `beta2` - Exponential decay rate for second moment
83    ///
84    /// # Returns
85    ///
86    /// Updated second moment
87    fn simd_adam_second_moment(v: &ArrayView1<T>, gradients: &ArrayView1<T>, beta2: T)
88        -> Array1<T>;
89
90    /// SIMD-accelerated Adam parameter update
91    ///
92    /// params = params - learning_rate * m_hat / (sqrt(v_hat) + epsilon)
93    ///
94    /// # Arguments
95    ///
96    /// * `params` - Parameter array
97    /// * `m_hat` - Bias-corrected first moment
98    /// * `v_hat` - Bias-corrected second moment
99    /// * `learning_rate` - Learning rate scalar
100    /// * `epsilon` - Small constant for numerical stability
101    ///
102    /// # Returns
103    ///
104    /// Updated parameters
105    fn simd_adam_update(
106        params: &ArrayView1<T>,
107        m_hat: &ArrayView1<T>,
108        v_hat: &ArrayView1<T>,
109        learning_rate: T,
110        epsilon: T,
111    ) -> Array1<T>;
112
113    /// SIMD-accelerated weight decay application
114    ///
115    /// gradients = gradients + weight_decay * params
116    ///
117    /// # Arguments
118    ///
119    /// * `gradients` - Gradient array
120    /// * `params` - Parameter array
121    /// * `weight_decay` - Weight decay coefficient
122    ///
123    /// # Returns
124    ///
125    /// Gradients with weight decay applied
126    fn simd_weight_decay(
127        gradients: &ArrayView1<T>,
128        params: &ArrayView1<T>,
129        weight_decay: T,
130    ) -> Array1<T>;
131
132    /// SIMD-accelerated gradient norm computation
133    ///
134    /// # Arguments
135    ///
136    /// * `gradients` - Gradient array
137    ///
138    /// # Returns
139    ///
140    /// L2 norm of gradients
141    fn simd_gradient_norm(gradients: &ArrayView1<T>) -> T;
142}
143
144/// Implementation of SIMD optimizer operations for f32
145impl SimdOptimizer<f32> for f32 {
146    fn simd_sgd_update(
147        params: &ArrayView1<f32>,
148        gradients: &ArrayView1<f32>,
149        learning_rate: f32,
150    ) -> Array1<f32> {
151        // Use SIMD for large arrays, scalar for small ones
152        if params.len() >= 16 {
153            // SIMD path: params - learning_rate * gradients
154            let scaled_grads = f32::simd_scalar_mul(gradients, learning_rate);
155            f32::simd_sub(params, &scaled_grads.view())
156        } else {
157            // Scalar path for small arrays
158            params
159                .iter()
160                .zip(gradients.iter())
161                .map(|(&p, &g)| p - learning_rate * g)
162                .collect()
163        }
164    }
165
166    fn simd_momentum_update(
167        params: &ArrayView1<f32>,
168        gradients: &ArrayView1<f32>,
169        velocity: &ArrayView1<f32>,
170        learning_rate: f32,
171        momentum: f32,
172    ) -> (Array1<f32>, Array1<f32>) {
173        if params.len() >= 16 {
174            // SIMD path
175            // velocity = momentum * velocity + learning_rate * gradient
176            let scaled_velocity = f32::simd_scalar_mul(velocity, momentum);
177            let scaled_gradients = f32::simd_scalar_mul(gradients, learning_rate);
178            let new_velocity = f32::simd_add(&scaled_velocity.view(), &scaled_gradients.view());
179
180            // params = params - velocity
181            let new_params = f32::simd_sub(params, &new_velocity.view());
182
183            (new_params, new_velocity)
184        } else {
185            // Scalar path
186            let new_velocity: Array1<f32> = velocity
187                .iter()
188                .zip(gradients.iter())
189                .map(|(&v, &g)| momentum * v + learning_rate * g)
190                .collect();
191
192            let new_params: Array1<f32> = params
193                .iter()
194                .zip(new_velocity.iter())
195                .map(|(&p, &v)| p - v)
196                .collect();
197
198            (new_params, new_velocity)
199        }
200    }
201
202    fn simd_adam_first_moment(
203        m: &ArrayView1<f32>,
204        gradients: &ArrayView1<f32>,
205        beta1: f32,
206    ) -> Array1<f32> {
207        if m.len() >= 16 {
208            // SIMD path: m = beta1 * m + (1 - beta1) * gradient
209            let scaled_m = f32::simd_scalar_mul(m, beta1);
210            let scaled_grads = f32::simd_scalar_mul(gradients, 1.0 - beta1);
211            f32::simd_add(&scaled_m.view(), &scaled_grads.view())
212        } else {
213            // Scalar path
214            m.iter()
215                .zip(gradients.iter())
216                .map(|(&m_val, &g)| beta1 * m_val + (1.0 - beta1) * g)
217                .collect()
218        }
219    }
220
221    fn simd_adam_second_moment(
222        v: &ArrayView1<f32>,
223        gradients: &ArrayView1<f32>,
224        beta2: f32,
225    ) -> Array1<f32> {
226        if v.len() >= 16 {
227            // SIMD path: v = beta2 * v + (1 - beta2) * gradient^2
228            let scaled_v = f32::simd_scalar_mul(v, beta2);
229            let grad_squared = f32::simd_mul(gradients, gradients);
230            let scaled_grad_squared = f32::simd_scalar_mul(&grad_squared.view(), 1.0 - beta2);
231            f32::simd_add(&scaled_v.view(), &scaled_grad_squared.view())
232        } else {
233            // Scalar path
234            v.iter()
235                .zip(gradients.iter())
236                .map(|(&v_val, &g)| beta2 * v_val + (1.0 - beta2) * g * g)
237                .collect()
238        }
239    }
240
241    fn simd_adam_update(
242        params: &ArrayView1<f32>,
243        m_hat: &ArrayView1<f32>,
244        v_hat: &ArrayView1<f32>,
245        learning_rate: f32,
246        epsilon: f32,
247    ) -> Array1<f32> {
248        if params.len() >= 16 {
249            // SIMD path: params - learning_rate * m_hat / (sqrt(v_hat) + epsilon)
250            // Compute sqrt(v_hat) + epsilon
251            let v_hat_sqrt: Array1<f32> = v_hat.iter().map(|&v| v.sqrt() + epsilon).collect();
252
253            // Compute m_hat / (sqrt(v_hat) + epsilon)
254            let step = f32::simd_div(m_hat, &v_hat_sqrt.view());
255
256            // Scale by learning rate
257            let scaled_step = f32::simd_scalar_mul(&step.view(), learning_rate);
258
259            // Update parameters
260            f32::simd_sub(params, &scaled_step.view())
261        } else {
262            // Scalar path
263            params
264                .iter()
265                .zip(m_hat.iter().zip(v_hat.iter()))
266                .map(|(&p, (&m, &v))| p - learning_rate * m / (v.sqrt() + epsilon))
267                .collect()
268        }
269    }
270
271    fn simd_weight_decay(
272        gradients: &ArrayView1<f32>,
273        params: &ArrayView1<f32>,
274        weight_decay: f32,
275    ) -> Array1<f32> {
276        if gradients.len() >= 16 {
277            // SIMD path: gradients + weight_decay * params
278            let scaled_params = f32::simd_scalar_mul(params, weight_decay);
279            f32::simd_add(gradients, &scaled_params.view())
280        } else {
281            // Scalar path
282            gradients
283                .iter()
284                .zip(params.iter())
285                .map(|(&g, &p)| g + weight_decay * p)
286                .collect()
287        }
288    }
289
290    fn simd_gradient_norm(gradients: &ArrayView1<f32>) -> f32 {
291        if gradients.len() >= 16 {
292            // SIMD path using optimized dot product
293            f32::simd_dot(gradients, gradients).sqrt()
294        } else {
295            // Scalar path
296            gradients.iter().map(|&x| x * x).sum::<f32>().sqrt()
297        }
298    }
299}
300
301/// Implementation of SIMD optimizer operations for f64
302impl SimdOptimizer<f64> for f64 {
303    fn simd_sgd_update(
304        params: &ArrayView1<f64>,
305        gradients: &ArrayView1<f64>,
306        learning_rate: f64,
307    ) -> Array1<f64> {
308        if params.len() >= 8 {
309            // SIMD path
310            let scaled_grads = f64::simd_scalar_mul(gradients, learning_rate);
311            f64::simd_sub(params, &scaled_grads.view())
312        } else {
313            // Scalar path
314            params
315                .iter()
316                .zip(gradients.iter())
317                .map(|(&p, &g)| p - learning_rate * g)
318                .collect()
319        }
320    }
321
322    fn simd_momentum_update(
323        params: &ArrayView1<f64>,
324        gradients: &ArrayView1<f64>,
325        velocity: &ArrayView1<f64>,
326        learning_rate: f64,
327        momentum: f64,
328    ) -> (Array1<f64>, Array1<f64>) {
329        if params.len() >= 8 {
330            // SIMD path
331            let scaled_velocity = f64::simd_scalar_mul(velocity, momentum);
332            let scaled_gradients = f64::simd_scalar_mul(gradients, learning_rate);
333            let new_velocity = f64::simd_add(&scaled_velocity.view(), &scaled_gradients.view());
334            let new_params = f64::simd_sub(params, &new_velocity.view());
335            (new_params, new_velocity)
336        } else {
337            // Scalar path
338            let new_velocity: Array1<f64> = velocity
339                .iter()
340                .zip(gradients.iter())
341                .map(|(&v, &g)| momentum * v + learning_rate * g)
342                .collect();
343            let new_params: Array1<f64> = params
344                .iter()
345                .zip(new_velocity.iter())
346                .map(|(&p, &v)| p - v)
347                .collect();
348            (new_params, new_velocity)
349        }
350    }
351
352    fn simd_adam_first_moment(
353        m: &ArrayView1<f64>,
354        gradients: &ArrayView1<f64>,
355        beta1: f64,
356    ) -> Array1<f64> {
357        if m.len() >= 8 {
358            // SIMD path
359            let scaled_m = f64::simd_scalar_mul(m, beta1);
360            let scaled_grads = f64::simd_scalar_mul(gradients, 1.0 - beta1);
361            f64::simd_add(&scaled_m.view(), &scaled_grads.view())
362        } else {
363            // Scalar path
364            m.iter()
365                .zip(gradients.iter())
366                .map(|(&m_val, &g)| beta1 * m_val + (1.0 - beta1) * g)
367                .collect()
368        }
369    }
370
371    fn simd_adam_second_moment(
372        v: &ArrayView1<f64>,
373        gradients: &ArrayView1<f64>,
374        beta2: f64,
375    ) -> Array1<f64> {
376        if v.len() >= 8 {
377            // SIMD path
378            let scaled_v = f64::simd_scalar_mul(v, beta2);
379            let grad_squared = f64::simd_mul(gradients, gradients);
380            let scaled_grad_squared = f64::simd_scalar_mul(&grad_squared.view(), 1.0 - beta2);
381            f64::simd_add(&scaled_v.view(), &scaled_grad_squared.view())
382        } else {
383            // Scalar path
384            v.iter()
385                .zip(gradients.iter())
386                .map(|(&v_val, &g)| beta2 * v_val + (1.0 - beta2) * g * g)
387                .collect()
388        }
389    }
390
391    fn simd_adam_update(
392        params: &ArrayView1<f64>,
393        m_hat: &ArrayView1<f64>,
394        v_hat: &ArrayView1<f64>,
395        learning_rate: f64,
396        epsilon: f64,
397    ) -> Array1<f64> {
398        if params.len() >= 8 {
399            // SIMD path
400            let v_hat_sqrt: Array1<f64> = v_hat.iter().map(|&v| v.sqrt() + epsilon).collect();
401            let step = f64::simd_div(m_hat, &v_hat_sqrt.view());
402            let scaled_step = f64::simd_scalar_mul(&step.view(), learning_rate);
403            f64::simd_sub(params, &scaled_step.view())
404        } else {
405            // Scalar path
406            params
407                .iter()
408                .zip(m_hat.iter().zip(v_hat.iter()))
409                .map(|(&p, (&m, &v))| p - learning_rate * m / (v.sqrt() + epsilon))
410                .collect()
411        }
412    }
413
414    fn simd_weight_decay(
415        gradients: &ArrayView1<f64>,
416        params: &ArrayView1<f64>,
417        weight_decay: f64,
418    ) -> Array1<f64> {
419        if gradients.len() >= 8 {
420            // SIMD path
421            let scaled_params = f64::simd_scalar_mul(params, weight_decay);
422            f64::simd_add(gradients, &scaled_params.view())
423        } else {
424            // Scalar path
425            gradients
426                .iter()
427                .zip(params.iter())
428                .map(|(&g, &p)| g + weight_decay * p)
429                .collect()
430        }
431    }
432
433    fn simd_gradient_norm(gradients: &ArrayView1<f64>) -> f64 {
434        if gradients.len() >= 8 {
435            // SIMD path
436            f64::simd_dot(gradients, gradients).sqrt()
437        } else {
438            // Scalar path
439            gradients.iter().map(|&x| x * x).sum::<f64>().sqrt()
440        }
441    }
442}
443
444/// Helper function to determine if SIMD should be used based on array size
445///
446/// # Arguments
447///
448/// * `size` - Size of the array
449/// * `dtype_size` - Size of the data type in bytes (4 for f32, 8 for f64)
450///
451/// # Returns
452///
453/// True if SIMD should be used, false otherwise
454pub fn should_use_simd(size: usize, dtype_size: usize) -> bool {
455    // Use SIMD for arrays with at least 16 f32 elements or 8 f64 elements
456    let min_simd_size = match dtype_size {
457        4 => 16,         // f32
458        8 => 8,          // f64
459        _ => usize::MAX, // Unknown type, don't use SIMD
460    };
461
462    size >= min_simd_size
463}
464
465#[cfg(test)]
466mod tests {
467    use super::*;
468    use approx::assert_relative_eq;
469
470    #[test]
471    fn test_simd_sgd_update_f32() {
472        let params = Array1::from_vec(vec![1.0f32, 2.0, 3.0, 4.0]);
473        let gradients = Array1::from_vec(vec![0.1, 0.2, 0.3, 0.4]);
474        let learning_rate = 0.1;
475
476        let result = f32::simd_sgd_update(&params.view(), &gradients.view(), learning_rate);
477
478        assert_relative_eq!(result[0], 0.99, epsilon = 1e-6);
479        assert_relative_eq!(result[1], 1.98, epsilon = 1e-6);
480        assert_relative_eq!(result[2], 2.97, epsilon = 1e-6);
481        assert_relative_eq!(result[3], 3.96, epsilon = 1e-6);
482    }
483
484    #[test]
485    fn test_simd_sgd_update_f64() {
486        let params = Array1::from_vec(vec![1.0f64, 2.0, 3.0, 4.0]);
487        let gradients = Array1::from_vec(vec![0.1, 0.2, 0.3, 0.4]);
488        let learning_rate = 0.1;
489
490        let result = f64::simd_sgd_update(&params.view(), &gradients.view(), learning_rate);
491
492        assert_relative_eq!(result[0], 0.99, epsilon = 1e-10);
493        assert_relative_eq!(result[1], 1.98, epsilon = 1e-10);
494        assert_relative_eq!(result[2], 2.97, epsilon = 1e-10);
495        assert_relative_eq!(result[3], 3.96, epsilon = 1e-10);
496    }
497
498    #[test]
499    fn test_simd_momentum_update() {
500        let params = Array1::from_vec(vec![1.0f32, 2.0, 3.0, 4.0]);
501        let gradients = Array1::from_vec(vec![0.1, 0.2, 0.3, 0.4]);
502        let velocity = Array1::from_vec(vec![0.01, 0.02, 0.03, 0.04]);
503        let learning_rate = 0.1;
504        let momentum = 0.9;
505
506        let (new_params, new_velocity) = f32::simd_momentum_update(
507            &params.view(),
508            &gradients.view(),
509            &velocity.view(),
510            learning_rate,
511            momentum,
512        );
513
514        // Check velocity: 0.9 * old_velocity + 0.1 * gradient
515        assert_relative_eq!(new_velocity[0], 0.9 * 0.01 + 0.1 * 0.1, epsilon = 1e-6);
516
517        // Check params: old_params - new_velocity
518        assert_relative_eq!(new_params[0], 1.0 - new_velocity[0], epsilon = 1e-6);
519    }
520
521    #[test]
522    fn test_simd_adam_first_moment() {
523        let m = Array1::from_vec(vec![0.01f32, 0.02, 0.03, 0.04]);
524        let gradients = Array1::from_vec(vec![0.1, 0.2, 0.3, 0.4]);
525        let beta1 = 0.9;
526
527        let result = f32::simd_adam_first_moment(&m.view(), &gradients.view(), beta1);
528
529        assert_relative_eq!(result[0], 0.9 * 0.01 + 0.1 * 0.1, epsilon = 1e-6);
530        assert_relative_eq!(result[1], 0.9 * 0.02 + 0.1 * 0.2, epsilon = 1e-6);
531    }
532
533    #[test]
534    fn test_simd_adam_second_moment() {
535        let v = Array1::from_vec(vec![0.001f32, 0.002, 0.003, 0.004]);
536        let gradients = Array1::from_vec(vec![0.1, 0.2, 0.3, 0.4]);
537        let beta2 = 0.999;
538
539        let result = f32::simd_adam_second_moment(&v.view(), &gradients.view(), beta2);
540
541        assert_relative_eq!(result[0], 0.999 * 0.001 + 0.001 * 0.1 * 0.1, epsilon = 1e-6);
542    }
543
544    #[test]
545    fn test_simd_weight_decay() {
546        let gradients = Array1::from_vec(vec![0.1f32, 0.2, 0.3, 0.4]);
547        let params = Array1::from_vec(vec![1.0, 2.0, 3.0, 4.0]);
548        let weight_decay = 0.01;
549
550        let result = f32::simd_weight_decay(&gradients.view(), &params.view(), weight_decay);
551
552        assert_relative_eq!(result[0], 0.1 + 0.01 * 1.0, epsilon = 1e-6);
553        assert_relative_eq!(result[1], 0.2 + 0.01 * 2.0, epsilon = 1e-6);
554    }
555
556    #[test]
557    fn test_simd_gradient_norm() {
558        let gradients = Array1::from_vec(vec![3.0f32, 4.0]);
559        let norm = f32::simd_gradient_norm(&gradients.view());
560        assert_relative_eq!(norm, 5.0, epsilon = 1e-6);
561
562        let gradients_f64 = Array1::from_vec(vec![3.0f64, 4.0]);
563        let norm_f64 = f64::simd_gradient_norm(&gradients_f64.view());
564        assert_relative_eq!(norm_f64, 5.0, epsilon = 1e-10);
565    }
566
567    #[test]
568    fn test_should_use_simd() {
569        // f32 tests
570        assert!(!should_use_simd(8, 4)); // Too small for f32 SIMD
571        assert!(should_use_simd(16, 4)); // Exactly at threshold
572        assert!(should_use_simd(100, 4)); // Large enough
573
574        // f64 tests
575        assert!(!should_use_simd(4, 8)); // Too small for f64 SIMD
576        assert!(should_use_simd(8, 8)); // Exactly at threshold
577        assert!(should_use_simd(100, 8)); // Large enough
578    }
579
580    #[test]
581    fn test_simd_large_array() {
582        // Test with a large array to ensure SIMD path is taken
583        let size = 1000;
584        let params: Array1<f32> = Array1::from_vec((0..size).map(|i| i as f32).collect());
585        let gradients: Array1<f32> = Array1::from_vec(vec![0.1; size]);
586        let learning_rate = 0.01;
587
588        let result = f32::simd_sgd_update(&params.view(), &gradients.view(), learning_rate);
589
590        for i in 0..size {
591            assert_relative_eq!(result[i], (i as f32) - learning_rate * 0.1, epsilon = 1e-6);
592        }
593    }
594}