Skip to main content

ruprim_host/simd/
kernels.rs

1//! Portable SIMD kernels using macerator.
2//!
3//! These kernels work across all architectures supported by macerator:
4//! - aarch64 (NEON)
5//! - x86_64 (AVX2, AVX512, SSE)
6//! - wasm32 (SIMD128)
7//! - Scalar fallback for embedded/other platforms
8
9use core::iter::Sum;
10use core::ops::AddAssign;
11
12use macerator::{
13    ReduceAdd, ReduceMax, ReduceMin, Simd, VAdd, VOrd, vload_unaligned, vstore_unaligned,
14};
15
16// ============================================================================
17// Sum reduction
18// ============================================================================
19
20/// Sum all elements in a f32 slice using SIMD with 4 accumulators.
21#[inline]
22pub fn sum_f32(data: &[f32]) -> f32 {
23    macerator_sum(data)
24}
25
26/// 8-accumulator SIMD sum. Independent accumulator chains let the CPU
27/// pipeline floating-point adds and hide L2 cache latency.
28#[macerator::with_simd]
29fn macerator_sum<S: Simd, F: VAdd + Sum + ReduceAdd>(mut xs: &[F]) -> F {
30    let lanes = F::lanes::<S>();
31    let stride = lanes * 8;
32    let zero = F::default().splat::<S>();
33    let (mut s0, mut s1, mut s2, mut s3) = (zero, zero, zero, zero);
34    let (mut s4, mut s5, mut s6, mut s7) = (zero, zero, zero, zero);
35
36    while xs.len() >= stride {
37        unsafe {
38            let p = xs.as_ptr();
39            s0 += vload_unaligned(p);
40            s1 += vload_unaligned(p.add(lanes));
41            s2 += vload_unaligned(p.add(lanes * 2));
42            s3 += vload_unaligned(p.add(lanes * 3));
43            s4 += vload_unaligned(p.add(lanes * 4));
44            s5 += vload_unaligned(p.add(lanes * 5));
45            s6 += vload_unaligned(p.add(lanes * 6));
46            s7 += vload_unaligned(p.add(lanes * 7));
47        }
48        xs = &xs[stride..];
49    }
50
51    // Combine 8 accumulators into one, then drain remaining full vectors
52    let mut sum = ((s0 + s1) + (s2 + s3)) + ((s4 + s5) + (s6 + s7));
53    while xs.len() >= lanes {
54        sum += unsafe { vload_unaligned(xs.as_ptr()) };
55        xs = &xs[lanes..];
56    }
57
58    sum.reduce_add() + xs.iter().copied().sum()
59}
60
61// ============================================================================
62// Scatter-add for dimension reductions
63// ============================================================================
64
65/// Scatter-add: for each row, add to corresponding output positions.
66/// Used for cache-friendly first-dim and middle-dim reductions.
67///
68/// # Arguments
69/// * `src` - Source data pointer
70/// * `dst` - Destination accumulator (must be pre-zeroed)
71/// * `num_rows` - Number of rows to sum
72/// * `row_len` - Length of each row (columns)
73/// * `src_row_stride` - Stride between source rows
74#[macerator::with_simd]
75pub fn scatter_add_f32<S: Simd, F: VAdd + AddAssign>(
76    src: &[F],
77    dst: &mut [F],
78    num_rows: usize,
79    row_len: usize,
80    src_row_stride: usize,
81) {
82    let lanes = F::lanes::<S>();
83
84    for row in 0..num_rows {
85        let row_start = row * src_row_stride;
86        let row_data = &src[row_start..row_start + row_len];
87
88        let simd_len = row_len / lanes * lanes;
89
90        // SIMD accumulate
91        let mut i = 0;
92        while i < simd_len {
93            unsafe {
94                let s = vload_unaligned(row_data.as_ptr().add(i));
95                let d = vload_unaligned(dst.as_ptr().add(i));
96                vstore_unaligned::<S, _>(dst.as_mut_ptr().add(i), d + s);
97            }
98            i += lanes;
99        }
100
101        // Scalar tail
102        for j in simd_len..row_len {
103            dst[j] += row_data[j];
104        }
105    }
106}
107
108/// Batched scatter-add for middle-dim reductions.
109/// For tensors like [B, M, K] reducing dim=1.
110#[macerator::with_simd]
111pub fn scatter_add_batched<S: Simd, F: VAdd + AddAssign>(
112    src: &[F],
113    dst: &mut [F],
114    num_batches: usize,
115    num_rows: usize,
116    row_len: usize,
117    batch_stride: usize,
118    row_stride: usize,
119) {
120    let lanes = F::lanes::<S>();
121
122    for batch in 0..num_batches {
123        let batch_src_start = batch * batch_stride;
124        let batch_dst_start = batch * row_len;
125        let batch_dst = &mut dst[batch_dst_start..batch_dst_start + row_len];
126
127        for row in 0..num_rows {
128            let row_start = batch_src_start + row * row_stride;
129            let row_data = &src[row_start..row_start + row_len];
130
131            let simd_len = row_len / lanes * lanes;
132
133            let mut i = 0;
134            while i < simd_len {
135                unsafe {
136                    let s = vload_unaligned(row_data.as_ptr().add(i));
137                    let d = vload_unaligned(batch_dst.as_ptr().add(i));
138                    vstore_unaligned::<S, _>(batch_dst.as_mut_ptr().add(i), d + s);
139                }
140                i += lanes;
141            }
142
143            for j in simd_len..row_len {
144                batch_dst[j] += row_data[j];
145            }
146        }
147    }
148}
149
150// ============================================================================
151// Row-wise sum (last-dim reduction)
152// ============================================================================
153
154/// Sum each row, storing results in output slice.
155/// Used for last-dim reductions.
156#[inline]
157pub fn sum_rows_f32(src: &[f32], dst: &mut [f32], num_rows: usize, row_len: usize) {
158    debug_assert_eq!(dst.len(), num_rows, "dst length must equal num_rows");
159    debug_assert!(
160        src.len() >= num_rows * row_len,
161        "src too short: need {} elements, got {}",
162        num_rows * row_len,
163        src.len()
164    );
165    for (row, dst_val) in dst.iter_mut().enumerate() {
166        let row_start = row * row_len;
167        let row_data = &src[row_start..row_start + row_len];
168        *dst_val = macerator_sum(row_data);
169    }
170}
171
172// ============================================================================
173// Max/Min reduction
174// ============================================================================
175
176/// Find the maximum element in a f32 slice using SIMD.
177#[inline]
178pub fn max_f32(data: &[f32]) -> f32 {
179    macerator_max(data, f32::NEG_INFINITY)
180}
181
182/// Find the minimum element in a f32 slice using SIMD.
183#[inline]
184pub fn min_f32(data: &[f32]) -> f32 {
185    macerator_min(data, f32::INFINITY)
186}
187
188#[macerator::with_simd]
189fn macerator_max<S: Simd, F: VOrd + ReduceMax + PartialOrd>(mut xs: &[F], init: F) -> F {
190    let lanes = F::lanes::<S>();
191    let mut acc = init.splat::<S>();
192
193    while xs.len() >= lanes {
194        let v = unsafe { vload_unaligned(xs.as_ptr()) };
195        acc = acc.max(v);
196        xs = &xs[lanes..];
197    }
198
199    let mut result = acc.reduce_max();
200    for &x in xs {
201        if x > result {
202            result = x;
203        }
204    }
205    result
206}
207
208#[macerator::with_simd]
209fn macerator_min<S: Simd, F: VOrd + ReduceMin + PartialOrd>(mut xs: &[F], init: F) -> F {
210    let lanes = F::lanes::<S>();
211    let mut acc = init.splat::<S>();
212
213    while xs.len() >= lanes {
214        let v = unsafe { vload_unaligned(xs.as_ptr()) };
215        acc = acc.min(v);
216        xs = &xs[lanes..];
217    }
218
219    let mut result = acc.reduce_min();
220    for &x in xs {
221        if x < result {
222            result = x;
223        }
224    }
225    result
226}
227
228#[cfg(test)]
229mod tests {
230    use super::*;
231
232    #[test]
233    fn test_sum_f32() {
234        let data: Vec<f32> = (0..1000).map(|i| i as f32).collect();
235        let expected: f32 = data.iter().sum();
236        let result = sum_f32(&data);
237        assert!((result - expected).abs() < 0.01);
238    }
239
240    #[test]
241    fn test_sum_f32_empty() {
242        let data: Vec<f32> = vec![];
243        assert_eq!(sum_f32(&data), 0.0);
244    }
245
246    #[test]
247    fn test_sum_f32_small() {
248        let data = vec![1.0, 2.0, 3.0];
249        assert_eq!(sum_f32(&data), 6.0);
250    }
251
252    #[test]
253    fn test_scatter_add_f32() {
254        // Simulate reducing [3, 4] along dim=0 -> [1, 4]
255        let src = vec![
256            1.0, 2.0, 3.0, 4.0, // row 0
257            5.0, 6.0, 7.0, 8.0, // row 1
258            9.0, 10.0, 11.0, 12.0, // row 2
259        ];
260        let mut dst = vec![0.0; 4];
261
262        scatter_add_f32(&src, &mut dst, 3, 4, 4);
263
264        assert_eq!(dst, vec![15.0, 18.0, 21.0, 24.0]);
265    }
266
267    #[test]
268    fn test_sum_rows_f32() {
269        // Simulate reducing [3, 4] along dim=1 -> [3, 1]
270        let src = vec![
271            1.0, 2.0, 3.0, 4.0, // row 0 -> 10
272            5.0, 6.0, 7.0, 8.0, // row 1 -> 26
273            9.0, 10.0, 11.0, 12.0, // row 2 -> 42
274        ];
275        let mut dst = vec![0.0; 3];
276
277        sum_rows_f32(&src, &mut dst, 3, 4);
278
279        assert_eq!(dst, vec![10.0, 26.0, 42.0]);
280    }
281
282    #[test]
283    fn test_max_f32() {
284        let data: Vec<f32> = (0..1000).map(|i| i as f32).collect();
285        assert_eq!(max_f32(&data), 999.0);
286    }
287
288    #[test]
289    fn test_max_f32_small() {
290        let data = vec![3.0, 1.0, 4.0, 1.0, 5.0];
291        assert_eq!(max_f32(&data), 5.0);
292    }
293
294    #[test]
295    fn test_max_f32_negative() {
296        let data = vec![-3.0, -1.0, -4.0, -1.0, -5.0];
297        assert_eq!(max_f32(&data), -1.0);
298    }
299
300    #[test]
301    fn test_min_f32() {
302        let data: Vec<f32> = (0..1000).map(|i| i as f32).collect();
303        assert_eq!(min_f32(&data), 0.0);
304    }
305
306    #[test]
307    fn test_min_f32_small() {
308        let data = vec![3.0, 1.0, 4.0, 1.0, 5.0];
309        assert_eq!(min_f32(&data), 1.0);
310    }
311
312    #[test]
313    fn test_min_f32_negative() {
314        let data = vec![-3.0, -1.0, -4.0, -1.0, -5.0];
315        assert_eq!(min_f32(&data), -5.0);
316    }
317}