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