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 $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
490pub 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
515macro_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#[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}