1#[allow(unused_imports)]
6use crate::prelude::*;
7use core::ops::{Add, Mul};
8use num_traits::Zero;
9
10pub fn simd_sum<T>(values: &[T]) -> T
12where
13 T: Copy + Add<Output = T> + Zero,
14{
15 if values.len() < 8 {
17 return values.iter().copied().fold(T::zero(), Add::add);
18 }
19
20 let chunk_size = 8;
22 let mut sum = T::zero();
23
24 for chunk in values.chunks(chunk_size) {
25 let chunk_sum = chunk.iter().copied().fold(T::zero(), Add::add);
26 sum = sum + chunk_sum;
27 }
28
29 sum
30}
31
32pub fn simd_dot_product<T>(a: &[T], b: &[T]) -> T
34where
35 T: Copy + Add<Output = T> + Mul<Output = T> + Zero,
36{
37 assert_eq!(a.len(), b.len(), "vectors must have same length");
38
39 if a.len() < 8 {
40 return a
41 .iter()
42 .zip(b.iter())
43 .map(|(&x, &y)| x * y)
44 .fold(T::zero(), Add::add);
45 }
46
47 let chunk_size = 8;
49 let mut sum = T::zero();
50
51 let chunks_a = a.chunks(chunk_size);
52 let chunks_b = b.chunks(chunk_size);
53
54 for (chunk_a, chunk_b) in chunks_a.zip(chunks_b) {
55 let chunk_sum = chunk_a
56 .iter()
57 .zip(chunk_b.iter())
58 .map(|(&x, &y)| x * y)
59 .fold(T::zero(), Add::add);
60 sum = sum + chunk_sum;
61 }
62
63 sum
64}
65
66pub fn simd_norm_squared<T>(values: &[T]) -> T
68where
69 T: Copy + Add<Output = T> + Mul<Output = T> + Zero,
70{
71 simd_dot_product(values, values)
72}
73
74pub fn simd_matrix_vec_mul<T>(matrix: &[Vec<T>], vec: &[T]) -> Vec<T>
76where
77 T: Copy + Add<Output = T> + Mul<Output = T> + Zero,
78{
79 assert!(!matrix.is_empty(), "matrix must not be empty");
80 assert_eq!(
81 matrix[0].len(),
82 vec.len(),
83 "matrix columns must match vector length"
84 );
85
86 matrix
87 .iter()
88 .map(|row| simd_dot_product(row, vec))
89 .collect()
90}
91
92pub fn simd_vec_add<T>(a: &[T], b: &[T]) -> Vec<T>
94where
95 T: Copy + Add<Output = T>,
96{
97 assert_eq!(a.len(), b.len(), "vectors must have same length");
98
99 a.iter().zip(b.iter()).map(|(&x, &y)| x + y).collect()
100}
101
102pub fn simd_vec_mul<T>(a: &[T], b: &[T]) -> Vec<T>
104where
105 T: Copy + Mul<Output = T>,
106{
107 assert_eq!(a.len(), b.len(), "vectors must have same length");
108
109 a.iter().zip(b.iter()).map(|(&x, &y)| x * y).collect()
110}
111
112pub fn simd_vec_scale<T>(vec: &[T], scalar: T) -> Vec<T>
114where
115 T: Copy + Mul<Output = T>,
116{
117 vec.iter().map(|&x| x * scalar).collect()
118}
119
120pub fn simd_weighted_sum<T>(values: &[T], weights: &[T]) -> T
122where
123 T: Copy + Add<Output = T> + Mul<Output = T> + Zero,
124{
125 simd_dot_product(values, weights)
126}
127
128pub fn simd_reduce<T, F>(values: &[T], init: T, op: F) -> T
130where
131 T: Copy,
132 F: Fn(T, T) -> T,
133{
134 if values.is_empty() {
135 return init;
136 }
137
138 let mut working = values.to_vec();
140
141 while working.len() > 1 {
142 let mut next = Vec::with_capacity(working.len().div_ceil(2));
143
144 for chunk in working.chunks(2) {
145 if chunk.len() == 2 {
146 next.push(op(chunk[0], chunk[1]));
147 } else {
148 next.push(chunk[0]);
149 }
150 }
151
152 working = next;
153 }
154
155 working[0]
156}
157
158pub fn simd_max<T>(values: &[T]) -> Option<T>
160where
161 T: Copy + PartialOrd,
162{
163 if values.is_empty() {
164 return None;
165 }
166
167 Some(simd_reduce(
168 values,
169 values[0],
170 |a, b| {
171 if a > b { a } else { b }
172 },
173 ))
174}
175
176pub fn simd_min<T>(values: &[T]) -> Option<T>
178where
179 T: Copy + PartialOrd,
180{
181 if values.is_empty() {
182 return None;
183 }
184
185 Some(simd_reduce(
186 values,
187 values[0],
188 |a, b| {
189 if a < b { a } else { b }
190 },
191 ))
192}
193
194pub fn simd_mean_f64(values: &[f64]) -> Option<f64> {
196 if values.is_empty() {
197 return None;
198 }
199
200 let sum = simd_sum(values);
201 Some(sum / values.len() as f64)
202}
203
204#[cfg(test)]
205mod tests {
206 use super::*;
207
208 #[test]
209 fn test_simd_sum() {
210 let values = vec![1.0, 2.0, 3.0, 4.0, 5.0];
211 let sum = simd_sum(&values);
212 assert_eq!(sum, 15.0);
213 }
214
215 #[test]
216 fn test_simd_dot_product() {
217 let a = vec![1.0, 2.0, 3.0, 4.0];
218 let b = vec![5.0, 6.0, 7.0, 8.0];
219 let dot = simd_dot_product(&a, &b);
220 assert_eq!(dot, 1.0 * 5.0 + 2.0 * 6.0 + 3.0 * 7.0 + 4.0 * 8.0);
221 }
222
223 #[test]
224 fn test_simd_norm_squared() {
225 let values = vec![3.0, 4.0];
226 let norm_sq = simd_norm_squared(&values);
227 assert_eq!(norm_sq, 25.0);
228 }
229
230 #[test]
231 fn test_simd_matrix_vec_mul() {
232 let matrix = vec![vec![1.0, 2.0], vec![3.0, 4.0]];
233 let vec = vec![5.0, 6.0];
234 let result = simd_matrix_vec_mul(&matrix, &vec);
235 assert_eq!(result, vec![17.0, 39.0]);
236 }
237
238 #[test]
239 fn test_simd_vec_add() {
240 let a = vec![1.0, 2.0, 3.0];
241 let b = vec![4.0, 5.0, 6.0];
242 let result = simd_vec_add(&a, &b);
243 assert_eq!(result, vec![5.0, 7.0, 9.0]);
244 }
245
246 #[test]
247 fn test_simd_max() {
248 let values = vec![3.0, 1.0, 4.0, 1.0, 5.0];
249 let max = simd_max(&values);
250 assert_eq!(max, Some(5.0));
251 }
252
253 #[test]
254 fn test_simd_min() {
255 let values = vec![3.0, 1.0, 4.0, 1.0, 5.0];
256 let min = simd_min(&values);
257 assert_eq!(min, Some(1.0));
258 }
259
260 #[test]
261 fn test_simd_reduce() {
262 let values = vec![1, 2, 3, 4, 5];
263 let sum = simd_reduce(&values, 0, |a, b| a + b);
264 assert_eq!(sum, 15);
265 }
266}