Skip to main content

wowlab_types/stats/
simd.rs

1//! SIMD-accelerated helpers for telemetry hot paths (WASM SIMD128 on `wasm32+simd128`, autovectorized scalar elsewhere).
2
3#[cfg(all(target_arch = "wasm32", target_feature = "simd128"))]
4const F64_LANES: usize = 2;
5const SQUARE_EXP: i32 = 2;
6
7/// Element-wise add `src` into `dst`; requires `dst.len() >= src.len()`.
8#[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/// Compute the sum of `(x - mean)^2` for all elements.
26#[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            // SAFETY: offset + F64_LANES <= len; v128_load/store handle unaligned access on wasm32.
70
71            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            // SAFETY: offset + F64_LANES <= len; the slice has at least `len` elements.
101
102            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}