Skip to main content

strided_basic/
simd.rs

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    // Avoid runtime-dispatch overhead for tiny loops (especially common for small-array cases).
16    // This is a heuristic; correctness does not depend on it.
17    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                    // SAFETY: the caller supplies len writable, nonaliasing
61                    // elements. MaybeUninit does not assert initialized output.
62                    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                    // Full vectors use ordinary stores. Pulp's partial-access
72                    // machinery is needed only for the final incomplete vector.
73                    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                        // SAFETY: the mask writes exactly the remaining lanes.
80                        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
498/// Trait for types that may have SIMD-accelerated sum/dot operations.
499///
500/// Default implementations return `None` (no SIMD available).
501/// f32/f64 override these with SIMD kernels when the `simd` feature is enabled.
502pub 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
523// Default (no-op) impls for integer types and Complex
524macro_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// f32/f64: SIMD-accelerated when feature enabled, no-op otherwise
550#[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;