1#[inline(always)]
2pub(crate) fn dispatch<R>(f: impl FnOnce() -> R) -> R {
3 #[cfg(feature = "simd")]
4 {
5 pulp::Arch::new().dispatch(f)
6 }
7 #[cfg(not(feature = "simd"))]
8 {
9 f()
10 }
11}
12
13#[inline(always)]
14pub(crate) fn dispatch_if_large<R>(len: usize, f: impl FnOnce() -> R) -> R {
15 if len >= 64 {
18 dispatch(f)
19 } else {
20 f()
21 }
22}
23
24#[cfg(feature = "simd")]
25#[inline(always)]
26unsafe fn cast_slice<T, U>(src: &[T]) -> &[U] {
27 debug_assert_eq!(std::mem::size_of::<T>(), std::mem::size_of::<U>());
28 unsafe { std::slice::from_raw_parts(src.as_ptr().cast::<U>(), src.len()) }
29}
30
31#[cfg(feature = "simd")]
32macro_rules! impl_simd_mul_ptr {
33 (
34 $mul_into:ident,
35 $ty:ty,
36 $split:ident,
37 $split_uninit:ident,
38 $load:ident,
39 $mask:ident,
40 $store_ptr:ident,
41 $mul:ident,
42 $mask_scale:expr
43 ) => {
44 unsafe fn $mul_into(dst: *mut $ty, len: usize, a: &[$ty], b: &[$ty]) {
45 struct Mul<'a> {
46 dst: *mut $ty,
47 len: usize,
48 a: &'a [$ty],
49 b: &'a [$ty],
50 }
51
52 impl<'a> pulp::WithSimd for Mul<'a> {
53 type Output = ();
54
55 #[inline(always)]
56 fn with_simd<S: pulp::Simd>(self, simd: S) -> Self::Output {
57 debug_assert_eq!(self.len, self.a.len());
58 debug_assert_eq!(self.len, self.b.len());
59
60 let output = unsafe {
63 core::slice::from_raw_parts_mut(
64 self.dst.cast::<core::mem::MaybeUninit<$ty>>(),
65 self.len,
66 )
67 };
68 let (dest, tail) = S::$split_uninit(output);
69 let (a, a_tail) = S::$split(self.a);
70 let (b, b_tail) = S::$split(self.b);
71 for ((dst, &a), &b) in dest.iter_mut().zip(a).zip(b) {
74 dst.write(simd.$mul(a, b));
75 }
76 if !tail.is_empty() {
77 let va = simd.$load(a_tail);
78 let vb = simd.$load(b_tail);
79 unsafe {
81 simd.$store_ptr(
82 simd.$mask(0, (tail.len() * $mask_scale) as _),
83 tail.as_mut_ptr().cast::<$ty>(),
84 simd.$mul(va, vb),
85 );
86 }
87 }
88 }
89 }
90
91 pulp::Arch::new().dispatch(Mul { dst, len, a, b });
92 }
93 };
94}
95
96#[cfg(feature = "simd")]
97impl_simd_mul_ptr!(
98 simd_mul_f32_into,
99 f32,
100 as_simd_f32s,
101 as_uninit_mut_simd_f32s,
102 partial_load_f32s,
103 mask_between_m32s,
104 mask_store_ptr_f32s,
105 mul_f32s,
106 1
107);
108
109#[cfg(feature = "simd")]
110impl_simd_mul_ptr!(
111 simd_mul_f64_into,
112 f64,
113 as_simd_f64s,
114 as_uninit_mut_simd_f64s,
115 partial_load_f64s,
116 mask_between_m64s,
117 mask_store_ptr_f64s,
118 mul_f64s,
119 1
120);
121
122#[cfg(feature = "simd")]
123impl_simd_mul_ptr!(
124 simd_mul_c32_into,
125 num_complex::Complex32,
126 as_simd_c32s,
127 as_uninit_mut_simd_c32s,
128 partial_load_c32s,
129 mask_between_m32s,
130 mask_store_ptr_c32s,
131 mul_e_c32s,
132 2
133);
134
135#[cfg(feature = "simd")]
136impl_simd_mul_ptr!(
137 simd_mul_c64_into,
138 num_complex::Complex64,
139 as_simd_c64s,
140 as_uninit_mut_simd_c64s,
141 partial_load_c64s,
142 mask_between_m64s,
143 mask_store_ptr_c64s,
144 mul_e_c64s,
145 2
146);
147
148#[cfg(feature = "simd")]
149#[inline]
150#[cfg_attr(not(test), allow(dead_code))]
151pub(crate) fn try_mul_contiguous<D: 'static, A: 'static, B: 'static>(
152 dst: &mut [D],
153 a: &[A],
154 b: &[B],
155) -> bool {
156 unsafe { try_mul_contiguous_ptr(dst.as_mut_ptr(), dst.len(), a, b) }
157}
158
159#[cfg(feature = "simd")]
160#[inline]
161pub(crate) unsafe fn try_mul_contiguous_ptr<D: 'static, A: 'static, B: 'static>(
162 dst: *mut D,
163 len: usize,
164 a: &[A],
165 b: &[B],
166) -> bool {
167 use std::any::TypeId;
168
169 macro_rules! try_same_type {
170 ($ty:ty, $mul_into:ident) => {
171 if TypeId::of::<D>() == TypeId::of::<$ty>()
172 && TypeId::of::<A>() == TypeId::of::<$ty>()
173 && TypeId::of::<B>() == TypeId::of::<$ty>()
174 {
175 unsafe { $mul_into(dst.cast::<$ty>(), len, cast_slice(a), cast_slice(b)) };
176 return true;
177 }
178 };
179 }
180
181 try_same_type!(f64, simd_mul_f64_into);
182 try_same_type!(f32, simd_mul_f32_into);
183 try_same_type!(num_complex::Complex64, simd_mul_c64_into);
184 try_same_type!(num_complex::Complex32, simd_mul_c32_into);
185
186 false
187}
188
189#[cfg(not(feature = "simd"))]
190#[inline]
191#[cfg_attr(not(test), allow(dead_code))]
192pub(crate) fn try_mul_contiguous<D: 'static, A: 'static, B: 'static>(
193 _dst: &mut [D],
194 _a: &[A],
195 _b: &[B],
196) -> bool {
197 false
198}
199
200#[cfg(not(feature = "simd"))]
201#[inline]
202pub(crate) unsafe fn try_mul_contiguous_ptr<D: 'static, A: 'static, B: 'static>(
203 _dst: *mut D,
204 _len: usize,
205 _a: &[A],
206 _b: &[B],
207) -> bool {
208 false
209}
210
211#[cfg(feature = "parallel")]
212const TRANSPOSE_TILE: usize = 8;
213
214#[cfg(feature = "parallel")]
215unsafe fn mul_transposed_scalar_rhs_source_contiguous<T>(
216 dst: *mut T,
217 src: *const T,
218 scalar: T,
219 inner_len: usize,
220 row_len: usize,
221 src_fast_stride: isize,
222 src_row_stride: isize,
223) where
224 T: Copy + std::ops::Mul<Output = T>,
225{
226 let mut inner0 = 0usize;
227
228 while inner0 < inner_len {
229 let inner_count = TRANSPOSE_TILE.min(inner_len - inner0);
230 let mut row0 = 0usize;
231
232 while row0 < row_len {
233 let row_count = TRANSPOSE_TILE.min(row_len - row0);
234
235 for inner in 0..inner_count {
236 let src_base = src.offset((inner0 + inner) as isize * src_fast_stride);
237 for row in 0..row_count {
238 let row_index = row0 + row;
239 let src_offset = row_index as isize * src_row_stride;
240 *dst.add(row_index * inner_len + inner0 + inner) =
241 *src_base.offset(src_offset) * scalar;
242 }
243 }
244
245 row0 += TRANSPOSE_TILE;
246 }
247
248 inner0 += TRANSPOSE_TILE;
249 }
250}
251
252#[cfg(feature = "parallel")]
253unsafe fn mul_transposed_scalar_rhs_dst_contiguous<T>(
254 dst: *mut T,
255 src: *const T,
256 scalar: T,
257 inner_len: usize,
258 row_len: usize,
259 src_fast_stride: isize,
260 src_row_stride: isize,
261) where
262 T: Copy + std::ops::Mul<Output = T>,
263{
264 for row in 0..row_len {
265 let dst_base = dst.add(row * inner_len);
266 let src_base = src.offset(row as isize * src_row_stride);
267 for inner in 0..inner_len {
268 *dst_base.add(inner) = *src_base.offset(inner as isize * src_fast_stride) * scalar;
269 }
270 }
271}
272
273#[cfg(feature = "parallel")]
274unsafe fn mul_transposed_scalar_lhs_source_contiguous<T>(
275 dst: *mut T,
276 scalar: T,
277 src: *const T,
278 inner_len: usize,
279 row_len: usize,
280 src_fast_stride: isize,
281 src_row_stride: isize,
282) where
283 T: Copy + std::ops::Mul<Output = T>,
284{
285 let mut inner0 = 0usize;
286
287 while inner0 < inner_len {
288 let inner_count = TRANSPOSE_TILE.min(inner_len - inner0);
289 let mut row0 = 0usize;
290
291 while row0 < row_len {
292 let row_count = TRANSPOSE_TILE.min(row_len - row0);
293
294 for inner in 0..inner_count {
295 let src_base = src.offset((inner0 + inner) as isize * src_fast_stride);
296 for row in 0..row_count {
297 let row_index = row0 + row;
298 let src_offset = row_index as isize * src_row_stride;
299 *dst.add(row_index * inner_len + inner0 + inner) =
300 scalar * *src_base.offset(src_offset);
301 }
302 }
303
304 row0 += TRANSPOSE_TILE;
305 }
306
307 inner0 += TRANSPOSE_TILE;
308 }
309}
310
311#[cfg(feature = "parallel")]
312unsafe fn mul_transposed_scalar_lhs_dst_contiguous<T>(
313 dst: *mut T,
314 scalar: T,
315 src: *const T,
316 inner_len: usize,
317 row_len: usize,
318 src_fast_stride: isize,
319 src_row_stride: isize,
320) where
321 T: Copy + std::ops::Mul<Output = T>,
322{
323 for row in 0..row_len {
324 let dst_base = dst.add(row * inner_len);
325 let src_base = src.offset(row as isize * src_row_stride);
326 for inner in 0..inner_len {
327 *dst_base.add(inner) = scalar * *src_base.offset(inner as isize * src_fast_stride);
328 }
329 }
330}
331
332#[cfg(feature = "parallel")]
333#[inline(always)]
334pub(crate) unsafe fn mul_transposed_scalar_rhs_2d_typed<T>(
335 dst: *mut T,
336 src: *const T,
337 scalar: T,
338 inner_len: usize,
339 row_len: usize,
340 src_fast_stride: isize,
341 src_row_stride: isize,
342) where
343 T: Copy + std::ops::Mul<Output = T>,
344{
345 if row_len > inner_len {
346 unsafe {
347 mul_transposed_scalar_rhs_source_contiguous(
348 dst,
349 src,
350 scalar,
351 inner_len,
352 row_len,
353 src_fast_stride,
354 src_row_stride,
355 );
356 }
357 } else {
358 unsafe {
359 mul_transposed_scalar_rhs_dst_contiguous(
360 dst,
361 src,
362 scalar,
363 inner_len,
364 row_len,
365 src_fast_stride,
366 src_row_stride,
367 );
368 }
369 }
370}
371
372#[cfg(feature = "parallel")]
373#[inline(always)]
374pub(crate) unsafe fn mul_transposed_scalar_lhs_2d_typed<T>(
375 dst: *mut T,
376 scalar: T,
377 src: *const T,
378 inner_len: usize,
379 row_len: usize,
380 src_fast_stride: isize,
381 src_row_stride: isize,
382) where
383 T: Copy + std::ops::Mul<Output = T>,
384{
385 if row_len > inner_len {
386 unsafe {
387 mul_transposed_scalar_lhs_source_contiguous(
388 dst,
389 scalar,
390 src,
391 inner_len,
392 row_len,
393 src_fast_stride,
394 src_row_stride,
395 );
396 }
397 } else {
398 unsafe {
399 mul_transposed_scalar_lhs_dst_contiguous(
400 dst,
401 scalar,
402 src,
403 inner_len,
404 row_len,
405 src_fast_stride,
406 src_row_stride,
407 );
408 }
409 }
410}
411
412#[inline]
413#[cfg(feature = "parallel")]
414pub(crate) unsafe fn try_mul_transposed_scalar_rhs_2d<D: 'static, A: 'static, B: 'static>(
415 dst: *mut D,
416 src: *const A,
417 scalar: *const B,
418 inner_len: usize,
419 row_len: usize,
420 src_fast_stride: isize,
421 src_row_stride: isize,
422) -> bool {
423 use std::any::TypeId;
424
425 macro_rules! try_same_type {
426 ($ty:ty) => {
427 if TypeId::of::<D>() == TypeId::of::<$ty>()
428 && TypeId::of::<A>() == TypeId::of::<$ty>()
429 && TypeId::of::<B>() == TypeId::of::<$ty>()
430 {
431 unsafe {
432 mul_transposed_scalar_rhs_2d_typed(
433 dst.cast::<$ty>(),
434 src.cast::<$ty>(),
435 *scalar.cast::<$ty>(),
436 inner_len,
437 row_len,
438 src_fast_stride,
439 src_row_stride,
440 );
441 }
442 return true;
443 }
444 };
445 }
446
447 try_same_type!(f64);
448 try_same_type!(f32);
449 try_same_type!(num_complex::Complex64);
450 try_same_type!(num_complex::Complex32);
451
452 false
453}
454
455#[inline]
456#[cfg(feature = "parallel")]
457pub(crate) unsafe fn try_mul_transposed_scalar_lhs_2d<D: 'static, A: 'static, B: 'static>(
458 dst: *mut D,
459 scalar: *const A,
460 src: *const B,
461 inner_len: usize,
462 row_len: usize,
463 src_fast_stride: isize,
464 src_row_stride: isize,
465) -> bool {
466 use std::any::TypeId;
467
468 macro_rules! try_same_type {
469 ($ty:ty) => {
470 if TypeId::of::<D>() == TypeId::of::<$ty>()
471 && TypeId::of::<A>() == TypeId::of::<$ty>()
472 && TypeId::of::<B>() == TypeId::of::<$ty>()
473 {
474 unsafe {
475 mul_transposed_scalar_lhs_2d_typed(
476 dst.cast::<$ty>(),
477 *scalar.cast::<$ty>(),
478 src.cast::<$ty>(),
479 inner_len,
480 row_len,
481 src_fast_stride,
482 src_row_stride,
483 );
484 }
485 return true;
486 }
487 };
488 }
489
490 try_same_type!(f64);
491 try_same_type!(f32);
492 try_same_type!(num_complex::Complex64);
493 try_same_type!(num_complex::Complex32);
494
495 false
496}
497
498pub trait MaybeSimdOps: Copy + Sized {
503 fn try_simd_sum(_src: &[Self]) -> Option<Self> {
504 None
505 }
506 fn try_simd_dot(_a: &[Self], _b: &[Self]) -> Option<Self> {
507 None
508 }
509}
510
511pub(crate) trait MaybeSimdProduct: Copy + Sized {
512 fn try_simd_product(_src: &[Self]) -> Option<Self> {
513 None
514 }
515}
516
517pub(crate) trait MaybeSimdSumSquares: Copy + Sized {
518 fn try_simd_sum_squares(_src: &[Self]) -> Option<Self> {
519 None
520 }
521}
522
523macro_rules! impl_no_simd {
525 ($($t:ty),*) => {
526 $(
527 impl MaybeSimdOps for $t {}
528 impl MaybeSimdProduct for $t {}
529 impl MaybeSimdSumSquares for $t {}
530 )*
531 };
532}
533
534impl_no_simd!(i8, i16, i32, i64, i128, isize, u8, u16, u32, u64, u128, usize);
535
536impl<T: num_traits::Num + Copy + Clone + std::ops::Neg<Output = T>> MaybeSimdOps
537 for num_complex::Complex<T>
538{
539}
540impl<T: num_traits::Num + Copy + Clone + std::ops::Neg<Output = T>> MaybeSimdProduct
541 for num_complex::Complex<T>
542{
543}
544impl<T: num_traits::Num + Copy + Clone + std::ops::Neg<Output = T>> MaybeSimdSumSquares
545 for num_complex::Complex<T>
546{
547}
548
549#[cfg(not(feature = "simd"))]
551impl MaybeSimdOps for f32 {}
552#[cfg(not(feature = "simd"))]
553impl MaybeSimdProduct for f32 {}
554#[cfg(not(feature = "simd"))]
555impl MaybeSimdSumSquares for f32 {}
556
557#[cfg(not(feature = "simd"))]
558impl MaybeSimdOps for f64 {}
559#[cfg(not(feature = "simd"))]
560impl MaybeSimdProduct for f64 {}
561#[cfg(not(feature = "simd"))]
562impl MaybeSimdSumSquares for f64 {}
563
564#[cfg(feature = "simd")]
565mod simd_impls {
566 use super::{MaybeSimdOps, MaybeSimdProduct, MaybeSimdSumSquares};
567 use pulp::{Simd, WithSimd};
568
569 impl MaybeSimdOps for f32 {
570 fn try_simd_sum(src: &[f32]) -> Option<f32> {
571 struct Sum<'a>(&'a [f32]);
572 impl<'a> WithSimd for Sum<'a> {
573 type Output = f32;
574
575 #[inline(always)]
576 fn with_simd<S: Simd>(self, simd: S) -> Self::Output {
577 let (head, tail) = S::as_simd_f32s(self.0);
578
579 let mut acc0 = simd.splat_f32s(0.0);
580 let mut acc1 = simd.splat_f32s(0.0);
581 let mut acc2 = simd.splat_f32s(0.0);
582 let mut acc3 = simd.splat_f32s(0.0);
583
584 let mut i = 0usize;
585 while i + 4 <= head.len() {
586 acc0 = simd.add_f32s(acc0, head[i]);
587 acc1 = simd.add_f32s(acc1, head[i + 1]);
588 acc2 = simd.add_f32s(acc2, head[i + 2]);
589 acc3 = simd.add_f32s(acc3, head[i + 3]);
590 i += 4;
591 }
592 for &v in &head[i..] {
593 acc0 = simd.add_f32s(acc0, v);
594 }
595
596 let acc = simd.add_f32s(simd.add_f32s(acc0, acc1), simd.add_f32s(acc2, acc3));
597 let mut sum = simd.reduce_sum_f32s(acc);
598 for &x in tail {
599 sum += x;
600 }
601 sum
602 }
603 }
604
605 Some(pulp::Arch::new().dispatch(Sum(src)))
606 }
607
608 fn try_simd_dot(a: &[f32], b: &[f32]) -> Option<f32> {
609 try_simd_dot_f32(a, b)
610 }
611 }
612
613 impl MaybeSimdProduct for f32 {
614 fn try_simd_product(src: &[f32]) -> Option<f32> {
615 struct Product<'a>(&'a [f32]);
616 impl<'a> WithSimd for Product<'a> {
617 type Output = f32;
618
619 #[inline(always)]
620 fn with_simd<S: Simd>(self, simd: S) -> Self::Output {
621 let (head, tail) = S::as_simd_f32s(self.0);
622
623 let mut acc0 = simd.splat_f32s(1.0);
624 let mut acc1 = simd.splat_f32s(1.0);
625 let mut acc2 = simd.splat_f32s(1.0);
626 let mut acc3 = simd.splat_f32s(1.0);
627
628 let mut i = 0usize;
629 while i + 4 <= head.len() {
630 acc0 = simd.mul_f32s(acc0, head[i]);
631 acc1 = simd.mul_f32s(acc1, head[i + 1]);
632 acc2 = simd.mul_f32s(acc2, head[i + 2]);
633 acc3 = simd.mul_f32s(acc3, head[i + 3]);
634 i += 4;
635 }
636 for &value in &head[i..] {
637 acc0 = simd.mul_f32s(acc0, value);
638 }
639
640 let acc = simd.mul_f32s(simd.mul_f32s(acc0, acc1), simd.mul_f32s(acc2, acc3));
641 let mut product = simd.reduce_product_f32s(acc);
642 for &value in tail {
643 product *= value;
644 }
645 product
646 }
647 }
648
649 Some(pulp::Arch::new().dispatch(Product(src)))
650 }
651 }
652
653 impl MaybeSimdSumSquares for f32 {
654 fn try_simd_sum_squares(src: &[f32]) -> Option<f32> {
655 struct SumSquares<'a>(&'a [f32]);
656 impl<'a> WithSimd for SumSquares<'a> {
657 type Output = f32;
658
659 #[inline(always)]
660 fn with_simd<S: Simd>(self, simd: S) -> Self::Output {
661 let (head, tail) = S::as_simd_f32s(self.0);
662 let mut acc0 = simd.splat_f32s(0.0);
663 let mut acc1 = simd.splat_f32s(0.0);
664 let mut acc2 = simd.splat_f32s(0.0);
665 let mut acc3 = simd.splat_f32s(0.0);
666
667 let mut i = 0usize;
668 while i + 4 <= head.len() {
669 let square0 = simd.mul_f32s(head[i], head[i]);
670 let square1 = simd.mul_f32s(head[i + 1], head[i + 1]);
671 let square2 = simd.mul_f32s(head[i + 2], head[i + 2]);
672 let square3 = simd.mul_f32s(head[i + 3], head[i + 3]);
673 acc0 = simd.add_f32s(acc0, square0);
674 acc1 = simd.add_f32s(acc1, square1);
675 acc2 = simd.add_f32s(acc2, square2);
676 acc3 = simd.add_f32s(acc3, square3);
677 i += 4;
678 }
679 for &value in &head[i..] {
680 let square = simd.mul_f32s(value, value);
681 acc0 = simd.add_f32s(acc0, square);
682 }
683
684 let acc = simd.add_f32s(simd.add_f32s(acc0, acc1), simd.add_f32s(acc2, acc3));
685 let mut sum = simd.reduce_sum_f32s(acc);
686 for &value in tail {
687 let square = value * value;
688 sum += square;
689 }
690 sum
691 }
692 }
693
694 Some(pulp::Arch::new().dispatch(SumSquares(src)))
695 }
696 }
697
698 fn try_simd_dot_f32(a: &[f32], b: &[f32]) -> Option<f32> {
699 struct Dot<'a> {
700 a: &'a [f32],
701 b: &'a [f32],
702 }
703 impl<'a> WithSimd for Dot<'a> {
704 type Output = f32;
705
706 #[inline(always)]
707 fn with_simd<S: Simd>(self, simd: S) -> Self::Output {
708 debug_assert_eq!(self.a.len(), self.b.len());
709 let (a_head, a_tail) = S::as_simd_f32s(self.a);
710 let (b_head, b_tail) = S::as_simd_f32s(self.b);
711 debug_assert_eq!(a_head.len(), b_head.len());
712 debug_assert_eq!(a_tail.len(), b_tail.len());
713
714 let mut acc0 = simd.splat_f32s(0.0);
715 let mut acc1 = simd.splat_f32s(0.0);
716 let mut acc2 = simd.splat_f32s(0.0);
717 let mut acc3 = simd.splat_f32s(0.0);
718
719 let mut i = 0usize;
720 while i + 4 <= a_head.len() {
721 acc0 = simd.mul_add_f32s(a_head[i], b_head[i], acc0);
722 acc1 = simd.mul_add_f32s(a_head[i + 1], b_head[i + 1], acc1);
723 acc2 = simd.mul_add_f32s(a_head[i + 2], b_head[i + 2], acc2);
724 acc3 = simd.mul_add_f32s(a_head[i + 3], b_head[i + 3], acc3);
725 i += 4;
726 }
727 for j in i..a_head.len() {
728 acc0 = simd.mul_add_f32s(a_head[j], b_head[j], acc0);
729 }
730
731 let acc = simd.add_f32s(simd.add_f32s(acc0, acc1), simd.add_f32s(acc2, acc3));
732 let mut sum = simd.reduce_sum_f32s(acc);
733 for (&x, &y) in a_tail.iter().zip(b_tail.iter()) {
734 sum += x * y;
735 }
736 sum
737 }
738 }
739
740 Some(pulp::Arch::new().dispatch(Dot { a, b }))
741 }
742
743 impl MaybeSimdOps for f64 {
744 fn try_simd_sum(src: &[f64]) -> Option<f64> {
745 struct Sum<'a>(&'a [f64]);
746 impl<'a> WithSimd for Sum<'a> {
747 type Output = f64;
748
749 #[inline(always)]
750 fn with_simd<S: Simd>(self, simd: S) -> Self::Output {
751 let (head, tail) = S::as_simd_f64s(self.0);
752
753 let mut acc0 = simd.splat_f64s(0.0);
754 let mut acc1 = simd.splat_f64s(0.0);
755 let mut acc2 = simd.splat_f64s(0.0);
756 let mut acc3 = simd.splat_f64s(0.0);
757
758 let mut i = 0usize;
759 while i + 4 <= head.len() {
760 acc0 = simd.add_f64s(acc0, head[i]);
761 acc1 = simd.add_f64s(acc1, head[i + 1]);
762 acc2 = simd.add_f64s(acc2, head[i + 2]);
763 acc3 = simd.add_f64s(acc3, head[i + 3]);
764 i += 4;
765 }
766 for &v in &head[i..] {
767 acc0 = simd.add_f64s(acc0, v);
768 }
769
770 let acc = simd.add_f64s(simd.add_f64s(acc0, acc1), simd.add_f64s(acc2, acc3));
771 let mut sum = simd.reduce_sum_f64s(acc);
772 for &x in tail {
773 sum += x;
774 }
775 sum
776 }
777 }
778
779 Some(pulp::Arch::new().dispatch(Sum(src)))
780 }
781
782 fn try_simd_dot(a: &[f64], b: &[f64]) -> Option<f64> {
783 try_simd_dot_f64(a, b)
784 }
785 }
786
787 impl MaybeSimdProduct for f64 {
788 fn try_simd_product(src: &[f64]) -> Option<f64> {
789 struct Product<'a>(&'a [f64]);
790 impl<'a> WithSimd for Product<'a> {
791 type Output = f64;
792
793 #[inline(always)]
794 fn with_simd<S: Simd>(self, simd: S) -> Self::Output {
795 let (head, tail) = S::as_simd_f64s(self.0);
796
797 let mut acc0 = simd.splat_f64s(1.0);
798 let mut acc1 = simd.splat_f64s(1.0);
799 let mut acc2 = simd.splat_f64s(1.0);
800 let mut acc3 = simd.splat_f64s(1.0);
801
802 let mut i = 0usize;
803 while i + 4 <= head.len() {
804 acc0 = simd.mul_f64s(acc0, head[i]);
805 acc1 = simd.mul_f64s(acc1, head[i + 1]);
806 acc2 = simd.mul_f64s(acc2, head[i + 2]);
807 acc3 = simd.mul_f64s(acc3, head[i + 3]);
808 i += 4;
809 }
810 for &value in &head[i..] {
811 acc0 = simd.mul_f64s(acc0, value);
812 }
813
814 let acc = simd.mul_f64s(simd.mul_f64s(acc0, acc1), simd.mul_f64s(acc2, acc3));
815 let mut product = simd.reduce_product_f64s(acc);
816 for &value in tail {
817 product *= value;
818 }
819 product
820 }
821 }
822
823 Some(pulp::Arch::new().dispatch(Product(src)))
824 }
825 }
826
827 impl MaybeSimdSumSquares for f64 {
828 fn try_simd_sum_squares(src: &[f64]) -> Option<f64> {
829 struct SumSquares<'a>(&'a [f64]);
830 impl<'a> WithSimd for SumSquares<'a> {
831 type Output = f64;
832
833 #[inline(always)]
834 fn with_simd<S: Simd>(self, simd: S) -> Self::Output {
835 let (head, tail) = S::as_simd_f64s(self.0);
836 let mut acc0 = simd.splat_f64s(0.0);
837 let mut acc1 = simd.splat_f64s(0.0);
838 let mut acc2 = simd.splat_f64s(0.0);
839 let mut acc3 = simd.splat_f64s(0.0);
840
841 let mut i = 0usize;
842 while i + 4 <= head.len() {
843 let square0 = simd.mul_f64s(head[i], head[i]);
844 let square1 = simd.mul_f64s(head[i + 1], head[i + 1]);
845 let square2 = simd.mul_f64s(head[i + 2], head[i + 2]);
846 let square3 = simd.mul_f64s(head[i + 3], head[i + 3]);
847 acc0 = simd.add_f64s(acc0, square0);
848 acc1 = simd.add_f64s(acc1, square1);
849 acc2 = simd.add_f64s(acc2, square2);
850 acc3 = simd.add_f64s(acc3, square3);
851 i += 4;
852 }
853 for &value in &head[i..] {
854 let square = simd.mul_f64s(value, value);
855 acc0 = simd.add_f64s(acc0, square);
856 }
857
858 let acc = simd.add_f64s(simd.add_f64s(acc0, acc1), simd.add_f64s(acc2, acc3));
859 let mut sum = simd.reduce_sum_f64s(acc);
860 for &value in tail {
861 let square = value * value;
862 sum += square;
863 }
864 sum
865 }
866 }
867
868 Some(pulp::Arch::new().dispatch(SumSquares(src)))
869 }
870 }
871
872 fn try_simd_dot_f64(a: &[f64], b: &[f64]) -> Option<f64> {
873 struct Dot<'a> {
874 a: &'a [f64],
875 b: &'a [f64],
876 }
877 impl<'a> WithSimd for Dot<'a> {
878 type Output = f64;
879
880 #[inline(always)]
881 fn with_simd<S: Simd>(self, simd: S) -> Self::Output {
882 debug_assert_eq!(self.a.len(), self.b.len());
883 let (a_head, a_tail) = S::as_simd_f64s(self.a);
884 let (b_head, b_tail) = S::as_simd_f64s(self.b);
885 debug_assert_eq!(a_head.len(), b_head.len());
886 debug_assert_eq!(a_tail.len(), b_tail.len());
887
888 let mut acc0 = simd.splat_f64s(0.0);
889 let mut acc1 = simd.splat_f64s(0.0);
890 let mut acc2 = simd.splat_f64s(0.0);
891 let mut acc3 = simd.splat_f64s(0.0);
892
893 let mut i = 0usize;
894 while i + 4 <= a_head.len() {
895 acc0 = simd.mul_add_f64s(a_head[i], b_head[i], acc0);
896 acc1 = simd.mul_add_f64s(a_head[i + 1], b_head[i + 1], acc1);
897 acc2 = simd.mul_add_f64s(a_head[i + 2], b_head[i + 2], acc2);
898 acc3 = simd.mul_add_f64s(a_head[i + 3], b_head[i + 3], acc3);
899 i += 4;
900 }
901 for j in i..a_head.len() {
902 acc0 = simd.mul_add_f64s(a_head[j], b_head[j], acc0);
903 }
904
905 let acc = simd.add_f64s(simd.add_f64s(acc0, acc1), simd.add_f64s(acc2, acc3));
906 let mut sum = simd.reduce_sum_f64s(acc);
907 for (&x, &y) in a_tail.iter().zip(b_tail.iter()) {
908 sum += x * y;
909 }
910 sum
911 }
912 }
913
914 Some(pulp::Arch::new().dispatch(Dot { a, b }))
915 }
916}
917
918#[cfg(test)]
919#[path = "simd/tests/tests.rs"]
920mod tests;