Skip to main content

moirai_utils/simd/
scalar.rs

1//! Sealed scalar contracts and scalar fallback kernels.
2
3use 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/// Native-precision scalar contract for SIMD-aware slice operations.
12///
13/// Implementations are sealed so native backend invariants stay under crate
14/// control. The hidden methods form the monomorphized backend dispatch surface
15/// used by the public functions in this module.
16#[allow(private_bounds)]
17pub trait SimdScalar:
18    sealed::Sealed + Copy + Send + Sync + Add<Output = Self> + Mul<Output = Self> + Sum<Self> + 'static
19{
20    /// Additive identity.
21    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/// Scalar contract for native-precision real-valued statistics.
75#[allow(private_bounds)]
76pub trait SimdReal: SimdScalar + Sub<Output = Self> + Div<Output = Self> {
77    /// Converts a non-zero slice length into the scalar's native representation.
78    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}