Skip to main content

strided_kernel/
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        $lanes:ident,
37        $load:ident,
38        $mask:ident,
39        $store_ptr:ident,
40        $mul:ident,
41        $mask_scale:expr
42    ) => {
43        unsafe fn $mul_into(dst: *mut $ty, len: usize, a: &[$ty], b: &[$ty]) {
44            struct Mul<'a> {
45                dst: *mut $ty,
46                len: usize,
47                a: &'a [$ty],
48                b: &'a [$ty],
49            }
50
51            impl<'a> pulp::WithSimd for Mul<'a> {
52                type Output = ();
53
54                #[inline(always)]
55                fn with_simd<S: pulp::Simd>(self, simd: S) -> Self::Output {
56                    debug_assert_eq!(self.len, self.a.len());
57                    debug_assert_eq!(self.len, self.b.len());
58
59                    let lanes = S::$lanes;
60                    let mut i = 0usize;
61                    while i + lanes <= self.len {
62                        let va = simd.$load(&self.a[i..i + lanes]);
63                        let vb = simd.$load(&self.b[i..i + lanes]);
64                        unsafe {
65                            simd.$store_ptr(
66                                simd.$mask(0, (lanes * $mask_scale) as _),
67                                self.dst.add(i),
68                                simd.$mul(va, vb),
69                            );
70                        }
71                        i += lanes;
72                    }
73                    if i < self.len {
74                        let va = simd.$load(&self.a[i..]);
75                        let vb = simd.$load(&self.b[i..]);
76                        unsafe {
77                            simd.$store_ptr(
78                                simd.$mask(0, ((self.len - i) * $mask_scale) as _),
79                                self.dst.add(i),
80                                simd.$mul(va, vb),
81                            );
82                        }
83                    }
84                }
85            }
86
87            pulp::Arch::new().dispatch(Mul { dst, len, a, b });
88        }
89    };
90}
91
92#[cfg(feature = "simd")]
93impl_simd_mul_ptr!(
94    simd_mul_f32_into,
95    f32,
96    F32_LANES,
97    partial_load_f32s,
98    mask_between_m32s,
99    mask_store_ptr_f32s,
100    mul_f32s,
101    1
102);
103
104#[cfg(feature = "simd")]
105impl_simd_mul_ptr!(
106    simd_mul_f64_into,
107    f64,
108    F64_LANES,
109    partial_load_f64s,
110    mask_between_m64s,
111    mask_store_ptr_f64s,
112    mul_f64s,
113    1
114);
115
116#[cfg(feature = "simd")]
117impl_simd_mul_ptr!(
118    simd_mul_c32_into,
119    num_complex::Complex32,
120    C32_LANES,
121    partial_load_c32s,
122    mask_between_m32s,
123    mask_store_ptr_c32s,
124    mul_e_c32s,
125    2
126);
127
128#[cfg(feature = "simd")]
129impl_simd_mul_ptr!(
130    simd_mul_c64_into,
131    num_complex::Complex64,
132    C64_LANES,
133    partial_load_c64s,
134    mask_between_m64s,
135    mask_store_ptr_c64s,
136    mul_e_c64s,
137    2
138);
139
140#[cfg(feature = "simd")]
141#[inline]
142#[cfg_attr(not(test), allow(dead_code))]
143pub(crate) fn try_mul_contiguous<D: 'static, A: 'static, B: 'static>(
144    dst: &mut [D],
145    a: &[A],
146    b: &[B],
147) -> bool {
148    unsafe { try_mul_contiguous_ptr(dst.as_mut_ptr(), dst.len(), a, b) }
149}
150
151#[cfg(feature = "simd")]
152#[inline]
153pub(crate) unsafe fn try_mul_contiguous_ptr<D: 'static, A: 'static, B: 'static>(
154    dst: *mut D,
155    len: usize,
156    a: &[A],
157    b: &[B],
158) -> bool {
159    use std::any::TypeId;
160
161    macro_rules! try_same_type {
162        ($ty:ty, $mul_into:ident) => {
163            if TypeId::of::<D>() == TypeId::of::<$ty>()
164                && TypeId::of::<A>() == TypeId::of::<$ty>()
165                && TypeId::of::<B>() == TypeId::of::<$ty>()
166            {
167                unsafe { $mul_into(dst.cast::<$ty>(), len, cast_slice(a), cast_slice(b)) };
168                return true;
169            }
170        };
171    }
172
173    try_same_type!(f64, simd_mul_f64_into);
174    try_same_type!(f32, simd_mul_f32_into);
175    try_same_type!(num_complex::Complex64, simd_mul_c64_into);
176    try_same_type!(num_complex::Complex32, simd_mul_c32_into);
177
178    false
179}
180
181#[cfg(not(feature = "simd"))]
182#[inline]
183#[cfg_attr(not(test), allow(dead_code))]
184pub(crate) fn try_mul_contiguous<D: 'static, A: 'static, B: 'static>(
185    _dst: &mut [D],
186    _a: &[A],
187    _b: &[B],
188) -> bool {
189    false
190}
191
192#[cfg(not(feature = "simd"))]
193#[inline]
194pub(crate) unsafe fn try_mul_contiguous_ptr<D: 'static, A: 'static, B: 'static>(
195    _dst: *mut D,
196    _len: usize,
197    _a: &[A],
198    _b: &[B],
199) -> bool {
200    false
201}
202
203#[cfg(feature = "parallel")]
204const TRANSPOSE_TILE: usize = 8;
205
206#[cfg(feature = "parallel")]
207unsafe fn mul_transposed_scalar_rhs_source_contiguous<T>(
208    dst: *mut T,
209    src: *const T,
210    scalar: T,
211    inner_len: usize,
212    row_len: usize,
213    src_fast_stride: isize,
214    src_row_stride: isize,
215) where
216    T: Copy + std::ops::Mul<Output = T>,
217{
218    let mut inner0 = 0usize;
219
220    while inner0 < inner_len {
221        let inner_count = TRANSPOSE_TILE.min(inner_len - inner0);
222        let mut row0 = 0usize;
223
224        while row0 < row_len {
225            let row_count = TRANSPOSE_TILE.min(row_len - row0);
226
227            for inner in 0..inner_count {
228                let src_base = src.offset((inner0 + inner) as isize * src_fast_stride);
229                for row in 0..row_count {
230                    let row_index = row0 + row;
231                    let src_offset = row_index as isize * src_row_stride;
232                    *dst.add(row_index * inner_len + inner0 + inner) =
233                        *src_base.offset(src_offset) * scalar;
234                }
235            }
236
237            row0 += TRANSPOSE_TILE;
238        }
239
240        inner0 += TRANSPOSE_TILE;
241    }
242}
243
244#[cfg(feature = "parallel")]
245unsafe fn mul_transposed_scalar_rhs_dst_contiguous<T>(
246    dst: *mut T,
247    src: *const T,
248    scalar: T,
249    inner_len: usize,
250    row_len: usize,
251    src_fast_stride: isize,
252    src_row_stride: isize,
253) where
254    T: Copy + std::ops::Mul<Output = T>,
255{
256    for row in 0..row_len {
257        let dst_base = dst.add(row * inner_len);
258        let src_base = src.offset(row as isize * src_row_stride);
259        for inner in 0..inner_len {
260            *dst_base.add(inner) = *src_base.offset(inner as isize * src_fast_stride) * scalar;
261        }
262    }
263}
264
265#[cfg(feature = "parallel")]
266unsafe fn mul_transposed_scalar_lhs_source_contiguous<T>(
267    dst: *mut T,
268    scalar: T,
269    src: *const T,
270    inner_len: usize,
271    row_len: usize,
272    src_fast_stride: isize,
273    src_row_stride: isize,
274) where
275    T: Copy + std::ops::Mul<Output = T>,
276{
277    let mut inner0 = 0usize;
278
279    while inner0 < inner_len {
280        let inner_count = TRANSPOSE_TILE.min(inner_len - inner0);
281        let mut row0 = 0usize;
282
283        while row0 < row_len {
284            let row_count = TRANSPOSE_TILE.min(row_len - row0);
285
286            for inner in 0..inner_count {
287                let src_base = src.offset((inner0 + inner) as isize * src_fast_stride);
288                for row in 0..row_count {
289                    let row_index = row0 + row;
290                    let src_offset = row_index as isize * src_row_stride;
291                    *dst.add(row_index * inner_len + inner0 + inner) =
292                        scalar * *src_base.offset(src_offset);
293                }
294            }
295
296            row0 += TRANSPOSE_TILE;
297        }
298
299        inner0 += TRANSPOSE_TILE;
300    }
301}
302
303#[cfg(feature = "parallel")]
304unsafe fn mul_transposed_scalar_lhs_dst_contiguous<T>(
305    dst: *mut T,
306    scalar: T,
307    src: *const T,
308    inner_len: usize,
309    row_len: usize,
310    src_fast_stride: isize,
311    src_row_stride: isize,
312) where
313    T: Copy + std::ops::Mul<Output = T>,
314{
315    for row in 0..row_len {
316        let dst_base = dst.add(row * inner_len);
317        let src_base = src.offset(row as isize * src_row_stride);
318        for inner in 0..inner_len {
319            *dst_base.add(inner) = scalar * *src_base.offset(inner as isize * src_fast_stride);
320        }
321    }
322}
323
324#[cfg(feature = "parallel")]
325#[inline(always)]
326pub(crate) unsafe fn mul_transposed_scalar_rhs_2d_typed<T>(
327    dst: *mut T,
328    src: *const T,
329    scalar: T,
330    inner_len: usize,
331    row_len: usize,
332    src_fast_stride: isize,
333    src_row_stride: isize,
334) where
335    T: Copy + std::ops::Mul<Output = T>,
336{
337    if row_len > inner_len {
338        unsafe {
339            mul_transposed_scalar_rhs_source_contiguous(
340                dst,
341                src,
342                scalar,
343                inner_len,
344                row_len,
345                src_fast_stride,
346                src_row_stride,
347            );
348        }
349    } else {
350        unsafe {
351            mul_transposed_scalar_rhs_dst_contiguous(
352                dst,
353                src,
354                scalar,
355                inner_len,
356                row_len,
357                src_fast_stride,
358                src_row_stride,
359            );
360        }
361    }
362}
363
364#[cfg(feature = "parallel")]
365#[inline(always)]
366pub(crate) unsafe fn mul_transposed_scalar_lhs_2d_typed<T>(
367    dst: *mut T,
368    scalar: T,
369    src: *const T,
370    inner_len: usize,
371    row_len: usize,
372    src_fast_stride: isize,
373    src_row_stride: isize,
374) where
375    T: Copy + std::ops::Mul<Output = T>,
376{
377    if row_len > inner_len {
378        unsafe {
379            mul_transposed_scalar_lhs_source_contiguous(
380                dst,
381                scalar,
382                src,
383                inner_len,
384                row_len,
385                src_fast_stride,
386                src_row_stride,
387            );
388        }
389    } else {
390        unsafe {
391            mul_transposed_scalar_lhs_dst_contiguous(
392                dst,
393                scalar,
394                src,
395                inner_len,
396                row_len,
397                src_fast_stride,
398                src_row_stride,
399            );
400        }
401    }
402}
403
404#[inline]
405#[cfg(feature = "parallel")]
406pub(crate) unsafe fn try_mul_transposed_scalar_rhs_2d<D: 'static, A: 'static, B: 'static>(
407    dst: *mut D,
408    src: *const A,
409    scalar: *const B,
410    inner_len: usize,
411    row_len: usize,
412    src_fast_stride: isize,
413    src_row_stride: isize,
414) -> bool {
415    use std::any::TypeId;
416
417    macro_rules! try_same_type {
418        ($ty:ty) => {
419            if TypeId::of::<D>() == TypeId::of::<$ty>()
420                && TypeId::of::<A>() == TypeId::of::<$ty>()
421                && TypeId::of::<B>() == TypeId::of::<$ty>()
422            {
423                unsafe {
424                    mul_transposed_scalar_rhs_2d_typed(
425                        dst.cast::<$ty>(),
426                        src.cast::<$ty>(),
427                        *scalar.cast::<$ty>(),
428                        inner_len,
429                        row_len,
430                        src_fast_stride,
431                        src_row_stride,
432                    );
433                }
434                return true;
435            }
436        };
437    }
438
439    try_same_type!(f64);
440    try_same_type!(f32);
441    try_same_type!(num_complex::Complex64);
442    try_same_type!(num_complex::Complex32);
443
444    false
445}
446
447#[inline]
448#[cfg(feature = "parallel")]
449pub(crate) unsafe fn try_mul_transposed_scalar_lhs_2d<D: 'static, A: 'static, B: 'static>(
450    dst: *mut D,
451    scalar: *const A,
452    src: *const B,
453    inner_len: usize,
454    row_len: usize,
455    src_fast_stride: isize,
456    src_row_stride: isize,
457) -> bool {
458    use std::any::TypeId;
459
460    macro_rules! try_same_type {
461        ($ty:ty) => {
462            if TypeId::of::<D>() == TypeId::of::<$ty>()
463                && TypeId::of::<A>() == TypeId::of::<$ty>()
464                && TypeId::of::<B>() == TypeId::of::<$ty>()
465            {
466                unsafe {
467                    mul_transposed_scalar_lhs_2d_typed(
468                        dst.cast::<$ty>(),
469                        *scalar.cast::<$ty>(),
470                        src.cast::<$ty>(),
471                        inner_len,
472                        row_len,
473                        src_fast_stride,
474                        src_row_stride,
475                    );
476                }
477                return true;
478            }
479        };
480    }
481
482    try_same_type!(f64);
483    try_same_type!(f32);
484    try_same_type!(num_complex::Complex64);
485    try_same_type!(num_complex::Complex32);
486
487    false
488}
489
490/// Trait for types that may have SIMD-accelerated sum/dot operations.
491///
492/// Default implementations return `None` (no SIMD available).
493/// f32/f64 override these with SIMD kernels when the `simd` feature is enabled.
494pub trait MaybeSimdOps: Copy + Sized {
495    fn try_simd_sum(_src: &[Self]) -> Option<Self> {
496        None
497    }
498    fn try_simd_dot(_a: &[Self], _b: &[Self]) -> Option<Self> {
499        None
500    }
501}
502
503pub(crate) trait MaybeSimdProduct: Copy + Sized {
504    fn try_simd_product(_src: &[Self]) -> Option<Self> {
505        None
506    }
507}
508
509pub(crate) trait MaybeSimdSumSquares: Copy + Sized {
510    fn try_simd_sum_squares(_src: &[Self]) -> Option<Self> {
511        None
512    }
513}
514
515// Default (no-op) impls for integer types and Complex
516macro_rules! impl_no_simd {
517    ($($t:ty),*) => {
518        $(
519            impl MaybeSimdOps for $t {}
520            impl MaybeSimdProduct for $t {}
521            impl MaybeSimdSumSquares for $t {}
522        )*
523    };
524}
525
526impl_no_simd!(i8, i16, i32, i64, i128, isize, u8, u16, u32, u64, u128, usize);
527
528impl<T: num_traits::Num + Copy + Clone + std::ops::Neg<Output = T>> MaybeSimdOps
529    for num_complex::Complex<T>
530{
531}
532impl<T: num_traits::Num + Copy + Clone + std::ops::Neg<Output = T>> MaybeSimdProduct
533    for num_complex::Complex<T>
534{
535}
536impl<T: num_traits::Num + Copy + Clone + std::ops::Neg<Output = T>> MaybeSimdSumSquares
537    for num_complex::Complex<T>
538{
539}
540
541// f32/f64: SIMD-accelerated when feature enabled, no-op otherwise
542#[cfg(not(feature = "simd"))]
543impl MaybeSimdOps for f32 {}
544#[cfg(not(feature = "simd"))]
545impl MaybeSimdProduct for f32 {}
546#[cfg(not(feature = "simd"))]
547impl MaybeSimdSumSquares for f32 {}
548
549#[cfg(not(feature = "simd"))]
550impl MaybeSimdOps for f64 {}
551#[cfg(not(feature = "simd"))]
552impl MaybeSimdProduct for f64 {}
553#[cfg(not(feature = "simd"))]
554impl MaybeSimdSumSquares for f64 {}
555
556#[cfg(feature = "simd")]
557mod simd_impls {
558    use super::{MaybeSimdOps, MaybeSimdProduct, MaybeSimdSumSquares};
559    use pulp::{Simd, WithSimd};
560
561    impl MaybeSimdOps for f32 {
562        fn try_simd_sum(src: &[f32]) -> Option<f32> {
563            struct Sum<'a>(&'a [f32]);
564            impl<'a> WithSimd for Sum<'a> {
565                type Output = f32;
566
567                #[inline(always)]
568                fn with_simd<S: Simd>(self, simd: S) -> Self::Output {
569                    let (head, tail) = S::as_simd_f32s(self.0);
570
571                    let mut acc0 = simd.splat_f32s(0.0);
572                    let mut acc1 = simd.splat_f32s(0.0);
573                    let mut acc2 = simd.splat_f32s(0.0);
574                    let mut acc3 = simd.splat_f32s(0.0);
575
576                    let mut i = 0usize;
577                    while i + 4 <= head.len() {
578                        acc0 = simd.add_f32s(acc0, head[i]);
579                        acc1 = simd.add_f32s(acc1, head[i + 1]);
580                        acc2 = simd.add_f32s(acc2, head[i + 2]);
581                        acc3 = simd.add_f32s(acc3, head[i + 3]);
582                        i += 4;
583                    }
584                    for &v in &head[i..] {
585                        acc0 = simd.add_f32s(acc0, v);
586                    }
587
588                    let acc = simd.add_f32s(simd.add_f32s(acc0, acc1), simd.add_f32s(acc2, acc3));
589                    let mut sum = simd.reduce_sum_f32s(acc);
590                    for &x in tail {
591                        sum += x;
592                    }
593                    sum
594                }
595            }
596
597            Some(pulp::Arch::new().dispatch(Sum(src)))
598        }
599
600        fn try_simd_dot(a: &[f32], b: &[f32]) -> Option<f32> {
601            try_simd_dot_f32(a, b)
602        }
603    }
604
605    impl MaybeSimdProduct for f32 {
606        fn try_simd_product(src: &[f32]) -> Option<f32> {
607            struct Product<'a>(&'a [f32]);
608            impl<'a> WithSimd for Product<'a> {
609                type Output = f32;
610
611                #[inline(always)]
612                fn with_simd<S: Simd>(self, simd: S) -> Self::Output {
613                    let (head, tail) = S::as_simd_f32s(self.0);
614
615                    let mut acc0 = simd.splat_f32s(1.0);
616                    let mut acc1 = simd.splat_f32s(1.0);
617                    let mut acc2 = simd.splat_f32s(1.0);
618                    let mut acc3 = simd.splat_f32s(1.0);
619
620                    let mut i = 0usize;
621                    while i + 4 <= head.len() {
622                        acc0 = simd.mul_f32s(acc0, head[i]);
623                        acc1 = simd.mul_f32s(acc1, head[i + 1]);
624                        acc2 = simd.mul_f32s(acc2, head[i + 2]);
625                        acc3 = simd.mul_f32s(acc3, head[i + 3]);
626                        i += 4;
627                    }
628                    for &value in &head[i..] {
629                        acc0 = simd.mul_f32s(acc0, value);
630                    }
631
632                    let acc = simd.mul_f32s(simd.mul_f32s(acc0, acc1), simd.mul_f32s(acc2, acc3));
633                    let mut product = simd.reduce_product_f32s(acc);
634                    for &value in tail {
635                        product *= value;
636                    }
637                    product
638                }
639            }
640
641            Some(pulp::Arch::new().dispatch(Product(src)))
642        }
643    }
644
645    impl MaybeSimdSumSquares for f32 {
646        fn try_simd_sum_squares(src: &[f32]) -> Option<f32> {
647            struct SumSquares<'a>(&'a [f32]);
648            impl<'a> WithSimd for SumSquares<'a> {
649                type Output = f32;
650
651                #[inline(always)]
652                fn with_simd<S: Simd>(self, simd: S) -> Self::Output {
653                    let (head, tail) = S::as_simd_f32s(self.0);
654                    let mut acc0 = simd.splat_f32s(0.0);
655                    let mut acc1 = simd.splat_f32s(0.0);
656                    let mut acc2 = simd.splat_f32s(0.0);
657                    let mut acc3 = simd.splat_f32s(0.0);
658
659                    let mut i = 0usize;
660                    while i + 4 <= head.len() {
661                        let square0 = simd.mul_f32s(head[i], head[i]);
662                        let square1 = simd.mul_f32s(head[i + 1], head[i + 1]);
663                        let square2 = simd.mul_f32s(head[i + 2], head[i + 2]);
664                        let square3 = simd.mul_f32s(head[i + 3], head[i + 3]);
665                        acc0 = simd.add_f32s(acc0, square0);
666                        acc1 = simd.add_f32s(acc1, square1);
667                        acc2 = simd.add_f32s(acc2, square2);
668                        acc3 = simd.add_f32s(acc3, square3);
669                        i += 4;
670                    }
671                    for &value in &head[i..] {
672                        let square = simd.mul_f32s(value, value);
673                        acc0 = simd.add_f32s(acc0, square);
674                    }
675
676                    let acc = simd.add_f32s(simd.add_f32s(acc0, acc1), simd.add_f32s(acc2, acc3));
677                    let mut sum = simd.reduce_sum_f32s(acc);
678                    for &value in tail {
679                        let square = value * value;
680                        sum += square;
681                    }
682                    sum
683                }
684            }
685
686            Some(pulp::Arch::new().dispatch(SumSquares(src)))
687        }
688    }
689
690    fn try_simd_dot_f32(a: &[f32], b: &[f32]) -> Option<f32> {
691        struct Dot<'a> {
692            a: &'a [f32],
693            b: &'a [f32],
694        }
695        impl<'a> WithSimd for Dot<'a> {
696            type Output = f32;
697
698            #[inline(always)]
699            fn with_simd<S: Simd>(self, simd: S) -> Self::Output {
700                debug_assert_eq!(self.a.len(), self.b.len());
701                let (a_head, a_tail) = S::as_simd_f32s(self.a);
702                let (b_head, b_tail) = S::as_simd_f32s(self.b);
703                debug_assert_eq!(a_head.len(), b_head.len());
704                debug_assert_eq!(a_tail.len(), b_tail.len());
705
706                let mut acc0 = simd.splat_f32s(0.0);
707                let mut acc1 = simd.splat_f32s(0.0);
708                let mut acc2 = simd.splat_f32s(0.0);
709                let mut acc3 = simd.splat_f32s(0.0);
710
711                let mut i = 0usize;
712                while i + 4 <= a_head.len() {
713                    acc0 = simd.mul_add_f32s(a_head[i], b_head[i], acc0);
714                    acc1 = simd.mul_add_f32s(a_head[i + 1], b_head[i + 1], acc1);
715                    acc2 = simd.mul_add_f32s(a_head[i + 2], b_head[i + 2], acc2);
716                    acc3 = simd.mul_add_f32s(a_head[i + 3], b_head[i + 3], acc3);
717                    i += 4;
718                }
719                for j in i..a_head.len() {
720                    acc0 = simd.mul_add_f32s(a_head[j], b_head[j], acc0);
721                }
722
723                let acc = simd.add_f32s(simd.add_f32s(acc0, acc1), simd.add_f32s(acc2, acc3));
724                let mut sum = simd.reduce_sum_f32s(acc);
725                for (&x, &y) in a_tail.iter().zip(b_tail.iter()) {
726                    sum += x * y;
727                }
728                sum
729            }
730        }
731
732        Some(pulp::Arch::new().dispatch(Dot { a, b }))
733    }
734
735    impl MaybeSimdOps for f64 {
736        fn try_simd_sum(src: &[f64]) -> Option<f64> {
737            struct Sum<'a>(&'a [f64]);
738            impl<'a> WithSimd for Sum<'a> {
739                type Output = f64;
740
741                #[inline(always)]
742                fn with_simd<S: Simd>(self, simd: S) -> Self::Output {
743                    let (head, tail) = S::as_simd_f64s(self.0);
744
745                    let mut acc0 = simd.splat_f64s(0.0);
746                    let mut acc1 = simd.splat_f64s(0.0);
747                    let mut acc2 = simd.splat_f64s(0.0);
748                    let mut acc3 = simd.splat_f64s(0.0);
749
750                    let mut i = 0usize;
751                    while i + 4 <= head.len() {
752                        acc0 = simd.add_f64s(acc0, head[i]);
753                        acc1 = simd.add_f64s(acc1, head[i + 1]);
754                        acc2 = simd.add_f64s(acc2, head[i + 2]);
755                        acc3 = simd.add_f64s(acc3, head[i + 3]);
756                        i += 4;
757                    }
758                    for &v in &head[i..] {
759                        acc0 = simd.add_f64s(acc0, v);
760                    }
761
762                    let acc = simd.add_f64s(simd.add_f64s(acc0, acc1), simd.add_f64s(acc2, acc3));
763                    let mut sum = simd.reduce_sum_f64s(acc);
764                    for &x in tail {
765                        sum += x;
766                    }
767                    sum
768                }
769            }
770
771            Some(pulp::Arch::new().dispatch(Sum(src)))
772        }
773
774        fn try_simd_dot(a: &[f64], b: &[f64]) -> Option<f64> {
775            try_simd_dot_f64(a, b)
776        }
777    }
778
779    impl MaybeSimdProduct for f64 {
780        fn try_simd_product(src: &[f64]) -> Option<f64> {
781            struct Product<'a>(&'a [f64]);
782            impl<'a> WithSimd for Product<'a> {
783                type Output = f64;
784
785                #[inline(always)]
786                fn with_simd<S: Simd>(self, simd: S) -> Self::Output {
787                    let (head, tail) = S::as_simd_f64s(self.0);
788
789                    let mut acc0 = simd.splat_f64s(1.0);
790                    let mut acc1 = simd.splat_f64s(1.0);
791                    let mut acc2 = simd.splat_f64s(1.0);
792                    let mut acc3 = simd.splat_f64s(1.0);
793
794                    let mut i = 0usize;
795                    while i + 4 <= head.len() {
796                        acc0 = simd.mul_f64s(acc0, head[i]);
797                        acc1 = simd.mul_f64s(acc1, head[i + 1]);
798                        acc2 = simd.mul_f64s(acc2, head[i + 2]);
799                        acc3 = simd.mul_f64s(acc3, head[i + 3]);
800                        i += 4;
801                    }
802                    for &value in &head[i..] {
803                        acc0 = simd.mul_f64s(acc0, value);
804                    }
805
806                    let acc = simd.mul_f64s(simd.mul_f64s(acc0, acc1), simd.mul_f64s(acc2, acc3));
807                    let mut product = simd.reduce_product_f64s(acc);
808                    for &value in tail {
809                        product *= value;
810                    }
811                    product
812                }
813            }
814
815            Some(pulp::Arch::new().dispatch(Product(src)))
816        }
817    }
818
819    impl MaybeSimdSumSquares for f64 {
820        fn try_simd_sum_squares(src: &[f64]) -> Option<f64> {
821            struct SumSquares<'a>(&'a [f64]);
822            impl<'a> WithSimd for SumSquares<'a> {
823                type Output = f64;
824
825                #[inline(always)]
826                fn with_simd<S: Simd>(self, simd: S) -> Self::Output {
827                    let (head, tail) = S::as_simd_f64s(self.0);
828                    let mut acc0 = simd.splat_f64s(0.0);
829                    let mut acc1 = simd.splat_f64s(0.0);
830                    let mut acc2 = simd.splat_f64s(0.0);
831                    let mut acc3 = simd.splat_f64s(0.0);
832
833                    let mut i = 0usize;
834                    while i + 4 <= head.len() {
835                        let square0 = simd.mul_f64s(head[i], head[i]);
836                        let square1 = simd.mul_f64s(head[i + 1], head[i + 1]);
837                        let square2 = simd.mul_f64s(head[i + 2], head[i + 2]);
838                        let square3 = simd.mul_f64s(head[i + 3], head[i + 3]);
839                        acc0 = simd.add_f64s(acc0, square0);
840                        acc1 = simd.add_f64s(acc1, square1);
841                        acc2 = simd.add_f64s(acc2, square2);
842                        acc3 = simd.add_f64s(acc3, square3);
843                        i += 4;
844                    }
845                    for &value in &head[i..] {
846                        let square = simd.mul_f64s(value, value);
847                        acc0 = simd.add_f64s(acc0, square);
848                    }
849
850                    let acc = simd.add_f64s(simd.add_f64s(acc0, acc1), simd.add_f64s(acc2, acc3));
851                    let mut sum = simd.reduce_sum_f64s(acc);
852                    for &value in tail {
853                        let square = value * value;
854                        sum += square;
855                    }
856                    sum
857                }
858            }
859
860            Some(pulp::Arch::new().dispatch(SumSquares(src)))
861        }
862    }
863
864    fn try_simd_dot_f64(a: &[f64], b: &[f64]) -> Option<f64> {
865        struct Dot<'a> {
866            a: &'a [f64],
867            b: &'a [f64],
868        }
869        impl<'a> WithSimd for Dot<'a> {
870            type Output = f64;
871
872            #[inline(always)]
873            fn with_simd<S: Simd>(self, simd: S) -> Self::Output {
874                debug_assert_eq!(self.a.len(), self.b.len());
875                let (a_head, a_tail) = S::as_simd_f64s(self.a);
876                let (b_head, b_tail) = S::as_simd_f64s(self.b);
877                debug_assert_eq!(a_head.len(), b_head.len());
878                debug_assert_eq!(a_tail.len(), b_tail.len());
879
880                let mut acc0 = simd.splat_f64s(0.0);
881                let mut acc1 = simd.splat_f64s(0.0);
882                let mut acc2 = simd.splat_f64s(0.0);
883                let mut acc3 = simd.splat_f64s(0.0);
884
885                let mut i = 0usize;
886                while i + 4 <= a_head.len() {
887                    acc0 = simd.mul_add_f64s(a_head[i], b_head[i], acc0);
888                    acc1 = simd.mul_add_f64s(a_head[i + 1], b_head[i + 1], acc1);
889                    acc2 = simd.mul_add_f64s(a_head[i + 2], b_head[i + 2], acc2);
890                    acc3 = simd.mul_add_f64s(a_head[i + 3], b_head[i + 3], acc3);
891                    i += 4;
892                }
893                for j in i..a_head.len() {
894                    acc0 = simd.mul_add_f64s(a_head[j], b_head[j], acc0);
895                }
896
897                let acc = simd.add_f64s(simd.add_f64s(acc0, acc1), simd.add_f64s(acc2, acc3));
898                let mut sum = simd.reduce_sum_f64s(acc);
899                for (&x, &y) in a_tail.iter().zip(b_tail.iter()) {
900                    sum += x * y;
901                }
902                sum
903            }
904        }
905
906        Some(pulp::Arch::new().dispatch(Dot { a, b }))
907    }
908}
909
910#[cfg(test)]
911mod tests {
912    #[cfg(feature = "simd")]
913    #[test]
914    fn sum_squares_simd_covers_unrolled_body_and_tail() {
915        const LEN: usize = 285;
916        let f32_values = vec![1.0_f32; LEN];
917        let f64_values = vec![1.0_f64; LEN];
918
919        assert_eq!(
920            <f32 as super::MaybeSimdSumSquares>::try_simd_sum_squares(&f32_values),
921            Some(LEN as f32)
922        );
923        assert_eq!(
924            <f64 as super::MaybeSimdSumSquares>::try_simd_sum_squares(&f64_values),
925            Some(LEN as f64)
926        );
927    }
928
929    #[cfg(feature = "simd")]
930    #[test]
931    fn test_try_mul_contiguous_complex64() {
932        let a = vec![
933            num_complex::Complex64::new(1.0, 2.0),
934            num_complex::Complex64::new(-3.0, 4.0),
935            num_complex::Complex64::new(0.5, -0.25),
936        ];
937        let b = vec![
938            num_complex::Complex64::new(5.0, -1.0),
939            num_complex::Complex64::new(2.0, 0.25),
940            num_complex::Complex64::new(-4.0, 3.0),
941        ];
942        let mut dst = vec![num_complex::Complex64::new(0.0, 0.0); a.len()];
943
944        assert!(super::try_mul_contiguous(&mut dst, &a, &b));
945        for i in 0..a.len() {
946            assert_eq!(dst[i], a[i] * b[i]);
947        }
948    }
949
950    #[cfg(feature = "simd")]
951    #[test]
952    fn test_try_mul_contiguous_complex32() {
953        let a = vec![
954            num_complex::Complex32::new(1.0, 2.0),
955            num_complex::Complex32::new(-3.0, 4.0),
956            num_complex::Complex32::new(0.5, -0.25),
957        ];
958        let b = vec![
959            num_complex::Complex32::new(5.0, -1.0),
960            num_complex::Complex32::new(2.0, 0.25),
961            num_complex::Complex32::new(-4.0, 3.0),
962        ];
963        let mut dst = vec![num_complex::Complex32::new(0.0, 0.0); a.len()];
964
965        assert!(super::try_mul_contiguous(&mut dst, &a, &b));
966        for i in 0..a.len() {
967            assert_eq!(dst[i], a[i] * b[i]);
968        }
969    }
970
971    #[cfg(feature = "parallel")]
972    #[test]
973    fn test_transposed_scalar_rhs_2d_f64_source_contiguous() {
974        let inner_len = 5usize;
975        let row_len = 7usize;
976        let src: Vec<f64> = (0..inner_len * row_len).map(|i| i as f64 + 0.25).collect();
977        let scalar = 2.0f64;
978        let mut dst = vec![0.0f64; inner_len * row_len];
979
980        let used = unsafe {
981            super::try_mul_transposed_scalar_rhs_2d::<f64, f64, f64>(
982                dst.as_mut_ptr(),
983                src.as_ptr(),
984                &scalar,
985                inner_len,
986                row_len,
987                row_len as isize,
988                1,
989            )
990        };
991
992        assert!(used);
993        for row in 0..row_len {
994            for inner in 0..inner_len {
995                assert_eq!(
996                    dst[row * inner_len + inner],
997                    src[inner * row_len + row] * scalar
998                );
999            }
1000        }
1001    }
1002
1003    #[cfg(feature = "parallel")]
1004    #[test]
1005    fn test_transposed_scalar_lhs_2d_f32_source_contiguous() {
1006        let inner_len = 5usize;
1007        let row_len = 7usize;
1008        let src: Vec<f32> = (0..inner_len * row_len)
1009            .map(|i| i as f32 * 0.5 + 1.0)
1010            .collect();
1011        let scalar = 3.0f32;
1012        let mut dst = vec![0.0f32; inner_len * row_len];
1013
1014        let used = unsafe {
1015            super::try_mul_transposed_scalar_lhs_2d::<f32, f32, f32>(
1016                dst.as_mut_ptr(),
1017                &scalar,
1018                src.as_ptr(),
1019                inner_len,
1020                row_len,
1021                row_len as isize,
1022                1,
1023            )
1024        };
1025
1026        assert!(used);
1027        for row in 0..row_len {
1028            for inner in 0..inner_len {
1029                assert_eq!(
1030                    dst[row * inner_len + inner],
1031                    scalar * src[inner * row_len + row]
1032                );
1033            }
1034        }
1035    }
1036
1037    #[cfg(feature = "parallel")]
1038    #[test]
1039    fn test_transposed_scalar_rhs_2d_handles_short_rows() {
1040        let inner_len = 8usize;
1041        let row_len = 3usize;
1042        let src: Vec<f64> = (0..inner_len * row_len).map(|i| i as f64 + 1.0).collect();
1043        let scalar = 2.0f64;
1044        let mut dst = vec![0.0f64; inner_len * row_len];
1045
1046        let used = unsafe {
1047            super::try_mul_transposed_scalar_rhs_2d::<f64, f64, f64>(
1048                dst.as_mut_ptr(),
1049                src.as_ptr(),
1050                &scalar,
1051                inner_len,
1052                row_len,
1053                row_len as isize,
1054                1,
1055            )
1056        };
1057
1058        assert!(used);
1059        for row in 0..row_len {
1060            for inner in 0..inner_len {
1061                assert_eq!(
1062                    dst[row * inner_len + inner],
1063                    src[inner * row_len + row] * scalar
1064                );
1065            }
1066        }
1067    }
1068
1069    #[cfg(feature = "parallel")]
1070    #[test]
1071    fn test_transposed_scalar_2d_handles_small_square_tiles() {
1072        let inner_len = 4usize;
1073        let row_len = 4usize;
1074        let src: Vec<f64> = (0..inner_len * row_len).map(|i| i as f64 + 1.0).collect();
1075        let scalar = 2.0f64;
1076        let mut rhs_dst = vec![0.0f64; inner_len * row_len];
1077        let mut lhs_dst = vec![0.0f64; inner_len * row_len];
1078
1079        let rhs_used = unsafe {
1080            super::try_mul_transposed_scalar_rhs_2d::<f64, f64, f64>(
1081                rhs_dst.as_mut_ptr(),
1082                src.as_ptr(),
1083                &scalar,
1084                inner_len,
1085                row_len,
1086                row_len as isize,
1087                1,
1088            )
1089        };
1090        let lhs_used = unsafe {
1091            super::try_mul_transposed_scalar_lhs_2d::<f64, f64, f64>(
1092                lhs_dst.as_mut_ptr(),
1093                &scalar,
1094                src.as_ptr(),
1095                inner_len,
1096                row_len,
1097                row_len as isize,
1098                1,
1099            )
1100        };
1101
1102        assert!(rhs_used);
1103        assert!(lhs_used);
1104        for row in 0..row_len {
1105            for inner in 0..inner_len {
1106                assert_eq!(
1107                    rhs_dst[row * inner_len + inner],
1108                    src[inner * row_len + row] * scalar
1109                );
1110                assert_eq!(
1111                    lhs_dst[row * inner_len + inner],
1112                    scalar * src[inner * row_len + row]
1113                );
1114            }
1115        }
1116    }
1117
1118    #[cfg(feature = "parallel")]
1119    #[test]
1120    fn test_transposed_scalar_rhs_2d_complex64_source_contiguous() {
1121        let inner_len = 5usize;
1122        let row_len = 7usize;
1123        let src: Vec<num_complex::Complex64> = (0..inner_len * row_len)
1124            .map(|i| num_complex::Complex64::new(i as f64 + 0.25, i as f64 * -0.5))
1125            .collect();
1126        let scalar = num_complex::Complex64::new(2.0, -0.25);
1127        let mut dst = vec![num_complex::Complex64::new(0.0, 0.0); inner_len * row_len];
1128
1129        let used = unsafe {
1130            super::try_mul_transposed_scalar_rhs_2d::<
1131                num_complex::Complex64,
1132                num_complex::Complex64,
1133                num_complex::Complex64,
1134            >(
1135                dst.as_mut_ptr(),
1136                src.as_ptr(),
1137                &scalar,
1138                inner_len,
1139                row_len,
1140                row_len as isize,
1141                1,
1142            )
1143        };
1144
1145        assert!(used);
1146        for row in 0..row_len {
1147            for inner in 0..inner_len {
1148                assert_eq!(
1149                    dst[row * inner_len + inner],
1150                    src[inner * row_len + row] * scalar
1151                );
1152            }
1153        }
1154    }
1155
1156    #[cfg(feature = "parallel")]
1157    #[test]
1158    fn test_transposed_scalar_lhs_2d_complex32_source_contiguous() {
1159        let inner_len = 5usize;
1160        let row_len = 7usize;
1161        let src: Vec<num_complex::Complex32> = (0..inner_len * row_len)
1162            .map(|i| num_complex::Complex32::new(i as f32 * 0.5 + 1.0, i as f32 * 0.25))
1163            .collect();
1164        let scalar = num_complex::Complex32::new(3.0, -0.5);
1165        let mut dst = vec![num_complex::Complex32::new(0.0, 0.0); inner_len * row_len];
1166
1167        let used = unsafe {
1168            super::try_mul_transposed_scalar_lhs_2d::<
1169                num_complex::Complex32,
1170                num_complex::Complex32,
1171                num_complex::Complex32,
1172            >(
1173                dst.as_mut_ptr(),
1174                &scalar,
1175                src.as_ptr(),
1176                inner_len,
1177                row_len,
1178                row_len as isize,
1179                1,
1180            )
1181        };
1182
1183        assert!(used);
1184        for row in 0..row_len {
1185            for inner in 0..inner_len {
1186                assert_eq!(
1187                    dst[row * inner_len + inner],
1188                    scalar * src[inner * row_len + row]
1189                );
1190            }
1191        }
1192    }
1193}