1use super::arch;
4use core::iter::Sum;
5use core::ops::{Add, Div, Mul, Sub};
6
7pub(crate) mod sealed {
8 pub trait Sealed {}
9}
10
11#[allow(private_bounds)]
17pub trait SimdScalar:
18 sealed::Sealed + Copy + Send + Sync + Add<Output = Self> + Mul<Output = Self> + Sum<Self> + 'static
19{
20 const ZERO: Self;
22
23 #[doc(hidden)]
24 #[inline]
25 fn native_vector_available() -> bool {
26 false
27 }
28
29 #[doc(hidden)]
30 #[inline]
31 fn uses_native_vector_path(len: usize) -> bool {
32 let _ = len;
33 false
34 }
35
36 #[doc(hidden)]
37 #[inline]
38 fn matrix_vector_path_available<const N: usize>() -> bool {
39 let _ = N;
40 false
41 }
42
43 #[doc(hidden)]
44 #[inline]
45 fn add_slices(left: &[Self], right: &[Self], result: &mut [Self]) {
46 scalar_add(left, right, result);
47 }
48
49 #[doc(hidden)]
50 #[inline]
51 fn mul_slices(left: &[Self], right: &[Self], result: &mut [Self]) {
52 scalar_mul(left, right, result);
53 }
54
55 #[doc(hidden)]
56 #[inline]
57 fn dot_slice(left: &[Self], right: &[Self]) -> Self {
58 scalar_dot(left, right)
59 }
60
61 #[doc(hidden)]
62 #[inline]
63 fn sum_slice(data: &[Self]) -> Self {
64 data.iter().copied().sum()
65 }
66
67 #[doc(hidden)]
68 #[inline]
69 fn matrix_mul_square<const N: usize>(left: &[Self], right: &[Self], result: &mut [Self]) {
70 scalar_matrix_mul_square::<Self, N>(left, right, result);
71 }
72}
73
74#[allow(private_bounds)]
76pub trait SimdReal: SimdScalar + Sub<Output = Self> + Div<Output = Self> {
77 fn from_len(len: usize) -> Self;
79
80 #[doc(hidden)]
81 #[inline]
82 fn mean_slice(data: &[Self]) -> Self {
83 Self::sum_slice(data) / Self::from_len(data.len())
84 }
85
86 #[doc(hidden)]
87 #[inline]
88 fn variance_slice(data: &[Self]) -> Self {
89 let mean = Self::mean_slice(data);
90 data.iter()
91 .copied()
92 .map(|value| {
93 let diff = value - mean;
94 diff * diff
95 })
96 .sum::<Self>()
97 / Self::from_len(data.len())
98 }
99}
100
101#[inline]
102fn scalar_add<T: SimdScalar>(left: &[T], right: &[T], result: &mut [T]) {
103 for ((left, right), output) in left.iter().zip(right.iter()).zip(result.iter_mut()) {
104 *output = *left + *right;
105 }
106}
107
108#[inline]
109fn scalar_mul<T: SimdScalar>(left: &[T], right: &[T], result: &mut [T]) {
110 for ((left, right), output) in left.iter().zip(right.iter()).zip(result.iter_mut()) {
111 *output = *left * *right;
112 }
113}
114
115#[inline]
116fn scalar_dot<T: SimdScalar>(left: &[T], right: &[T]) -> T {
117 left.iter()
118 .copied()
119 .zip(right.iter().copied())
120 .fold(T::ZERO, |acc, (left, right)| acc + left * right)
121}
122
123#[inline]
124fn scalar_matrix_mul_square<T: SimdScalar, const N: usize>(
125 left: &[T],
126 right: &[T],
127 result: &mut [T],
128) {
129 assert!(N != 0, "matrix dimension must be non-zero");
130 let expected = N.checked_mul(N).expect("matrix dimension overflow");
131 assert_eq!(left.len(), expected, "left matrix size must equal N * N");
132 assert_eq!(right.len(), expected, "right matrix size must equal N * N");
133 assert_eq!(
134 result.len(),
135 expected,
136 "result matrix size must equal N * N"
137 );
138
139 for row in 0..N {
140 for col in 0..N {
141 let mut acc = T::ZERO;
142 for index in 0..N {
143 acc = acc + left[row * N + index] * right[index * N + col];
144 }
145 result[row * N + col] = acc;
146 }
147 }
148}
149
150impl sealed::Sealed for f32 {}
151impl SimdScalar for f32 {
152 const ZERO: Self = 0.0;
153
154 #[inline]
155 fn native_vector_available() -> bool {
156 native_vector_available()
157 }
158
159 #[inline]
160 fn uses_native_vector_path(len: usize) -> bool {
161 uses_native_vector_path(len)
162 }
163
164 #[inline]
165 fn matrix_vector_path_available<const N: usize>() -> bool {
166 N == 4 && native_vector_available()
167 }
168
169 #[inline]
170 fn add_slices(left: &[Self], right: &[Self], result: &mut [Self]) {
171 let len = left.len();
172 #[cfg(any(target_arch = "x86_64", target_arch = "aarch64"))]
173 {
174 if let Some(chunk_len) = native_vector_chunk_len(len) {
175 unsafe {
176 arch::add(
177 &left[..chunk_len],
178 &right[..chunk_len],
179 &mut result[..chunk_len],
180 );
181 }
182 if chunk_len < len {
183 scalar_add(
184 &left[chunk_len..],
185 &right[chunk_len..],
186 &mut result[chunk_len..],
187 );
188 }
189 return;
190 }
191 }
192 scalar_add(left, right, result);
193 }
194
195 #[inline]
196 fn mul_slices(left: &[Self], right: &[Self], result: &mut [Self]) {
197 let len = left.len();
198 #[cfg(any(target_arch = "x86_64", target_arch = "aarch64"))]
199 {
200 if let Some(chunk_len) = native_vector_chunk_len(len) {
201 unsafe {
202 arch::mul(
203 &left[..chunk_len],
204 &right[..chunk_len],
205 &mut result[..chunk_len],
206 );
207 }
208 if chunk_len < len {
209 scalar_mul(
210 &left[chunk_len..],
211 &right[chunk_len..],
212 &mut result[chunk_len..],
213 );
214 }
215 return;
216 }
217 }
218 scalar_mul(left, right, result);
219 }
220
221 #[inline]
222 fn dot_slice(left: &[Self], right: &[Self]) -> Self {
223 let len = left.len();
224 #[cfg(any(target_arch = "x86_64", target_arch = "aarch64"))]
225 {
226 if let Some(chunk_len) = native_vector_chunk_len(len) {
227 let mut sum = unsafe { arch::dot(&left[..chunk_len], &right[..chunk_len]) };
228 if chunk_len < len {
229 sum += scalar_dot(&left[chunk_len..], &right[chunk_len..]);
230 }
231 return sum;
232 }
233 }
234 scalar_dot(left, right)
235 }
236
237 #[inline]
238 fn sum_slice(data: &[Self]) -> Self {
239 let len = data.len();
240 #[cfg(any(target_arch = "x86_64", target_arch = "aarch64"))]
241 {
242 if let Some(chunk_len) = native_vector_chunk_len(len) {
243 let mut total = unsafe { arch::sum(&data[..chunk_len]) };
244 if chunk_len < len {
245 total += data[chunk_len..].iter().copied().sum::<Self>();
246 }
247 return total;
248 }
249 }
250 data.iter().copied().sum()
251 }
252
253 #[inline]
254 fn matrix_mul_square<const N: usize>(left: &[Self], right: &[Self], result: &mut [Self]) {
255 if N == 4 && native_vector_available() {
256 scalar_matrix_shape::<N>(left, right, result);
257 unsafe {
258 arch::matrix_mul_square(left, right, result);
259 }
260 } else {
261 scalar_matrix_mul_square::<Self, N>(left, right, result);
262 }
263 }
264}
265
266impl SimdReal for f32 {
267 #[inline]
268 fn from_len(len: usize) -> Self {
269 len as Self
270 }
271
272 #[inline]
273 fn variance_slice(data: &[Self]) -> Self {
274 let len = data.len();
275 #[cfg(any(target_arch = "x86_64", target_arch = "aarch64"))]
276 {
277 if let Some(chunk_len) = native_vector_chunk_len(len) {
278 let mean = Self::mean_slice(data);
279 let mut total = unsafe { arch::squared_diff_sum(&data[..chunk_len], mean) };
280 if chunk_len < len {
281 total += data[chunk_len..]
282 .iter()
283 .copied()
284 .map(|value| {
285 let diff = value - mean;
286 diff * diff
287 })
288 .sum::<Self>();
289 }
290 return total / Self::from_len(len);
291 }
292 }
293
294 let mean = Self::mean_slice(data);
295 data.iter()
296 .copied()
297 .map(|value| {
298 let diff = value - mean;
299 diff * diff
300 })
301 .sum::<Self>()
302 / Self::from_len(len)
303 }
304}
305
306impl sealed::Sealed for f64 {}
307impl SimdScalar for f64 {
308 const ZERO: Self = 0.0;
309
310 #[inline]
311 fn native_vector_available() -> bool {
312 native_wide_vector_available()
313 }
314
315 #[inline]
316 fn uses_native_vector_path(len: usize) -> bool {
317 uses_native_wide_vector_path(len)
318 }
319
320 #[inline]
321 fn add_slices(left: &[Self], right: &[Self], result: &mut [Self]) {
322 let len = left.len();
323 #[cfg(target_arch = "x86_64")]
324 {
325 if let Some(chunk_len) = native_vector_chunk_len(len) {
326 unsafe {
327 arch::add_wide(
328 &left[..chunk_len],
329 &right[..chunk_len],
330 &mut result[..chunk_len],
331 );
332 }
333 if chunk_len < len {
334 scalar_add(
335 &left[chunk_len..],
336 &right[chunk_len..],
337 &mut result[chunk_len..],
338 );
339 }
340 return;
341 }
342 }
343 scalar_add(left, right, result);
344 }
345
346 #[inline]
347 fn mul_slices(left: &[Self], right: &[Self], result: &mut [Self]) {
348 let len = left.len();
349 #[cfg(target_arch = "x86_64")]
350 {
351 if let Some(chunk_len) = native_vector_chunk_len(len) {
352 unsafe {
353 arch::mul_wide(
354 &left[..chunk_len],
355 &right[..chunk_len],
356 &mut result[..chunk_len],
357 );
358 }
359 if chunk_len < len {
360 scalar_mul(
361 &left[chunk_len..],
362 &right[chunk_len..],
363 &mut result[chunk_len..],
364 );
365 }
366 return;
367 }
368 }
369 scalar_mul(left, right, result);
370 }
371
372 #[inline]
373 fn dot_slice(left: &[Self], right: &[Self]) -> Self {
374 let len = left.len();
375 #[cfg(target_arch = "x86_64")]
376 {
377 if let Some(chunk_len) = native_vector_chunk_len(len) {
378 let mut sum = unsafe { arch::dot_wide(&left[..chunk_len], &right[..chunk_len]) };
379 if chunk_len < len {
380 sum += scalar_dot(&left[chunk_len..], &right[chunk_len..]);
381 }
382 return sum;
383 }
384 }
385 scalar_dot(left, right)
386 }
387
388 #[inline]
389 fn sum_slice(data: &[Self]) -> Self {
390 let len = data.len();
391 #[cfg(target_arch = "x86_64")]
392 {
393 if let Some(chunk_len) = native_vector_chunk_len(len) {
394 let mut total = unsafe { arch::sum_wide(&data[..chunk_len]) };
395 if chunk_len < len {
396 total += data[chunk_len..].iter().copied().sum::<Self>();
397 }
398 return total;
399 }
400 }
401 data.iter().copied().sum()
402 }
403}
404
405impl SimdReal for f64 {
406 #[inline]
407 fn from_len(len: usize) -> Self {
408 len as Self
409 }
410
411 #[inline]
412 fn variance_slice(data: &[Self]) -> Self {
413 let len = data.len();
414 #[cfg(target_arch = "x86_64")]
415 {
416 if let Some(chunk_len) = native_vector_chunk_len(len) {
417 let mean = Self::mean_slice(data);
418 let mut total = unsafe { arch::squared_diff_sum_wide(&data[..chunk_len], mean) };
419 if chunk_len < len {
420 total += data[chunk_len..]
421 .iter()
422 .copied()
423 .map(|value| {
424 let diff = value - mean;
425 diff * diff
426 })
427 .sum::<Self>();
428 }
429 return total / Self::from_len(len);
430 }
431 }
432
433 let mean = Self::mean_slice(data);
434 data.iter()
435 .copied()
436 .map(|value| {
437 let diff = value - mean;
438 diff * diff
439 })
440 .sum::<Self>()
441 / Self::from_len(len)
442 }
443}
444
445impl sealed::Sealed for i32 {}
446impl SimdScalar for i32 {
447 const ZERO: Self = 0;
448}
449
450impl sealed::Sealed for i64 {}
451impl SimdScalar for i64 {
452 const ZERO: Self = 0;
453}
454
455impl sealed::Sealed for u32 {}
456impl SimdScalar for u32 {
457 const ZERO: Self = 0;
458}
459
460impl sealed::Sealed for u64 {}
461impl SimdScalar for u64 {
462 const ZERO: Self = 0;
463}
464
465impl sealed::Sealed for usize {}
466impl SimdScalar for usize {
467 const ZERO: Self = 0;
468}
469
470#[inline]
471fn native_vector_available() -> bool {
472 #[cfg(target_arch = "x86_64")]
473 {
474 if arch::has_avx2_support() {
475 return true;
476 }
477 }
478
479 #[cfg(target_arch = "aarch64")]
480 {
481 if arch::has_neon_support() {
482 return true;
483 }
484 }
485
486 false
487}
488
489#[inline]
490fn uses_native_vector_path(len: usize) -> bool {
491 #[cfg(target_arch = "x86_64")]
492 {
493 if arch::has_avx2_support() && len >= arch::LANES {
494 return true;
495 }
496 }
497
498 #[cfg(target_arch = "aarch64")]
499 {
500 if arch::has_neon_support() && len >= arch::LANES {
501 return true;
502 }
503 }
504
505 false
506}
507
508#[inline]
509fn native_wide_vector_available() -> bool {
510 #[cfg(target_arch = "x86_64")]
511 {
512 if arch::has_avx2_support() {
513 return true;
514 }
515 }
516
517 false
518}
519
520#[inline]
521fn uses_native_wide_vector_path(len: usize) -> bool {
522 #[cfg(target_arch = "x86_64")]
523 {
524 if arch::has_avx2_support() && len >= arch::LANES {
525 return true;
526 }
527 }
528
529 let _ = len;
530 false
531}
532
533#[cfg(any(target_arch = "x86_64", target_arch = "aarch64"))]
534#[inline]
535fn native_vector_chunk_len(len: usize) -> Option<usize> {
536 native_vector_available()
537 .then_some((len / arch::LANES) * arch::LANES)
538 .filter(|chunk_len| *chunk_len != 0)
539}
540
541#[inline]
542fn scalar_matrix_shape<const N: usize>(left: &[f32], right: &[f32], result: &mut [f32]) {
543 assert!(N != 0, "matrix dimension must be non-zero");
544 let expected = N.checked_mul(N).expect("matrix dimension overflow");
545 assert_eq!(left.len(), expected, "left matrix size must equal N * N");
546 assert_eq!(right.len(), expected, "right matrix size must equal N * N");
547 assert_eq!(
548 result.len(),
549 expected,
550 "result matrix size must equal N * N"
551 );
552}