1use scirs2_core::ndarray::{Array1, Array2, Axis, NdFloat};
4use scirs2_core::numeric::FromPrimitive;
5
6pub fn mean_axis<'a, D>(array: &'a Array2<D>, axis: Axis) -> Array1<D>
8where
9 D: NdFloat + FromPrimitive + 'a,
10{
11 array.mean_axis(axis).expect("operation should succeed")
12}
13
14pub fn var_axis<'a, D>(array: &'a Array2<D>, axis: Axis, ddof: usize) -> Array1<D>
16where
17 D: NdFloat + FromPrimitive + 'a,
18{
19 let mean = array.mean_axis(axis).expect("operation should succeed");
20 let n = array.len_of(axis);
21
22 if axis == Axis(0) {
23 let mut var = Array1::zeros(array.ncols());
25 for j in 0..array.ncols() {
26 let col = array.column(j);
27 let m = mean[j];
28 let sum_sq: D = col.mapv(|x| (x - m).powi(2)).sum();
29 var[j] = sum_sq / D::from(n - ddof).expect("operation should succeed");
30 }
31 var
32 } else {
33 let mut var = Array1::zeros(array.nrows());
35 for i in 0..array.nrows() {
36 let row = array.row(i);
37 let m = mean[i];
38 let sum_sq: D = row.mapv(|x| (x - m).powi(2)).sum();
39 var[i] = sum_sq / D::from(n - ddof).expect("operation should succeed");
40 }
41 var
42 }
43}
44
45pub fn std_axis<'a, D>(array: &'a Array2<D>, axis: Axis, ddof: usize) -> Array1<D>
47where
48 D: NdFloat + FromPrimitive + 'a,
49{
50 var_axis(array, axis, ddof).mapv(|v| v.sqrt())
51}
52
53pub fn covariance<D>(x: &Array2<D>, ddof: usize) -> Array2<D>
55where
56 D: NdFloat + FromPrimitive,
57{
58 let n_samples = x.nrows();
59
60 let mean = x.mean_axis(Axis(0)).expect("operation should succeed");
62 let centered = x - &mean;
63
64 let cov =
66 centered.t().dot(¢ered) / D::from(n_samples - ddof).expect("operation should succeed");
67 cov
68}
69
70#[allow(non_snake_case)]
71#[cfg(test)]
72mod tests {
73 use super::*;
74 use approx::assert_abs_diff_eq;
75 use scirs2_core::ndarray::array;
76
77 #[test]
78 fn test_mean_axis() {
79 let x = array![[1.0, 2.0, 3.0], [4.0, 5.0, 6.0], [7.0, 8.0, 9.0]];
80
81 let mean_rows = mean_axis(&x, Axis(0));
82 assert_abs_diff_eq!(mean_rows[0], 4.0, epsilon = 1e-10);
83 assert_abs_diff_eq!(mean_rows[1], 5.0, epsilon = 1e-10);
84 assert_abs_diff_eq!(mean_rows[2], 6.0, epsilon = 1e-10);
85
86 let mean_cols = mean_axis(&x, Axis(1));
87 assert_abs_diff_eq!(mean_cols[0], 2.0, epsilon = 1e-10);
88 assert_abs_diff_eq!(mean_cols[1], 5.0, epsilon = 1e-10);
89 assert_abs_diff_eq!(mean_cols[2], 8.0, epsilon = 1e-10);
90 }
91
92 #[test]
93 fn test_var_axis() {
94 let x = array![[1.0, 2.0, 3.0], [4.0, 5.0, 6.0], [7.0, 8.0, 9.0]];
95
96 let var_rows = var_axis(&x, Axis(0), 0);
97 assert_abs_diff_eq!(var_rows[0], 6.0, epsilon = 1e-10);
98 assert_abs_diff_eq!(var_rows[1], 6.0, epsilon = 1e-10);
99 assert_abs_diff_eq!(var_rows[2], 6.0, epsilon = 1e-10);
100 }
101}