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
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/// Native-precision scalar contract for SIMD-aware slice operations.
21///
22/// Implementations are sealed so native backend invariants stay under crate
23/// control. The hidden methods form the monomorphized backend dispatch surface
24/// used by the public functions in this module.
25#[allow(private_bounds)]
26pub trait SimdScalar:
27    sealed::Sealed + Copy + Send + Sync + Add<Output = Self> + Mul<Output = Self> + Sum<Self> + 'static
28{
29    /// Additive identity.
30    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/// Scalar contract for native-precision real-valued statistics.
84#[allow(private_bounds)]
85pub trait SimdReal: SimdScalar + Sub<Output = Self> + Div<Output = Self> {
86    /// Converts a non-zero slice length into the scalar's native representation.
87    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                // SAFETY: `chunk_len` is a LANES multiple <= len with the ISA
185                // feature probed, so the sliced arguments satisfy the arch
186                // fn contract.
187                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                // SAFETY: `chunk_len` is a LANES multiple <= len with the ISA
214                // feature probed, so the sliced arguments satisfy the arch
215                // fn contract.
216                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            // SAFETY: N == 4 fixes the 16-element shape and the feature
273            // probe passed, satisfying the arch fn contract.
274            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                // SAFETY: `chunk_len` is an 8-lane multiple <= len with AVX2
344                // probed, so the sliced arguments satisfy the arch fn
345                // contract.
346                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                // SAFETY: `chunk_len` is an 8-lane multiple <= len with AVX2
373                // probed, so the sliced arguments satisfy the arch fn
374                // contract.
375                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}