1use scirs2_core::ndarray::{Array1, ArrayView1};
10use scirs2_core::numeric::Float;
11use scirs2_core::simd_ops::SimdUnifiedOps;
12
13pub trait SimdOptimizer<T: Float> {
18 fn simd_sgd_update(
30 params: &ArrayView1<T>,
31 gradients: &ArrayView1<T>,
32 learning_rate: T,
33 ) -> Array1<T>;
34
35 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 fn simd_adam_first_moment(m: &ArrayView1<T>, gradients: &ArrayView1<T>, beta1: T) -> Array1<T>;
73
74 fn simd_adam_second_moment(v: &ArrayView1<T>, gradients: &ArrayView1<T>, beta2: T)
88 -> Array1<T>;
89
90 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 fn simd_weight_decay(
127 gradients: &ArrayView1<T>,
128 params: &ArrayView1<T>,
129 weight_decay: T,
130 ) -> Array1<T>;
131
132 fn simd_gradient_norm(gradients: &ArrayView1<T>) -> T;
142}
143
144impl 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 if params.len() >= 16 {
153 let scaled_grads = f32::simd_scalar_mul(gradients, learning_rate);
155 f32::simd_sub(params, &scaled_grads.view())
156 } else {
157 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 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 let new_params = f32::simd_sub(params, &new_velocity.view());
182
183 (new_params, new_velocity)
184 } else {
185 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 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 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 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 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 let v_hat_sqrt: Array1<f32> = v_hat.iter().map(|&v| v.sqrt() + epsilon).collect();
252
253 let step = f32::simd_div(m_hat, &v_hat_sqrt.view());
255
256 let scaled_step = f32::simd_scalar_mul(&step.view(), learning_rate);
258
259 f32::simd_sub(params, &scaled_step.view())
261 } else {
262 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 let scaled_params = f32::simd_scalar_mul(params, weight_decay);
279 f32::simd_add(gradients, &scaled_params.view())
280 } else {
281 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 f32::simd_dot(gradients, gradients).sqrt()
294 } else {
295 gradients.iter().map(|&x| x * x).sum::<f32>().sqrt()
297 }
298 }
299}
300
301impl 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 let scaled_grads = f64::simd_scalar_mul(gradients, learning_rate);
311 f64::simd_sub(params, &scaled_grads.view())
312 } else {
313 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 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 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 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 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 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 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 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 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 let scaled_params = f64::simd_scalar_mul(params, weight_decay);
422 f64::simd_add(gradients, &scaled_params.view())
423 } else {
424 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 f64::simd_dot(gradients, gradients).sqrt()
437 } else {
438 gradients.iter().map(|&x| x * x).sum::<f64>().sqrt()
440 }
441 }
442}
443
444pub fn should_use_simd(size: usize, dtype_size: usize) -> bool {
455 let min_simd_size = match dtype_size {
457 4 => 16, 8 => 8, _ => usize::MAX, };
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(¶ms.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(¶ms.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 ¶ms.view(),
508 &gradients.view(),
509 &velocity.view(),
510 learning_rate,
511 momentum,
512 );
513
514 assert_relative_eq!(new_velocity[0], 0.9 * 0.01 + 0.1 * 0.1, epsilon = 1e-6);
516
517 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(), ¶ms.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 assert!(!should_use_simd(8, 4)); assert!(should_use_simd(16, 4)); assert!(should_use_simd(100, 4)); assert!(!should_use_simd(4, 8)); assert!(should_use_simd(8, 8)); assert!(should_use_simd(100, 8)); }
579
580 #[test]
581 fn test_simd_large_array() {
582 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(¶ms.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}