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#[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 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 pub fn count_positives(&self) -> ArrayBase<OwnedRepr<f32>, S> {
94 self.npos.clone()
95 }
96
97 pub fn sum(&self) -> ArrayBase<OwnedRepr<f32>, S> {
99 self.s1.clone()
100 }
101
102 pub fn mean(&self) -> ArrayBase<OwnedRepr<f32>, S> {
104 self.s1.clone() / &self.s0.mapv(Self::_add_pseudo_count)
105 }
106
107 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 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 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 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}