wowlab_types/stats/
simd.rs1#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
4const F64_LANES: usize = 2;
5const SQUARE_EXP: i32 = 2;
6
7#[inline]
9pub fn add_f64_slices(dst: &mut [f64], src: &[f64]) {
10 debug_assert!(
11 dst.len() >= src.len(),
12 "destination must fit every source element"
13 );
14 #[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
15 {
16 simd_add_f64_slices(dst, src);
17 }
18
19 #[cfg(not(all(target_arch = "wasm32", target_feature = "simd128")))]
20 {
21 scalar_add_f64_slices(dst, src);
22 }
23}
24
25#[inline]
27#[must_use]
28pub fn sum_squared_deviations(values: &[f64], mean: f64) -> f64 {
29 #[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
30 {
31 simd_sum_squared_deviations(values, mean)
32 }
33
34 #[cfg(not(all(target_arch = "wasm32", target_feature = "simd128")))]
35 {
36 scalar_sum_squared_deviations(values, mean)
37 }
38}
39
40#[inline]
41fn scalar_add_f64_slices(dst: &mut [f64], src: &[f64]) {
42 for (d, &s) in dst.iter_mut().zip(src.iter()) {
43 *d += s;
44 }
45}
46
47#[inline]
48fn scalar_sum_squared_deviations(values: &[f64], mean: f64) -> f64 {
49 values.iter().map(|&x| (x - mean).powi(SQUARE_EXP)).sum()
50}
51
52#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
53mod wasm_simd {
54 use core::arch::wasm32::{
55 f64x2_add, f64x2_extract_lane, f64x2_mul, f64x2_splat, f64x2_sub, v128, v128_load,
56 v128_store,
57 };
58
59 use super::F64_LANES;
60
61 #[inline]
62 pub(super) fn simd_add_f64_slices(dst: &mut [f64], src: &[f64]) {
63 let len = src.len();
64 let chunks = len / F64_LANES;
65 let remainder = len % F64_LANES;
66
67 for i in 0..chunks {
68 let offset = i * F64_LANES;
69 unsafe {
72 let a = v128_load(dst.as_ptr().add(offset) as *const v128);
73 let b = v128_load(src.as_ptr().add(offset) as *const v128);
74 let sum = f64x2_add(a, b);
75
76 v128_store(dst.as_mut_ptr().add(offset) as *mut v128, sum);
77 }
78 }
79
80 if remainder > 0 {
81 let tail = chunks * F64_LANES;
82
83 if let (Some(d), Some(&s)) = (dst.get_mut(tail), src.get(tail)) {
84 *d += s;
85 }
86 }
87 }
88
89 #[inline]
90 pub(super) fn simd_sum_squared_deviations(values: &[f64], mean: f64) -> f64 {
91 let len = values.len();
92 let chunks = len / F64_LANES;
93 let remainder = len % F64_LANES;
94
95 let mean_v = f64x2_splat(mean);
96 let mut acc = f64x2_splat(0.0);
97
98 for i in 0..chunks {
99 let offset = i * F64_LANES;
100 unsafe {
103 let v = v128_load(values.as_ptr().add(offset) as *const v128);
104 let diff = f64x2_sub(v, mean_v);
105 let sq = f64x2_mul(diff, diff);
106
107 acc = f64x2_add(acc, sq);
108 }
109 }
110
111 let mut total = f64x2_extract_lane::<0>(acc) + f64x2_extract_lane::<1>(acc);
112
113 if remainder > 0 {
114 if let Some(&val) = values.get(chunks * F64_LANES) {
115 let x = val - mean;
116
117 total += x * x;
118 }
119 }
120
121 total
122 }
123}
124
125#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
126use wasm_simd::{simd_add_f64_slices, simd_sum_squared_deviations};
127
128#[cfg(test)]
129mod tests {
130 use googletest::prelude::*;
131 use wowlab_test_support::near_tol;
132
133 use super::{add_f64_slices, sum_squared_deviations};
134
135 #[gtest]
136 fn add_slices_basic() -> Result<()> {
137 let mut dst = vec![1.0, 2.0, 3.0, 4.0, 5.0];
138 let src = vec![10.0, 20.0, 30.0, 40.0, 50.0];
139
140 add_f64_slices(&mut dst, &src);
141
142 verify_that!(dst, container_eq([11.0, 22.0, 33.0, 44.0, 55.0]))
143 }
144
145 #[gtest]
146 fn add_slices_dst_longer() -> Result<()> {
147 let mut dst = vec![1.0, 2.0, 3.0, 4.0, 5.0];
148 let src = vec![10.0, 20.0, 30.0];
149
150 add_f64_slices(&mut dst, &src);
151
152 verify_that!(dst, container_eq([11.0, 22.0, 33.0, 4.0, 5.0]))
153 }
154
155 #[gtest]
156 fn add_slices_empty() -> Result<()> {
157 let mut dst = vec![1.0, 2.0];
158 let src: Vec<f64> = vec![];
159
160 add_f64_slices(&mut dst, &src);
161
162 verify_that!(dst, container_eq([1.0, 2.0]))
163 }
164
165 #[gtest]
166 fn add_slices_odd_length() -> Result<()> {
167 let mut dst = vec![1.0, 2.0, 3.0];
168 let src = vec![10.0, 20.0, 30.0];
169
170 add_f64_slices(&mut dst, &src);
171
172 verify_that!(dst, container_eq([11.0, 22.0, 33.0]))
173 }
174
175 #[gtest]
176 fn sum_squared_deviations_basic() -> Result<()> {
177 let vals = vec![2.0, 4.0, 6.0, 8.0];
178 let mean = 5.0;
179 let result = sum_squared_deviations(&vals, mean);
180
181 verify_that!(result, near_tol(20.0))
182 }
183
184 #[gtest]
185 fn sum_squared_deviations_single() -> Result<()> {
186 let vals = vec![3.0];
187 let result = sum_squared_deviations(&vals, 3.0);
188
189 verify_that!(result, near_tol(0.0))
190 }
191
192 #[gtest]
193 fn sum_squared_deviations_empty() -> Result<()> {
194 let vals: Vec<f64> = vec![];
195 let result = sum_squared_deviations(&vals, 0.0);
196
197 verify_that!(result, near_tol(0.0))
198 }
199}