Skip to main content

legume_numeric/matrix/
ndarray_stat.rs

1use crate::matrix::{
2    common_io::{file_ext, write_lines},
3    traits::{IoOps, RunningStatOps},
4};
5use ndarray::{stack, ArrayBase, Axis, Data, Dimension, NdIndex, OwnedRepr, RemoveAxis};
6
7/// A container to keep track of sufficient statistics of an arbitrary
8/// shape `ndarray`
9///
10/// # Type parameters
11/// - `S` : The shape of the array
12///
13#[derive(Clone)]
14pub struct RunningStatistics<S>
15where
16    S: Dimension + RemoveAxis,
17{
18    npos: ArrayBase<OwnedRepr<f32>, S>,
19    s0: ArrayBase<OwnedRepr<f32>, S>,
20    s1: ArrayBase<OwnedRepr<f32>, S>,
21    s2: ArrayBase<OwnedRepr<f32>, S>,
22}
23
24impl<S> RunningStatistics<S>
25where
26    S: Dimension + RemoveAxis,
27{
28    /// Create a new RunningStatistics object
29    ///
30    /// # Arguments
31    ///
32    /// * `shape` - The shape of the array
33    ///
34    /// # Examples
35    ///
36    /// ```
37    /// use legume_numeric::matrix::ndarray_stat::RunningStatistics;
38    /// use ndarray::Ix1;
39    /// let nrow = 10;
40    /// RunningStatistics::new(Ix1(nrow));
41    /// ```
42    ///
43    pub fn new(shape: S) -> Self {
44        let npos = ArrayBase::zeros(shape.clone());
45        let s0 = ArrayBase::zeros(shape.clone());
46        let s1 = ArrayBase::zeros(shape.clone());
47        let s2 = ArrayBase::zeros(shape);
48
49        RunningStatistics { npos, s0, s1, s2 }
50    }
51
52    pub fn add<V>(&mut self, xx: &ArrayBase<V, S>)
53    where
54        V: Data<Elem = f32>,
55    {
56        self.npos += &xx.mapv(Self::_is_positive);
57        self.s0 += &xx.mapv(Self::_is_finite);
58        self.s1 += &xx.mapv(Self::_finite);
59        self.s2 += &xx.mapv(Self::_finite).mapv(|v| v * v);
60    }
61
62    pub fn add_element<I>(&mut self, idx: &I, val: f32)
63    where
64        I: NdIndex<S> + Clone,
65    {
66        fn get<'a, S, I>(mat: &'a mut ArrayBase<OwnedRepr<f32>, S>, idx: &'a I) -> &'a mut f32
67        where
68            S: Dimension + RemoveAxis,
69            I: NdIndex<S> + Clone,
70        {
71            mat.get_mut(idx.clone()).expect("failed to access matrix")
72        }
73
74        let idx_clone = idx.clone();
75
76        *get(&mut self.npos, &idx_clone) += Self::_is_positive(val);
77        *get(&mut self.s0, &idx_clone) += Self::_is_finite(val);
78        let safe_val = Self::_finite(val);
79        *get(&mut self.s1, &idx_clone) += safe_val;
80        *get(&mut self.s2, &idx_clone) += safe_val * safe_val;
81    }
82
83    pub fn clear(&mut self) {
84        self.npos.fill(0.0);
85        self.s0.fill(0.0);
86        self.s1.fill(0.0);
87        self.s2.fill(0.0);
88    }
89
90    /// Frequency of positive values. For a sparse count matrix, this
91    /// will reflect the number of non-zero values
92    ///
93    pub fn count_positives(&self) -> ArrayBase<OwnedRepr<f32>, S> {
94        self.npos.clone()
95    }
96
97    /// Sum of values
98    pub fn sum(&self) -> ArrayBase<OwnedRepr<f32>, S> {
99        self.s1.clone()
100    }
101
102    /// Average statistic
103    pub fn mean(&self) -> ArrayBase<OwnedRepr<f32>, S> {
104        self.s1.clone() / &self.s0.mapv(Self::_add_pseudo_count)
105    }
106
107    /// Variance
108    pub fn variance(&self) -> ArrayBase<OwnedRepr<f32>, S> {
109        let mean = self.mean();
110        let nn = &self.s0.mapv(Self::_add_pseudo_count);
111
112        &self.s2 / nn - &mean * &mean
113    }
114
115    /// Standard deviation
116    pub fn std(&self) -> ArrayBase<OwnedRepr<f32>, S> {
117        self.variance().mapv(f32::sqrt)
118    }
119
120    pub fn shape(&self) -> &[usize] {
121        self.s0.shape()
122    }
123
124    //////////////////////
125    // helper functions //
126    //////////////////////
127
128    fn _finite(x: f32) -> f32 {
129        if x.is_finite() {
130            x
131        } else {
132            0_f32
133        }
134    }
135
136    fn _is_finite(x: f32) -> f32 {
137        if x.is_finite() {
138            1_f32
139        } else {
140            0_f32
141        }
142    }
143
144    fn _is_positive(x: f32) -> f32 {
145        if x.is_finite() && x > 0_f32 {
146            1_f32
147        } else {
148            0_f32
149        }
150    }
151
152    fn _add_pseudo_count(x: f32) -> f32 {
153        x + 1e-8
154    }
155
156    /// Save the statistics to a file
157    /// # Arguments
158    /// * `filename` - The name of the file to save the statistics to
159    /// * `names` - The names of the statistics
160    /// * `sep` - Separator for text formats
161    /// * `row_column_name` - Name for the row column in parquet format (defaults to "stat")
162    pub fn save(
163        &self,
164        filename: &str,
165        names: &[Box<str>],
166        sep: &str,
167        row_column_name: Option<&str>,
168    ) -> anyhow::Result<()> {
169        match file_ext(filename).unwrap_or(Box::from("")).as_ref() {
170            "parquet" => {
171                let nnz = &self.count_positives();
172                let tot = &self.s1;
173                let mu = &self.mean();
174                let sig = &self.std();
175
176                let n = nnz.len();
177                let nnz_col = nnz.clone().into_shape_with_order((n,)).unwrap();
178                let tot_col = tot.clone().into_shape_with_order((n,)).unwrap();
179                let mu_col = mu.clone().into_shape_with_order((n,)).unwrap();
180                let sig_col = sig.clone().into_shape_with_order((n,)).unwrap();
181
182                let stacked = stack(
183                    Axis(1),
184                    &[
185                        nnz_col.view(),
186                        tot_col.view(),
187                        mu_col.view(),
188                        sig_col.view(),
189                    ],
190                )
191                .unwrap();
192
193                let column_names: Vec<Box<str>> = vec!["nnz", "tot", "mu", "sig"]
194                    .into_iter()
195                    .map(|s| s.into())
196                    .collect();
197
198                let row_col = row_column_name.or(Some("stat"));
199                stacked.to_parquet_with_names(
200                    filename,
201                    (Some(names), row_col),
202                    Some(&column_names),
203                )?;
204            }
205            _ => {
206                let mut out = self.to_string_vec(names, sep)?;
207                let header = format!("#name{}nnz{}tot{}mu{}sig", sep, sep, sep, sep);
208                out.insert(0, header.into_boxed_str());
209                write_lines(&out, filename)?;
210            }
211        };
212
213        Ok(())
214    }
215
216    pub fn to_string_vec(&self, names: &[Box<str>], sep: &str) -> anyhow::Result<Vec<Box<str>>> {
217        if names.len() != self.shape()[0] {
218            anyhow::bail!(
219                "The number of names does not match the number of the first dimension of the statistics"
220            );
221        }
222
223        let nnz_: Vec<Box<str>> = to_string_vec(&self.count_positives(), sep);
224        let tot_ = to_string_vec(&self.s1, sep);
225        let mu_: Vec<Box<str>> = to_string_vec(&self.mean(), sep);
226        let sig_: Vec<Box<str>> = to_string_vec(&self.std(), sep);
227
228        let out: Vec<Box<str>> = (0..self.shape()[0])
229            .map(|i| {
230                format!(
231                    "{}{}{}{}{}{}{}{}{}",
232                    names[i], sep, nnz_[i], sep, tot_[i], sep, mu_[i], sep, sig_[i]
233                )
234                .into_boxed_str()
235            })
236            .collect();
237        Ok(out)
238    }
239}
240
241impl<S> RunningStatOps<f32> for RunningStatistics<S>
242where
243    S: Dimension + RemoveAxis,
244{
245    type Output = ArrayBase<OwnedRepr<f32>, S>;
246
247    fn clear(&mut self) {
248        self.npos.fill(0.0);
249        self.s0.fill(0.0);
250        self.s1.fill(0.0);
251        self.s2.fill(0.0);
252    }
253
254    fn count_positives(&self) -> Self::Output {
255        self.npos.clone()
256    }
257
258    fn sum(&self) -> Self::Output {
259        self.s1.clone()
260    }
261
262    fn mean(&self) -> Self::Output {
263        self.s1.clone() / &self.s0.mapv(Self::_add_pseudo_count)
264    }
265
266    fn variance(&self) -> Self::Output {
267        let mean = <Self as RunningStatOps<f32>>::mean(self);
268        let nn = &self.s0.mapv(Self::_add_pseudo_count);
269        &self.s2 / nn - &mean * &mean
270    }
271
272    fn std(&self) -> Self::Output {
273        <Self as RunningStatOps<f32>>::variance(self).mapv(f32::sqrt)
274    }
275}
276
277fn to_string_vec<S>(xx: &ArrayBase<OwnedRepr<f32>, S>, sep: &str) -> Vec<Box<str>>
278where
279    S: Dimension + RemoveAxis,
280{
281    xx.axis_iter(Axis(0))
282        .map(|m| {
283            m.iter()
284                .map(|v| {
285                    if *v > 1e-4 {
286                        format!("{:.4}", v)
287                            .trim_end_matches('0')
288                            .trim_end_matches('.')
289                            .to_string()
290                    } else if *v > 1e-20 {
291                        format!("{:.4e}", v)
292                    } else {
293                        "0".to_string()
294                    }
295                })
296                .collect::<Vec<String>>()
297                .join(sep)
298                .clone()
299                .into_boxed_str()
300        })
301        .collect()
302}