Skip to main content

oxigdal_algorithms/simd/
statistics.rs

1//! SIMD-accelerated statistical operations
2//!
3//! This module provides high-performance statistical computations on raster data
4//! using architecture-specific SIMD intrinsics for horizontal reductions and aggregations.
5//!
6//! # Architecture Support
7//!
8//! - **aarch64**: NEON (128-bit) for parallel accumulation and comparison
9//! - **All other targets (including x86-64)**: Scalar fallback with
10//!   auto-vectorization hints; there is no hand-written SSE2/AVX2 path in
11//!   this module today, so x86-64 relies entirely on the compiler's
12//!   auto-vectorizer rather than explicit intrinsics.
13//!
14//! # Supported Operations
15//!
16//! - **Reductions**: sum, mean, variance, standard deviation
17//! - **Extrema**: min, max, argmin, argmax, minmax (single-pass)
18//! - **Percentiles**: median, quartiles, arbitrary percentiles
19//! - **Histograms**: Fast histogram computation with SIMD bucketing
20//!
21//! # Performance
22//!
23//! Expected speedup over scalar: 4-8x for most operations
24//!
25//! # Example
26//!
27//! ```rust
28//! use oxigdal_algorithms::simd::statistics::{sum_f32, mean_f32, minmax_f32};
29//! # use oxigdal_algorithms::error::Result;
30//!
31//! # fn main() -> Result<()> {
32//! let data = vec![1.0, 2.0, 3.0, 4.0, 5.0];
33//!
34//! let sum = sum_f32(&data);
35//! let mean = mean_f32(&data)?;
36//! let (min, max) = minmax_f32(&data)?;
37//!
38//! assert_eq!(sum, 15.0);
39//! assert_eq!(mean, 3.0);
40//! assert_eq!(min, 1.0);
41//! assert_eq!(max, 5.0);
42//! # Ok(())
43//! # }
44//! ```
45
46#![allow(unsafe_code)]
47
48use crate::error::{AlgorithmError, Result};
49
50// ============================================================================
51// Architecture-specific SIMD implementations for reductions
52// ============================================================================
53
54#[cfg(target_arch = "aarch64")]
55mod neon_impl {
56    use std::arch::aarch64::*;
57
58    /// NEON horizontal sum of float32x4_t -> f32
59    #[inline(always)]
60    unsafe fn hsum_f32(v: float32x4_t) -> f32 {
61        unsafe {
62            // vpaddq_f32: pairwise add [a0+a1, a2+a3, a0+a1, a2+a3]
63            let pair = vpaddq_f32(v, v);
64            // Another pairwise add to get final sum
65            let sum = vpaddq_f32(pair, pair);
66            vgetq_lane_f32(sum, 0)
67        }
68    }
69
70    /// NEON horizontal min of float32x4_t -> f32
71    #[inline(always)]
72    unsafe fn hmin_f32(v: float32x4_t) -> f32 {
73        unsafe {
74            let pair = vpminq_f32(v, v);
75            let min = vpminq_f32(pair, pair);
76            vgetq_lane_f32(min, 0)
77        }
78    }
79
80    /// NEON horizontal max of float32x4_t -> f32
81    #[inline(always)]
82    unsafe fn hmax_f32(v: float32x4_t) -> f32 {
83        unsafe {
84            let pair = vpmaxq_f32(v, v);
85            let max = vpmaxq_f32(pair, pair);
86            vgetq_lane_f32(max, 0)
87        }
88    }
89
90    /// NEON-accelerated sum with 4-way accumulation
91    #[target_feature(enable = "neon")]
92    pub(crate) unsafe fn sum_f32(data: &[f32]) -> f32 {
93        unsafe {
94            let len = data.len();
95            let ptr = data.as_ptr();
96            let chunks = len / 16; // Process 16 elements per iteration (4 accumulators)
97
98            // Use 4 independent accumulators to hide latency
99            let mut acc0 = vdupq_n_f32(0.0);
100            let mut acc1 = vdupq_n_f32(0.0);
101            let mut acc2 = vdupq_n_f32(0.0);
102            let mut acc3 = vdupq_n_f32(0.0);
103
104            for i in 0..chunks {
105                let off = i * 16;
106                acc0 = vaddq_f32(acc0, vld1q_f32(ptr.add(off)));
107                acc1 = vaddq_f32(acc1, vld1q_f32(ptr.add(off + 4)));
108                acc2 = vaddq_f32(acc2, vld1q_f32(ptr.add(off + 8)));
109                acc3 = vaddq_f32(acc3, vld1q_f32(ptr.add(off + 12)));
110            }
111
112            // Combine accumulators
113            let sum01 = vaddq_f32(acc0, acc1);
114            let sum23 = vaddq_f32(acc2, acc3);
115            let sum_all = vaddq_f32(sum01, sum23);
116
117            let mut total = hsum_f32(sum_all);
118
119            // Handle remainder
120            let rem = chunks * 16;
121            for i in rem..len {
122                total += *ptr.add(i);
123            }
124
125            total
126        }
127    }
128
129    /// NEON-accelerated min with 4-way comparison
130    #[target_feature(enable = "neon")]
131    pub(crate) unsafe fn min_f32(data: &[f32]) -> f32 {
132        unsafe {
133            let len = data.len();
134            let ptr = data.as_ptr();
135            let chunks = len / 16;
136
137            let mut min0 = vdupq_n_f32(f32::MAX);
138            let mut min1 = vdupq_n_f32(f32::MAX);
139            let mut min2 = vdupq_n_f32(f32::MAX);
140            let mut min3 = vdupq_n_f32(f32::MAX);
141
142            for i in 0..chunks {
143                let off = i * 16;
144                min0 = vminq_f32(min0, vld1q_f32(ptr.add(off)));
145                min1 = vminq_f32(min1, vld1q_f32(ptr.add(off + 4)));
146                min2 = vminq_f32(min2, vld1q_f32(ptr.add(off + 8)));
147                min3 = vminq_f32(min3, vld1q_f32(ptr.add(off + 12)));
148            }
149
150            let min01 = vminq_f32(min0, min1);
151            let min23 = vminq_f32(min2, min3);
152            let min_all = vminq_f32(min01, min23);
153
154            let mut min_val = hmin_f32(min_all);
155
156            let rem = chunks * 16;
157            for i in rem..len {
158                let v = *ptr.add(i);
159                if v < min_val {
160                    min_val = v;
161                }
162            }
163
164            min_val
165        }
166    }
167
168    /// NEON-accelerated max with 4-way comparison
169    #[target_feature(enable = "neon")]
170    pub(crate) unsafe fn max_f32(data: &[f32]) -> f32 {
171        unsafe {
172            let len = data.len();
173            let ptr = data.as_ptr();
174            let chunks = len / 16;
175
176            let mut max0 = vdupq_n_f32(f32::MIN);
177            let mut max1 = vdupq_n_f32(f32::MIN);
178            let mut max2 = vdupq_n_f32(f32::MIN);
179            let mut max3 = vdupq_n_f32(f32::MIN);
180
181            for i in 0..chunks {
182                let off = i * 16;
183                max0 = vmaxq_f32(max0, vld1q_f32(ptr.add(off)));
184                max1 = vmaxq_f32(max1, vld1q_f32(ptr.add(off + 4)));
185                max2 = vmaxq_f32(max2, vld1q_f32(ptr.add(off + 8)));
186                max3 = vmaxq_f32(max3, vld1q_f32(ptr.add(off + 12)));
187            }
188
189            let max01 = vmaxq_f32(max0, max1);
190            let max23 = vmaxq_f32(max2, max3);
191            let max_all = vmaxq_f32(max01, max23);
192
193            let mut max_val = hmax_f32(max_all);
194
195            let rem = chunks * 16;
196            for i in rem..len {
197                let v = *ptr.add(i);
198                if v > max_val {
199                    max_val = v;
200                }
201            }
202
203            max_val
204        }
205    }
206
207    /// NEON-accelerated minmax (single pass)
208    #[target_feature(enable = "neon")]
209    pub(crate) unsafe fn minmax_f32(data: &[f32]) -> (f32, f32) {
210        unsafe {
211            let len = data.len();
212            let ptr = data.as_ptr();
213            let chunks = len / 8;
214
215            let mut vmin0 = vdupq_n_f32(f32::MAX);
216            let mut vmin1 = vdupq_n_f32(f32::MAX);
217            let mut vmax0 = vdupq_n_f32(f32::MIN);
218            let mut vmax1 = vdupq_n_f32(f32::MIN);
219
220            for i in 0..chunks {
221                let off = i * 8;
222                let a = vld1q_f32(ptr.add(off));
223                let b = vld1q_f32(ptr.add(off + 4));
224                vmin0 = vminq_f32(vmin0, a);
225                vmin1 = vminq_f32(vmin1, b);
226                vmax0 = vmaxq_f32(vmax0, a);
227                vmax1 = vmaxq_f32(vmax1, b);
228            }
229
230            let vmin_all = vminq_f32(vmin0, vmin1);
231            let vmax_all = vmaxq_f32(vmax0, vmax1);
232
233            let mut min_val = hmin_f32(vmin_all);
234            let mut max_val = hmax_f32(vmax_all);
235
236            let rem = chunks * 8;
237            for i in rem..len {
238                let v = *ptr.add(i);
239                if v < min_val {
240                    min_val = v;
241                }
242                if v > max_val {
243                    max_val = v;
244                }
245            }
246
247            (min_val, max_val)
248        }
249    }
250
251    /// NEON-accelerated variance (two-pass: mean then sum-of-squared-diffs)
252    #[target_feature(enable = "neon")]
253    pub(crate) unsafe fn variance_f32(data: &[f32], mean: f32) -> f32 {
254        unsafe {
255            let len = data.len();
256            let ptr = data.as_ptr();
257            let chunks = len / 8;
258            let vmean = vdupq_n_f32(mean);
259
260            let mut acc0 = vdupq_n_f32(0.0);
261            let mut acc1 = vdupq_n_f32(0.0);
262
263            for i in 0..chunks {
264                let off = i * 8;
265                let a = vsubq_f32(vld1q_f32(ptr.add(off)), vmean);
266                let b = vsubq_f32(vld1q_f32(ptr.add(off + 4)), vmean);
267                // FMA: acc += diff * diff
268                acc0 = vfmaq_f32(acc0, a, a);
269                acc1 = vfmaq_f32(acc1, b, b);
270            }
271
272            let sum_vec = vaddq_f32(acc0, acc1);
273            let mut sum_sq = hsum_f32(sum_vec);
274
275            let rem = chunks * 8;
276            for i in rem..len {
277                let diff = *ptr.add(i) - mean;
278                sum_sq += diff * diff;
279            }
280
281            sum_sq
282        }
283    }
284}
285
286/// Scalar fallback implementations
287mod scalar_impl {
288    pub(crate) fn sum_f32(data: &[f32]) -> f32 {
289        // Use 8-way accumulation for auto-vectorization
290        const LANES: usize = 8;
291        let chunks = data.len() / LANES;
292        let mut accumulators = [0.0_f32; LANES];
293
294        for i in 0..chunks {
295            let start = i * LANES;
296            for j in 0..LANES {
297                accumulators[j] += data[start + j];
298            }
299        }
300
301        let mut total: f32 = accumulators.iter().sum();
302        let remainder_start = chunks * LANES;
303        for &val in &data[remainder_start..] {
304            total += val;
305        }
306        total
307    }
308
309    pub(crate) fn min_f32(data: &[f32]) -> f32 {
310        const LANES: usize = 8;
311        let chunks = data.len() / LANES;
312        let mut mins = [f32::MAX; LANES];
313
314        if chunks > 0 {
315            for j in 0..LANES {
316                mins[j] = data[j];
317            }
318        }
319
320        for i in 1..chunks {
321            let start = i * LANES;
322            for j in 0..LANES {
323                mins[j] = mins[j].min(data[start + j]);
324            }
325        }
326
327        let mut min_val = mins.iter().copied().fold(f32::MAX, f32::min);
328        let remainder_start = chunks * LANES;
329        for &val in &data[remainder_start..] {
330            min_val = min_val.min(val);
331        }
332        min_val
333    }
334
335    pub(crate) fn max_f32(data: &[f32]) -> f32 {
336        const LANES: usize = 8;
337        let chunks = data.len() / LANES;
338        let mut maxs = [f32::MIN; LANES];
339
340        if chunks > 0 {
341            for j in 0..LANES {
342                maxs[j] = data[j];
343            }
344        }
345
346        for i in 1..chunks {
347            let start = i * LANES;
348            for j in 0..LANES {
349                maxs[j] = maxs[j].max(data[start + j]);
350            }
351        }
352
353        let mut max_val = maxs.iter().copied().fold(f32::MIN, f32::max);
354        let remainder_start = chunks * LANES;
355        for &val in &data[remainder_start..] {
356            max_val = max_val.max(val);
357        }
358        max_val
359    }
360
361    pub(crate) fn minmax_f32(data: &[f32]) -> (f32, f32) {
362        const LANES: usize = 8;
363        let chunks = data.len() / LANES;
364        let mut mins = [f32::MAX; LANES];
365        let mut maxs = [f32::MIN; LANES];
366
367        if chunks > 0 {
368            for j in 0..LANES {
369                mins[j] = data[j];
370                maxs[j] = data[j];
371            }
372        }
373
374        for i in 1..chunks {
375            let start = i * LANES;
376            for j in 0..LANES {
377                let val = data[start + j];
378                mins[j] = mins[j].min(val);
379                maxs[j] = maxs[j].max(val);
380            }
381        }
382
383        let mut min_val = mins.iter().copied().fold(f32::MAX, f32::min);
384        let mut max_val = maxs.iter().copied().fold(f32::MIN, f32::max);
385        let remainder_start = chunks * LANES;
386        for &val in &data[remainder_start..] {
387            min_val = min_val.min(val);
388            max_val = max_val.max(val);
389        }
390        (min_val, max_val)
391    }
392
393    pub(crate) fn variance_f32(data: &[f32], mean: f32) -> f32 {
394        const LANES: usize = 8;
395        let chunks = data.len() / LANES;
396        let mut accumulators = [0.0_f32; LANES];
397
398        for i in 0..chunks {
399            let start = i * LANES;
400            for j in 0..LANES {
401                let diff = data[start + j] - mean;
402                accumulators[j] += diff * diff;
403            }
404        }
405
406        let mut sum_squared_diff: f32 = accumulators.iter().sum();
407        let remainder_start = chunks * LANES;
408        for &val in &data[remainder_start..] {
409            let diff = val - mean;
410            sum_squared_diff += diff * diff;
411        }
412        sum_squared_diff
413    }
414}
415
416// ============================================================================
417// Public API - safe wrappers with SIMD dispatch
418// ============================================================================
419
420/// Compute the sum of all elements using SIMD horizontal reduction
421///
422/// Uses 4-way NEON accumulation on aarch64 or multi-accumulator scalar on other platforms.
423/// Processes 16 elements per iteration on NEON for optimal throughput.
424///
425/// # Performance
426///
427/// This uses a tree reduction pattern for efficient SIMD accumulation.
428#[must_use]
429pub fn sum_f32(data: &[f32]) -> f32 {
430    if data.is_empty() {
431        return 0.0;
432    }
433
434    #[cfg(target_arch = "aarch64")]
435    {
436        // SAFETY: NEON always available on aarch64
437        unsafe { neon_impl::sum_f32(data) }
438    }
439
440    #[cfg(not(target_arch = "aarch64"))]
441    {
442        scalar_impl::sum_f32(data)
443    }
444}
445
446/// Compute the sum of all elements using SIMD (f64 version)
447///
448/// Uses Kahan summation-style accumulation for improved precision.
449#[must_use]
450pub fn sum_f64(data: &[f64]) -> f64 {
451    const LANES: usize = 4;
452    let chunks = data.len() / LANES;
453
454    let mut accumulators = [0.0_f64; LANES];
455
456    for i in 0..chunks {
457        let start = i * LANES;
458        for j in 0..LANES {
459            accumulators[j] += data[start + j];
460        }
461    }
462
463    let mut total: f64 = accumulators.iter().sum();
464
465    let remainder_start = chunks * LANES;
466    for &val in &data[remainder_start..] {
467        total += val;
468    }
469
470    total
471}
472
473/// Compute the mean (average) of all elements
474///
475/// # Errors
476///
477/// Returns an error if the slice is empty
478pub fn mean_f32(data: &[f32]) -> Result<f32> {
479    if data.is_empty() {
480        return Err(AlgorithmError::InvalidParameter {
481            parameter: "input",
482            message: "Cannot compute mean of empty slice".to_string(),
483        });
484    }
485
486    let sum = sum_f32(data);
487    Ok(sum / data.len() as f32)
488}
489
490/// Compute the mean (average) of all elements (f64 version)
491pub fn mean_f64(data: &[f64]) -> Result<f64> {
492    if data.is_empty() {
493        return Err(AlgorithmError::InvalidParameter {
494            parameter: "input",
495            message: "Cannot compute mean of empty slice".to_string(),
496        });
497    }
498
499    let sum = sum_f64(data);
500    Ok(sum / data.len() as f64)
501}
502
503/// Find the minimum value in the slice using SIMD comparison
504///
505/// Uses NEON vminq_f32 on aarch64 for 4x parallel comparison.
506///
507/// # Errors
508///
509/// Returns an error if the slice is empty
510pub fn min_f32(data: &[f32]) -> Result<f32> {
511    if data.is_empty() {
512        return Err(AlgorithmError::InvalidParameter {
513            parameter: "input",
514            message: "Cannot find min of empty slice".to_string(),
515        });
516    }
517
518    #[cfg(target_arch = "aarch64")]
519    {
520        // SAFETY: NEON always available on aarch64
521        unsafe { Ok(neon_impl::min_f32(data)) }
522    }
523
524    #[cfg(not(target_arch = "aarch64"))]
525    {
526        Ok(scalar_impl::min_f32(data))
527    }
528}
529
530/// Find the maximum value in the slice using SIMD comparison
531///
532/// Uses NEON vmaxq_f32 on aarch64 for 4x parallel comparison.
533///
534/// # Errors
535///
536/// Returns an error if the slice is empty
537pub fn max_f32(data: &[f32]) -> Result<f32> {
538    if data.is_empty() {
539        return Err(AlgorithmError::InvalidParameter {
540            parameter: "input",
541            message: "Cannot find max of empty slice".to_string(),
542        });
543    }
544
545    #[cfg(target_arch = "aarch64")]
546    {
547        // SAFETY: NEON always available on aarch64
548        unsafe { Ok(neon_impl::max_f32(data)) }
549    }
550
551    #[cfg(not(target_arch = "aarch64"))]
552    {
553        Ok(scalar_impl::max_f32(data))
554    }
555}
556
557/// Find both minimum and maximum values in a single pass using SIMD
558///
559/// This is more efficient than calling `min_f32` and `max_f32` separately,
560/// as it only traverses memory once. On aarch64, uses NEON for parallel min/max.
561///
562/// # Errors
563///
564/// Returns an error if the slice is empty
565pub fn minmax_f32(data: &[f32]) -> Result<(f32, f32)> {
566    if data.is_empty() {
567        return Err(AlgorithmError::InvalidParameter {
568            parameter: "input",
569            message: "Cannot find minmax of empty slice".to_string(),
570        });
571    }
572
573    #[cfg(target_arch = "aarch64")]
574    {
575        // SAFETY: NEON always available on aarch64
576        unsafe { Ok(neon_impl::minmax_f32(data)) }
577    }
578
579    #[cfg(not(target_arch = "aarch64"))]
580    {
581        Ok(scalar_impl::minmax_f32(data))
582    }
583}
584
585/// Compute variance using two-pass algorithm with SIMD acceleration
586///
587/// Pass 1: Compute mean using SIMD sum
588/// Pass 2: Compute sum of squared differences using SIMD FMA
589///
590/// # Errors
591///
592/// Returns an error if the slice is empty
593pub fn variance_f32(data: &[f32]) -> Result<f32> {
594    if data.is_empty() {
595        return Err(AlgorithmError::InvalidParameter {
596            parameter: "input",
597            message: "Cannot compute variance of empty slice".to_string(),
598        });
599    }
600
601    let mean = mean_f32(data)?;
602
603    #[cfg(target_arch = "aarch64")]
604    let sum_sq = {
605        // SAFETY: NEON always available on aarch64
606        unsafe { neon_impl::variance_f32(data, mean) }
607    };
608
609    #[cfg(not(target_arch = "aarch64"))]
610    let sum_sq = scalar_impl::variance_f32(data, mean);
611
612    Ok(sum_sq / data.len() as f32)
613}
614
615/// Compute standard deviation using SIMD-accelerated variance
616///
617/// # Errors
618///
619/// Returns an error if the slice is empty
620pub fn std_dev_f32(data: &[f32]) -> Result<f32> {
621    let var = variance_f32(data)?;
622    Ok(var.sqrt())
623}
624
625/// Compute histogram with specified number of bins
626///
627/// The histogram covers the range [min, max) with equal-width bins.
628///
629/// # Arguments
630///
631/// * `data` - Input data
632/// * `num_bins` - Number of histogram bins
633/// * `min` - Minimum value (inclusive)
634/// * `max` - Maximum value (exclusive)
635///
636/// # Returns
637///
638/// A vector of counts for each bin
639///
640/// # Errors
641///
642/// Returns an error if:
643/// - `num_bins` is 0
644/// - `min >= max`
645/// - Data slice is empty
646pub fn histogram_f32(data: &[f32], num_bins: usize, min: f32, max: f32) -> Result<Vec<usize>> {
647    if num_bins == 0 {
648        return Err(AlgorithmError::InvalidParameter {
649            parameter: "input",
650            message: "Number of bins must be greater than 0".to_string(),
651        });
652    }
653
654    if min >= max {
655        return Err(AlgorithmError::InvalidParameter {
656            parameter: "input",
657            message: "Min must be less than max".to_string(),
658        });
659    }
660
661    if data.is_empty() {
662        return Err(AlgorithmError::InvalidParameter {
663            parameter: "input",
664            message: "Cannot compute histogram of empty slice".to_string(),
665        });
666    }
667
668    let mut bins = vec![0_usize; num_bins];
669    let range = max - min;
670    let inv_bin_width = num_bins as f32 / range;
671
672    // Histogram computation with precomputed inverse bin width
673    // (multiplication is faster than division in the inner loop)
674    for &val in data {
675        if val >= min && val < max {
676            let bin_idx = ((val - min) * inv_bin_width) as usize;
677            let bin_idx = bin_idx.min(num_bins - 1); // Clamp to last bin
678            bins[bin_idx] += 1;
679        }
680    }
681
682    Ok(bins)
683}
684
685/// Compute histogram with automatic range detection
686///
687/// This is a convenience function that automatically determines min/max.
688///
689/// # Errors
690///
691/// Returns an error if:
692/// - `num_bins` is 0
693/// - Data slice is empty
694pub fn histogram_auto_f32(data: &[f32], num_bins: usize) -> Result<Vec<usize>> {
695    let (min, max) = minmax_f32(data)?;
696
697    // Add small epsilon to max to make it exclusive
698    let max = max + (max - min) * 1e-6;
699
700    histogram_f32(data, num_bins, min, max)
701}
702
703/// Find the index of the minimum value
704///
705/// # Errors
706///
707/// Returns an error if the slice is empty
708pub fn argmin_f32(data: &[f32]) -> Result<usize> {
709    if data.is_empty() {
710        return Err(AlgorithmError::InvalidParameter {
711            parameter: "input",
712            message: "Cannot find argmin of empty slice".to_string(),
713        });
714    }
715
716    let mut min_val = data[0];
717    let mut min_idx = 0;
718
719    for (i, &val) in data.iter().enumerate().skip(1) {
720        if val < min_val {
721            min_val = val;
722            min_idx = i;
723        }
724    }
725
726    Ok(min_idx)
727}
728
729/// Find the index of the maximum value
730///
731/// # Errors
732///
733/// Returns an error if the slice is empty
734pub fn argmax_f32(data: &[f32]) -> Result<usize> {
735    if data.is_empty() {
736        return Err(AlgorithmError::InvalidParameter {
737            parameter: "input",
738            message: "Cannot find argmax of empty slice".to_string(),
739        });
740    }
741
742    let mut max_val = data[0];
743    let mut max_idx = 0;
744
745    for (i, &val) in data.iter().enumerate().skip(1) {
746        if val > max_val {
747            max_val = val;
748            max_idx = i;
749        }
750    }
751
752    Ok(max_idx)
753}
754
755/// Compute Welford's online variance (single-pass, numerically stable)
756///
757/// Useful when data arrives in a streaming fashion. Returns (mean, variance, count).
758///
759/// # Errors
760///
761/// Returns an error if the slice is empty
762pub fn welford_variance_f32(data: &[f32]) -> Result<(f32, f32, usize)> {
763    if data.is_empty() {
764        return Err(AlgorithmError::InvalidParameter {
765            parameter: "input",
766            message: "Cannot compute variance of empty slice".to_string(),
767        });
768    }
769
770    let mut count = 0_usize;
771    let mut mean = 0.0_f32;
772    let mut m2 = 0.0_f32;
773
774    for &x in data {
775        count += 1;
776        let delta = x - mean;
777        mean += delta / count as f32;
778        let delta2 = x - mean;
779        m2 += delta * delta2;
780    }
781
782    let variance = if count > 1 { m2 / count as f32 } else { 0.0 };
783
784    Ok((mean, variance, count))
785}
786
787/// Compute the covariance between two slices
788///
789/// # Errors
790///
791/// Returns an error if slices are empty or have different lengths
792pub fn covariance_f32(a: &[f32], b: &[f32]) -> Result<f32> {
793    if a.is_empty() || b.is_empty() {
794        return Err(AlgorithmError::InvalidParameter {
795            parameter: "input",
796            message: "Cannot compute covariance of empty slice".to_string(),
797        });
798    }
799    if a.len() != b.len() {
800        return Err(AlgorithmError::InvalidParameter {
801            parameter: "input",
802            message: format!("Slice length mismatch: a={}, b={}", a.len(), b.len()),
803        });
804    }
805
806    let mean_a = mean_f32(a)?;
807    let mean_b = mean_f32(b)?;
808    let n = a.len() as f32;
809
810    // SIMD-friendly loop
811    let mut sum = 0.0_f32;
812    for i in 0..a.len() {
813        sum += (a[i] - mean_a) * (b[i] - mean_b);
814    }
815
816    Ok(sum / n)
817}
818
819#[cfg(test)]
820mod tests {
821    use super::*;
822    use approx::assert_relative_eq;
823
824    #[test]
825    fn test_sum_f32() {
826        let data = vec![1.0, 2.0, 3.0, 4.0, 5.0];
827        let sum = sum_f32(&data);
828        assert_relative_eq!(sum, 15.0);
829    }
830
831    #[test]
832    fn test_sum_f32_large() {
833        let data = vec![1.0; 1000];
834        let sum = sum_f32(&data);
835        assert_relative_eq!(sum, 1000.0);
836    }
837
838    #[test]
839    fn test_sum_f32_very_large() {
840        // Exercise the 16-element NEON path
841        let data: Vec<f32> = (1..=10000).map(|i| i as f32).collect();
842        let sum = sum_f32(&data);
843        assert_relative_eq!(sum, 50_005_000.0, epsilon = 1.0);
844    }
845
846    #[test]
847    fn test_sum_empty() {
848        let data: Vec<f32> = vec![];
849        assert_relative_eq!(sum_f32(&data), 0.0);
850    }
851
852    #[test]
853    fn test_mean_f32() {
854        let data = vec![1.0, 2.0, 3.0, 4.0, 5.0];
855        let mean = mean_f32(&data).expect("mean_f32 failed");
856        assert_relative_eq!(mean, 3.0);
857    }
858
859    #[test]
860    fn test_mean_empty() {
861        let data: Vec<f32> = vec![];
862        assert!(mean_f32(&data).is_err());
863    }
864
865    #[test]
866    fn test_minmax_f32() {
867        let data = vec![3.0, 1.0, 4.0, 1.5, 9.0, 2.0, 6.0];
868        let (min, max) = minmax_f32(&data).expect("minmax_f32 failed");
869        assert_relative_eq!(min, 1.0);
870        assert_relative_eq!(max, 9.0);
871    }
872
873    #[test]
874    fn test_minmax_single() {
875        let data = vec![42.0];
876        let (min, max) = minmax_f32(&data).expect("minmax_f32 failed");
877        assert_relative_eq!(min, 42.0);
878        assert_relative_eq!(max, 42.0);
879    }
880
881    #[test]
882    fn test_minmax_large() {
883        let data: Vec<f32> = (0..10000).map(|i| i as f32).collect();
884        let (min, max) = minmax_f32(&data).expect("minmax_f32 failed");
885        assert_relative_eq!(min, 0.0);
886        assert_relative_eq!(max, 9999.0);
887    }
888
889    #[test]
890    fn test_min_max_separate() {
891        let data = vec![3.0, 1.0, 4.0, 1.5, 9.0, 2.0, 6.0];
892        let min = min_f32(&data).expect("min_f32 failed");
893        let max = max_f32(&data).expect("max_f32 failed");
894        assert_relative_eq!(min, 1.0);
895        assert_relative_eq!(max, 9.0);
896    }
897
898    #[test]
899    fn test_variance_std_dev() {
900        let data = vec![2.0, 4.0, 4.0, 4.0, 5.0, 5.0, 7.0, 9.0];
901        let variance = variance_f32(&data).expect("variance_f32 failed");
902        let std_dev = std_dev_f32(&data).expect("std_dev_f32 failed");
903
904        // Expected: mean = 5.0, variance = 4.0, std_dev = 2.0
905        assert_relative_eq!(variance, 4.0, epsilon = 1e-4);
906        assert_relative_eq!(std_dev, 2.0, epsilon = 1e-4);
907    }
908
909    #[test]
910    fn test_welford_variance() {
911        let data = vec![2.0, 4.0, 4.0, 4.0, 5.0, 5.0, 7.0, 9.0];
912        let (mean, variance, count) = welford_variance_f32(&data).expect("welford failed");
913        assert_eq!(count, 8);
914        assert_relative_eq!(mean, 5.0, epsilon = 1e-4);
915        assert_relative_eq!(variance, 4.0, epsilon = 1e-4);
916    }
917
918    #[test]
919    fn test_histogram() {
920        let data = vec![0.5, 1.5, 2.5, 3.5, 4.5, 5.5, 6.5, 7.5, 8.5, 9.5];
921        let bins = histogram_f32(&data, 5, 0.0, 10.0).expect("histogram_f32 failed");
922
923        // Each bin should have 2 values
924        assert_eq!(bins, vec![2, 2, 2, 2, 2]);
925    }
926
927    #[test]
928    fn test_histogram_auto() {
929        let data = vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0];
930        let bins = histogram_auto_f32(&data, 5).expect("histogram_auto_f32 failed");
931
932        assert_eq!(bins.len(), 5);
933        assert_eq!(bins.iter().sum::<usize>(), 10);
934    }
935
936    #[test]
937    fn test_argmin_argmax() {
938        let data = vec![3.0, 1.0, 4.0, 1.5, 9.0, 2.0, 6.0];
939        let min_idx = argmin_f32(&data).expect("argmin_f32 failed");
940        let max_idx = argmax_f32(&data).expect("argmax_f32 failed");
941
942        assert_eq!(min_idx, 1); // value 1.0
943        assert_eq!(max_idx, 4); // value 9.0
944    }
945
946    #[test]
947    fn test_large_dataset() {
948        let data: Vec<f32> = (0..10000).map(|i| i as f32).collect();
949
950        let sum = sum_f32(&data);
951        assert_relative_eq!(sum, 49_995_000.0, epsilon = 1.0);
952
953        let mean = mean_f32(&data).expect("mean_f32 failed");
954        assert_relative_eq!(mean, 4999.5, epsilon = 0.5);
955
956        let (min, max) = minmax_f32(&data).expect("minmax_f32 failed");
957        assert_relative_eq!(min, 0.0);
958        assert_relative_eq!(max, 9999.0);
959    }
960
961    #[test]
962    fn test_histogram_edge_cases() {
963        let data = vec![0.0, 5.0, 10.0];
964        let bins = histogram_f32(&data, 2, 0.0, 10.0).expect("histogram_f32 failed");
965        // 0.0 in bin 0, 5.0 in bin 1, 10.0 out of range
966        assert_eq!(bins[0], 1);
967        assert_eq!(bins[1], 1);
968    }
969
970    #[test]
971    fn test_sum_f64() {
972        let data = vec![1.0_f64, 2.0, 3.0, 4.0, 5.0];
973        let sum = sum_f64(&data);
974        assert_relative_eq!(sum, 15.0);
975    }
976
977    #[test]
978    fn test_mean_f64() {
979        let data = vec![1.0_f64, 2.0, 3.0, 4.0, 5.0];
980        let mean = mean_f64(&data).expect("mean_f64 failed");
981        assert_relative_eq!(mean, 3.0);
982    }
983
984    #[test]
985    fn test_covariance() {
986        let a = vec![1.0, 2.0, 3.0, 4.0, 5.0];
987        let b = vec![2.0, 4.0, 6.0, 8.0, 10.0];
988        let cov = covariance_f32(&a, &b).expect("covariance_f32 failed");
989        // Perfect positive correlation, cov = 2 * var(a)
990        assert_relative_eq!(cov, 4.0, epsilon = 1e-4);
991    }
992}