1use core::iter::Sum;
10use core::ops::AddAssign;
11
12use macerator::{
13 ReduceAdd, ReduceMax, ReduceMin, Simd, VAdd, VOrd, vload_unaligned, vstore_unaligned,
14};
15
16#[inline]
22pub fn sum_f32(data: &[f32]) -> f32 {
23 macerator_sum(data)
24}
25
26#[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 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#[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 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 for j in simd_len..row_len {
103 dst[j] += row_data[j];
104 }
105 }
106}
107
108#[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#[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#[inline]
178pub fn max_f32(data: &[f32]) -> f32 {
179 macerator_max(data, f32::NEG_INFINITY)
180}
181
182#[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 let src = vec![
256 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0, ];
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 let src = vec![
271 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0, 11.0, 12.0, ];
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}