Skip to main content

legume_numeric/matrix/
sparse_stat.rs

1use crate::matrix::common_io::{file_ext, write_lines};
2use crate::matrix::parquet::{
3    parquet_add_numeric_column, parquet_add_string_column, ParquetWriter,
4};
5use crate::matrix::traits::RunningStatOps;
6use nalgebra_sparse::CscMatrix;
7use num_traits::{Float, ToPrimitive, Zero};
8use parquet::basic::Type as ParquetType;
9use std::fmt::Display;
10use std::iter::Sum;
11use std::ops::AddAssign;
12
13const STAT_COLUMN_NAMES: [&str; 4] = ["nnz", "tot", "mu", "sig"];
14
15/// Denominator for mean/variance: clamp to a tiny positive to avoid
16/// divide-by-zero when no samples have been accumulated yet.
17fn safe_denom<T: Float>(n: usize) -> T {
18    let n = T::from(n).unwrap_or(T::one());
19    if n > T::zero() {
20        n
21    } else {
22        T::from(1e-8).unwrap_or(T::one())
23    }
24}
25
26/// Running statistics that accepts sparse column input but stores
27/// sufficient statistics in dense vectors.
28///
29/// This is more efficient than `RunningStatistics<Ix1>` when the input
30/// data is sparse and has many rows, as it avoids materializing dense
31/// matrices during reads.
32///
33#[derive(Clone)]
34pub struct SparseRunningStatistics<T>
35where
36    T: Float,
37{
38    nrows: usize,
39    ncols_processed: usize,
40    npos: Vec<T>,
41    s1: Vec<T>,
42    s2: Vec<T>,
43}
44
45impl<T> SparseRunningStatistics<T>
46where
47    T: Float + AddAssign + Sum + Zero,
48{
49    /// Create a new SparseRunningStatistics object
50    ///
51    /// # Arguments
52    /// * `nrows` - Number of rows (features)
53    ///
54    pub fn new(nrows: usize) -> Self {
55        SparseRunningStatistics {
56            nrows,
57            ncols_processed: 0,
58            npos: vec![T::zero(); nrows],
59            s1: vec![T::zero(); nrows],
60            s2: vec![T::zero(); nrows],
61        }
62    }
63
64    /// Add a sparse column (row indices + values) to the running
65    /// statistics. The column advances `ncols_processed` by one.
66    pub fn add_sparse_column(&mut self, row_indices: &[usize], values: &[T]) {
67        debug_assert_eq!(row_indices.len(), values.len());
68
69        for (&row, &val) in row_indices.iter().zip(values.iter()) {
70            if val.is_finite() {
71                if val > T::zero() {
72                    self.npos[row] += T::one();
73                }
74                self.s1[row] += val;
75                self.s2[row] += val * val;
76            }
77        }
78        self.ncols_processed += 1;
79    }
80
81    pub fn nrows(&self) -> usize {
82        self.nrows
83    }
84
85    pub fn ncols_processed(&self) -> usize {
86        self.ncols_processed
87    }
88
89    fn denom(&self) -> T {
90        safe_denom::<T>(self.ncols_processed)
91    }
92
93    /// Convert to owned vectors (npos, sum, mean, std)
94    pub fn to_vecs(&self) -> (Vec<T>, Vec<T>, Vec<T>, Vec<T>) {
95        (self.npos.clone(), self.s1.clone(), self.mean(), self.std())
96    }
97
98    /// Add columns from a CscMatrix
99    ///
100    /// # Arguments
101    /// * `csc` - Sparse matrix in CSC format
102    ///
103    pub fn add_csc(&mut self, csc: &CscMatrix<T>) {
104        for col in csc.col_iter() {
105            let rows = col.row_indices();
106            let vals = col.values();
107            self.add_sparse_column(rows, vals);
108        }
109    }
110
111    /// Add a dense column directly. Skips zero / non-finite entries
112    /// so `npos` stays equivalent to what `add_sparse_column` would
113    /// produce — useful when an upstream coarsening step yields a dense
114    /// `[D, n]` intermediate we don't want to re-sparsify.
115    ///
116    /// Loop is structured for autovectorization: hoists the running-sum
117    /// and running-square accumulations through `zip_eq` over independent
118    /// `&mut` slices (no aliasing), and uses a branchless `1/0` mask
119    /// for the `npos` increment so the compiler can emit `cmov` /
120    /// `vmaskmovps` instead of a control-flow branch per element.
121    pub fn add_dense_column(&mut self, values: &[T]) {
122        debug_assert_eq!(values.len(), self.nrows);
123        let zero = T::zero();
124        let one = T::one();
125        for ((v_in, npos), (s1, s2)) in values
126            .iter()
127            .zip(self.npos.iter_mut())
128            .zip(self.s1.iter_mut().zip(self.s2.iter_mut()))
129        {
130            let v = *v_in;
131            // Mask non-finite to zero so finite-only invariant holds
132            // through the rest without branching the `is_finite` check.
133            let v = if v.is_finite() { v } else { zero };
134            // Branchless `+= (v > 0) as T`: compiler emits cmov.
135            let pos = if v > zero { one } else { zero };
136            *npos += pos;
137            *s1 += v;
138            *s2 += v * v;
139        }
140        self.ncols_processed += 1;
141    }
142
143    /// [`Self::add_dense_column`] with every value scaled by `scale` on the
144    /// way in. Folding the multiply into the accumulation loop touches each
145    /// element once, where scale-into-a-buffer-then-add would write and
146    /// re-read the whole column.
147    pub fn add_dense_column_scaled(&mut self, values: &[T], scale: T) {
148        debug_assert_eq!(values.len(), self.nrows);
149        let zero = T::zero();
150        let one = T::one();
151        for ((v_in, npos), (s1, s2)) in values
152            .iter()
153            .zip(self.npos.iter_mut())
154            .zip(self.s1.iter_mut().zip(self.s2.iter_mut()))
155        {
156            let v = *v_in * scale;
157            let v = if v.is_finite() { v } else { zero };
158            let pos = if v > zero { one } else { zero };
159            *npos += pos;
160            *s1 += v;
161            *s2 += v * v;
162        }
163        self.ncols_processed += 1;
164    }
165
166    /// Add every column of a dense `[D, n]` matrix in column-major
167    /// order. Calls `add_dense_column` per column so the inner loop
168    /// stays vectorizable; the per-column overhead is negligible
169    /// (one `+= 1` for `ncols_processed`) compared to the per-element
170    /// accumulation work.
171    pub fn add_dense_columns(&mut self, dense: &nalgebra::DMatrix<T>)
172    where
173        T: nalgebra::Scalar,
174    {
175        debug_assert_eq!(dense.nrows(), self.nrows);
176        for j in 0..dense.ncols() {
177            let col = dense.column(j);
178            self.add_dense_column(col.as_slice());
179        }
180    }
181
182    /// Combine another `SparseRunningStatistics` into this one. Used to
183    /// reduce per-thread accumulators back to a single result without
184    /// holding a global lock during the streaming pass.
185    pub fn merge(&mut self, other: &Self) {
186        debug_assert_eq!(self.nrows, other.nrows);
187        for (a, b) in self.npos.iter_mut().zip(other.npos.iter()) {
188            *a += *b;
189        }
190        for (a, b) in self.s1.iter_mut().zip(other.s1.iter()) {
191            *a += *b;
192        }
193        for (a, b) in self.s2.iter_mut().zip(other.s2.iter()) {
194            *a += *b;
195        }
196        self.ncols_processed += other.ncols_processed;
197    }
198}
199
200impl<T> SparseRunningStatistics<T>
201where
202    T: Float + AddAssign + Sum + Zero + Display + ToPrimitive,
203{
204    /// Save the statistics to a file (parquet if `filename` ends with
205    /// `.parquet`, otherwise a separator-delimited text file).
206    pub fn save(&self, filename: &str, names: &[Box<str>], sep: &str) -> anyhow::Result<()> {
207        let (nnz, tot, mu, sig) = self.to_f32_vecs();
208        write_stat_file(
209            filename,
210            names,
211            sep,
212            StatColumns {
213                nnz: &nnz,
214                tot: &tot,
215                mu: &mu,
216                sig: &sig,
217            },
218        )
219    }
220
221    /// Get statistics as vectors (nnz, tot, mu, sig) converted to f32
222    pub fn to_f32_vecs(&self) -> (Vec<f32>, Vec<f32>, Vec<f32>, Vec<f32>) {
223        let to_f32_slice =
224            |v: &[T]| -> Vec<f32> { v.iter().map(|x| x.to_f32().unwrap_or(0.0)).collect() };
225        let nnz = to_f32_slice(&self.npos);
226        let tot = to_f32_slice(&self.s1);
227        let mu = to_f32_slice(&self.mean());
228        let sig = to_f32_slice(&self.std());
229        (nnz, tot, mu, sig)
230    }
231
232    /// Convert statistics to string vectors for output
233    pub fn to_string_vec(&self, names: &[Box<str>], sep: &str) -> anyhow::Result<Vec<Box<str>>> {
234        if names.len() != self.nrows {
235            anyhow::bail!(
236                "The number of names ({}) does not match nrows ({})",
237                names.len(),
238                self.nrows
239            );
240        }
241
242        let nnz = self.count_positives();
243        let tot = self.sum();
244        let mu = self.mean();
245        let sig = self.std();
246
247        let out: Vec<Box<str>> = (0..self.nrows)
248            .map(|i| {
249                format!(
250                    "{}{}{}{}{}{}{}{}{}",
251                    names[i],
252                    sep,
253                    format_value(nnz[i]),
254                    sep,
255                    format_value(tot[i]),
256                    sep,
257                    format_value(mu[i]),
258                    sep,
259                    format_value(sig[i])
260                )
261                .into_boxed_str()
262            })
263            .collect();
264        Ok(out)
265    }
266}
267
268/// Running statistics computed per-column directly from sparse (CSC) input.
269///
270/// Unlike [`SparseRunningStatistics`] which tracks per-row statistics across
271/// columns, this tracks per-column statistics: nnz, sum, and sum-of-squares
272/// for each column. Mean/variance use `nrows` as the denominator (implicit
273/// zeros contribute zero values).
274///
275/// CSC slabs are consumed directly — no triplet materialization — and each
276/// slab only needs to know its starting column offset in the global index
277/// space.
278#[derive(Clone)]
279pub struct SparseColumnRunningStatistics<T>
280where
281    T: Float,
282{
283    nrows: usize,
284    npos: Vec<T>,
285    s1: Vec<T>,
286    s2: Vec<T>,
287}
288
289impl<T> SparseColumnRunningStatistics<T>
290where
291    T: Float + AddAssign + Sum + Zero,
292{
293    /// Create a new per-column running statistics accumulator.
294    ///
295    /// # Arguments
296    /// * `ncols` — total number of columns whose statistics will be tracked
297    /// * `nrows` — row denominator used for mean/variance (number of rows
298    ///   considered; for a row-filtered scan this should be the number of
299    ///   rows kept)
300    pub fn new(ncols: usize, nrows: usize) -> Self {
301        Self {
302            nrows,
303            npos: vec![T::zero(); ncols],
304            s1: vec![T::zero(); ncols],
305            s2: vec![T::zero(); ncols],
306        }
307    }
308
309    /// Accumulate statistics from a CSC slab. `col_offset` is the global
310    /// column index of the slab's local column 0.
311    pub fn add_csc(&mut self, csc: &CscMatrix<T>, col_offset: usize) {
312        self.add_csc_inner(csc, col_offset, None);
313    }
314
315    /// Accumulate statistics from a CSC slab, keeping only rows where
316    /// `row_mask[row]` is `true`. `row_mask` must cover every row index
317    /// appearing in the CSC (length ≥ csc.nrows()).
318    pub fn add_csc_masked(&mut self, csc: &CscMatrix<T>, col_offset: usize, row_mask: &[bool]) {
319        debug_assert!(row_mask.len() >= csc.nrows());
320        self.add_csc_inner(csc, col_offset, Some(row_mask));
321    }
322
323    fn add_csc_inner(&mut self, csc: &CscMatrix<T>, col_offset: usize, row_mask: Option<&[bool]>) {
324        for (local_col, col) in csc.col_iter().enumerate() {
325            let c = col_offset + local_col;
326            let rows = col.row_indices();
327            let vals = col.values();
328            for (&row, &val) in rows.iter().zip(vals.iter()) {
329                if let Some(mask) = row_mask {
330                    if !mask[row] {
331                        continue;
332                    }
333                }
334                if !val.is_finite() {
335                    continue;
336                }
337                if val > T::zero() {
338                    self.npos[c] += T::one();
339                }
340                self.s1[c] += val;
341                self.s2[c] += val * val;
342            }
343        }
344    }
345
346    pub fn ncols(&self) -> usize {
347        self.npos.len()
348    }
349
350    pub fn nrows(&self) -> usize {
351        self.nrows
352    }
353
354    fn denom(&self) -> T {
355        safe_denom::<T>(self.nrows)
356    }
357}
358
359impl<T> SparseColumnRunningStatistics<T>
360where
361    T: Float + AddAssign + Sum + Zero + Display + ToPrimitive,
362{
363    /// Save the per-column statistics to a file (parquet if `filename`
364    /// ends with `.parquet`, otherwise a separator-delimited text file).
365    pub fn save(&self, filename: &str, names: &[Box<str>], sep: &str) -> anyhow::Result<()> {
366        let (nnz, tot, mu, sig) = self.to_f32_vecs();
367        write_stat_file(
368            filename,
369            names,
370            sep,
371            StatColumns {
372                nnz: &nnz,
373                tot: &tot,
374                mu: &mu,
375                sig: &sig,
376            },
377        )
378    }
379
380    /// Get statistics as vectors (nnz, tot, mu, sig) converted to f32
381    pub fn to_f32_vecs(&self) -> (Vec<f32>, Vec<f32>, Vec<f32>, Vec<f32>) {
382        let to_f32_slice =
383            |v: &[T]| -> Vec<f32> { v.iter().map(|x| x.to_f32().unwrap_or(0.0)).collect() };
384        let nnz = to_f32_slice(&self.npos);
385        let tot = to_f32_slice(&self.s1);
386        let mu = to_f32_slice(&self.mean());
387        let sig = to_f32_slice(&self.std());
388        (nnz, tot, mu, sig)
389    }
390}
391
392impl<T> RunningStatOps<T> for SparseColumnRunningStatistics<T>
393where
394    T: Float + AddAssign + Sum + Zero,
395{
396    type Output = Vec<T>;
397
398    fn clear(&mut self) {
399        self.npos.fill(T::zero());
400        self.s1.fill(T::zero());
401        self.s2.fill(T::zero());
402    }
403
404    fn count_positives(&self) -> Vec<T> {
405        self.npos.clone()
406    }
407
408    fn sum(&self) -> Vec<T> {
409        self.s1.clone()
410    }
411
412    fn mean(&self) -> Vec<T> {
413        let n = self.denom();
414        self.s1.iter().map(|&s| s / n).collect()
415    }
416
417    fn variance(&self) -> Vec<T> {
418        let n = self.denom();
419        self.s1
420            .iter()
421            .zip(self.s2.iter())
422            .map(|(&s1, &s2)| {
423                let mu = s1 / n;
424                s2 / n - mu * mu
425            })
426            .collect()
427    }
428
429    fn std(&self) -> Vec<T> {
430        self.variance().into_iter().map(|v| v.sqrt()).collect()
431    }
432}
433
434/// Four-column stat table (nnz, tot, mu, sig) to be written by
435/// [`write_stat_file`]. All slices must have the same length.
436struct StatColumns<'a> {
437    nnz: &'a [f32],
438    tot: &'a [f32],
439    mu: &'a [f32],
440    sig: &'a [f32],
441}
442
443impl StatColumns<'_> {
444    fn len(&self) -> usize {
445        self.nnz.len()
446    }
447}
448
449/// Shared parquet/TSV writer for per-entity stat tables with the
450/// standard (name, nnz, tot, mu, sig) schema.
451fn write_stat_file(
452    filename: &str,
453    names: &[Box<str>],
454    sep: &str,
455    stats: StatColumns<'_>,
456) -> anyhow::Result<()> {
457    let n = stats.len();
458    if names.len() != n {
459        anyhow::bail!(
460            "The number of names ({}) does not match stat length ({})",
461            names.len(),
462            n
463        );
464    }
465
466    match file_ext(filename).unwrap_or(Box::from("")).as_ref() {
467        "parquet" => {
468            let column_names: Vec<Box<str>> =
469                STAT_COLUMN_NAMES.iter().map(|s| (*s).into()).collect();
470            let column_types = vec![ParquetType::FLOAT; STAT_COLUMN_NAMES.len()];
471
472            let parquet_writer = ParquetWriter::new(
473                filename,
474                (n, 4),
475                (Some(names), Some(&column_names)),
476                Some(&column_types),
477                None,
478            )?;
479
480            let mut writer = parquet_writer.get_writer()?;
481            let mut row_group_writer = writer.next_row_group()?;
482
483            parquet_add_string_column(&mut row_group_writer, names)?;
484            parquet_add_numeric_column(&mut row_group_writer, stats.nnz)?;
485            parquet_add_numeric_column(&mut row_group_writer, stats.tot)?;
486            parquet_add_numeric_column(&mut row_group_writer, stats.mu)?;
487            parquet_add_numeric_column(&mut row_group_writer, stats.sig)?;
488
489            row_group_writer.close()?;
490            writer.close()?;
491        }
492        _ => {
493            let mut out: Vec<Box<str>> = (0..n)
494                .map(|i| {
495                    format!(
496                        "{}{}{}{}{}{}{}{}{}",
497                        names[i],
498                        sep,
499                        format_value(stats.nnz[i]),
500                        sep,
501                        format_value(stats.tot[i]),
502                        sep,
503                        format_value(stats.mu[i]),
504                        sep,
505                        format_value(stats.sig[i])
506                    )
507                    .into_boxed_str()
508                })
509                .collect();
510            let header = format!("#name{}nnz{}tot{}mu{}sig", sep, sep, sep, sep);
511            out.insert(0, header.into_boxed_str());
512            write_lines(&out, filename)?;
513        }
514    }
515    Ok(())
516}
517
518/// Save multiple group statistics to a single parquet file with a group column
519pub fn save_grouped_stats_parquet(
520    filename: &str,
521    names: &[Box<str>],
522    group_names: &[Box<str>],
523    group_stats: &[SparseRunningStatistics<f32>],
524) -> anyhow::Result<()> {
525    save_grouped_stats_parquet_cols(filename, &[("name", names)], group_names, group_stats)
526}
527
528/// Like [`save_grouped_stats_parquet`] but with arbitrary leading **string key
529/// columns** instead of the single `name` — e.g. split a `gene/modality/detail`
530/// feature name into `gene` / `modality` / `component` columns. Each
531/// `key_cols[j].1` is a per-feature vector aligned to the stats' row order
532/// (length must equal the feature-row count); the output is long format, one row
533/// per (feature, group): `<key cols…>, group, nnz, tot, mu, sig`.
534pub fn save_grouped_stats_parquet_cols(
535    filename: &str,
536    key_cols: &[(&str, &[Box<str>])],
537    group_names: &[Box<str>],
538    group_stats: &[SparseRunningStatistics<f32>],
539) -> anyhow::Result<()> {
540    use crate::matrix::parquet::{write_named_table, Column};
541
542    if group_names.len() != group_stats.len() {
543        anyhow::bail!(
544            "Number of group names ({}) does not match number of group stats ({})",
545            group_names.len(),
546            group_stats.len()
547        );
548    }
549    anyhow::ensure!(
550        !key_cols.is_empty(),
551        "save_grouped_stats_parquet_cols: need at least one key column"
552    );
553    let n_features = group_stats.first().map_or(0, |s| s.nrows());
554    for &(name, vals) in key_cols {
555        anyhow::ensure!(
556            vals.len() == n_features,
557            "key column '{name}' has {} entries but there are {n_features} feature rows",
558            vals.len(),
559        );
560    }
561
562    // Expand to long format (one row per (feature, group)): each key column is its
563    // per-feature vector repeated for every group; `group` + stats vary per row.
564    // (`vec![Vec::with_capacity(..); n]` would clone the empty Vec, losing the
565    // reservation on every column but the first — reserve per column instead.)
566    let total_rows = n_features * group_names.len();
567    let mut keys: Vec<Vec<Box<str>>> = (0..key_cols.len())
568        .map(|_| Vec::with_capacity(total_rows))
569        .collect();
570    let mut all_groups: Vec<Box<str>> = Vec::with_capacity(total_rows);
571    let mut all_nnz: Vec<f32> = Vec::with_capacity(total_rows);
572    let mut all_tot: Vec<f32> = Vec::with_capacity(total_rows);
573    let mut all_mu: Vec<f32> = Vec::with_capacity(total_rows);
574    let mut all_sig: Vec<f32> = Vec::with_capacity(total_rows);
575    for (group_name, stat) in group_names.iter().zip(group_stats.iter()) {
576        let (nnz, tot, mu, sig) = stat.to_f32_vecs();
577        // Append each column for this group in bulk (one block of `n_features`).
578        for (j, &(_, vals)) in key_cols.iter().enumerate() {
579            keys[j].extend(vals.iter().cloned());
580        }
581        all_groups.resize(all_groups.len() + n_features, group_name.clone());
582        all_nnz.extend_from_slice(&nnz);
583        all_tot.extend_from_slice(&tot);
584        all_mu.extend_from_slice(&mu);
585        all_sig.extend_from_slice(&sig);
586    }
587
588    // First key column is the leading row column; remaining key columns, then
589    // `group` and the four stats, follow. Delegate schema + writing to the shared
590    // tidy-table writer instead of hand-rolling it.
591    let mut columns: Vec<(Box<str>, Column)> = Vec::with_capacity(key_cols.len() + 4);
592    for (&(name, _), col) in key_cols.iter().zip(keys.iter()).skip(1) {
593        columns.push((name.into(), Column::Str(col.as_slice())));
594    }
595    columns.push(("group".into(), Column::Str(all_groups.as_slice())));
596    columns.push(("nnz".into(), Column::F32(all_nnz.as_slice())));
597    columns.push(("tot".into(), Column::F32(all_tot.as_slice())));
598    columns.push(("mu".into(), Column::F32(all_mu.as_slice())));
599    columns.push(("sig".into(), Column::F32(all_sig.as_slice())));
600
601    write_named_table(filename, key_cols[0].0, &keys[0], &columns)
602}
603
604fn format_value<T: Float + Display>(v: T) -> String {
605    let v_f64 = v.to_f64().unwrap_or(0.0);
606    if v_f64.abs() > 1e-4 {
607        format!("{:.4}", v_f64)
608            .trim_end_matches('0')
609            .trim_end_matches('.')
610            .to_string()
611    } else if v_f64.abs() > 1e-20 {
612        format!("{:.4e}", v_f64)
613    } else {
614        "0".to_string()
615    }
616}
617
618impl<T> RunningStatOps<T> for SparseRunningStatistics<T>
619where
620    T: Float + AddAssign + Sum + Zero,
621{
622    type Output = Vec<T>;
623
624    fn clear(&mut self) {
625        self.ncols_processed = 0;
626        self.npos.fill(T::zero());
627        self.s1.fill(T::zero());
628        self.s2.fill(T::zero());
629    }
630
631    /// Count of positive (non-zero) values per row
632    fn count_positives(&self) -> Vec<T> {
633        self.npos.clone()
634    }
635
636    /// Sum per row
637    fn sum(&self) -> Vec<T> {
638        self.s1.clone()
639    }
640
641    /// Mean per row
642    /// Uses ncols_processed as the denominator (implicit zeros count)
643    fn mean(&self) -> Vec<T> {
644        let n = self.denom();
645        self.s1.iter().map(|&s| s / n).collect()
646    }
647
648    /// Variance per row
649    fn variance(&self) -> Vec<T> {
650        let n = self.denom();
651        self.s1
652            .iter()
653            .zip(self.s2.iter())
654            .map(|(&s1, &s2)| {
655                let mu = s1 / n;
656                s2 / n - mu * mu
657            })
658            .collect()
659    }
660
661    /// Standard deviation per row
662    fn std(&self) -> Vec<T> {
663        self.variance().into_iter().map(|v| v.sqrt()).collect()
664    }
665}
666
667#[cfg(test)]
668mod tests {
669    use super::*;
670
671    #[test]
672    fn test_sparse_running_stat_basic() {
673        let mut stat = SparseRunningStatistics::<f32>::new(4);
674
675        // Column 0: [1, 0, 2, 0]
676        stat.add_sparse_column(&[0, 2], &[1.0, 2.0]);
677
678        // Column 1: [0, 3, 0, 4]
679        stat.add_sparse_column(&[1, 3], &[3.0, 4.0]);
680
681        assert_eq!(stat.ncols_processed(), 2);
682
683        // npos: [1, 1, 1, 1]
684        assert_eq!(stat.count_positives(), vec![1.0, 1.0, 1.0, 1.0]);
685
686        // sum: [1, 3, 2, 4]
687        assert_eq!(stat.sum(), vec![1.0, 3.0, 2.0, 4.0]);
688
689        // mean: [0.5, 1.5, 1.0, 2.0]
690        let mean = stat.mean();
691        assert!((mean[0] - 0.5).abs() < 1e-6);
692        assert!((mean[1] - 1.5).abs() < 1e-6);
693        assert!((mean[2] - 1.0).abs() < 1e-6);
694        assert!((mean[3] - 2.0).abs() < 1e-6);
695    }
696
697    #[test]
698    fn test_sparse_running_stat_csc() {
699        use nalgebra_sparse::CooMatrix;
700
701        let mut stat = SparseRunningStatistics::<f32>::new(3);
702
703        // Create a 3x2 sparse matrix:
704        // [1, 0]
705        // [0, 2]
706        // [3, 0]
707        let mut coo: CooMatrix<f32> = CooMatrix::new(3, 2);
708        coo.push(0, 0, 1.0);
709        coo.push(1, 1, 2.0);
710        coo.push(2, 0, 3.0);
711        let csc = CscMatrix::from(&coo);
712
713        stat.add_csc(&csc);
714
715        assert_eq!(stat.ncols_processed(), 2);
716        assert_eq!(stat.count_positives(), vec![1.0, 1.0, 1.0]);
717        assert_eq!(stat.sum(), vec![1.0, 2.0, 3.0]);
718    }
719
720    #[test]
721    fn test_sparse_running_stat_f64() {
722        let mut stat = SparseRunningStatistics::<f64>::new(2);
723
724        stat.add_sparse_column(&[0, 1], &[1.0, 2.0]);
725        stat.add_sparse_column(&[0], &[3.0]);
726
727        assert_eq!(stat.ncols_processed(), 2);
728        assert_eq!(stat.sum(), vec![4.0, 2.0]);
729
730        let mean = stat.mean();
731        assert!((mean[0] - 2.0).abs() < 1e-10);
732        assert!((mean[1] - 1.0).abs() < 1e-10);
733    }
734
735    #[test]
736    fn test_sparse_column_running_stat_csc() {
737        use nalgebra_sparse::CooMatrix;
738
739        // 3 rows x 4 columns:
740        // col0 = [1, 0, 3]
741        // col1 = [0, 2, 0]
742        // col2 = [0, 0, 0]
743        // col3 = [4, 5, 0]
744        let mut coo: CooMatrix<f32> = CooMatrix::new(3, 4);
745        coo.push(0, 0, 1.0);
746        coo.push(2, 0, 3.0);
747        coo.push(1, 1, 2.0);
748        coo.push(0, 3, 4.0);
749        coo.push(1, 3, 5.0);
750        let csc = CscMatrix::from(&coo);
751
752        let mut stat = SparseColumnRunningStatistics::<f32>::new(4, 3);
753        stat.add_csc(&csc, 0);
754
755        assert_eq!(stat.count_positives(), vec![2.0, 1.0, 0.0, 2.0]);
756        assert_eq!(stat.sum(), vec![4.0, 2.0, 0.0, 9.0]);
757
758        // mean = sum / nrows (denominator = 3)
759        let mean = stat.mean();
760        assert!((mean[0] - 4.0 / 3.0).abs() < 1e-6);
761        assert!((mean[1] - 2.0 / 3.0).abs() < 1e-6);
762        assert!((mean[2] - 0.0).abs() < 1e-6);
763        assert!((mean[3] - 9.0 / 3.0).abs() < 1e-6);
764    }
765
766    #[test]
767    fn test_sparse_column_running_stat_block_offset() {
768        use nalgebra_sparse::CooMatrix;
769
770        // Simulate processing columns in two blocks of a 3x4 matrix:
771        // Block A (cols 0..2): col0=[1,0,3], col1=[0,2,0]
772        // Block B (cols 2..4): col2=[0,0,0], col3=[4,5,0]
773        let mut coo_a: CooMatrix<f32> = CooMatrix::new(3, 2);
774        coo_a.push(0, 0, 1.0);
775        coo_a.push(2, 0, 3.0);
776        coo_a.push(1, 1, 2.0);
777        let csc_a = CscMatrix::from(&coo_a);
778
779        let mut coo_b: CooMatrix<f32> = CooMatrix::new(3, 2);
780        coo_b.push(0, 1, 4.0);
781        coo_b.push(1, 1, 5.0);
782        let csc_b = CscMatrix::from(&coo_b);
783
784        let mut stat = SparseColumnRunningStatistics::<f32>::new(4, 3);
785        stat.add_csc(&csc_a, 0);
786        stat.add_csc(&csc_b, 2);
787
788        assert_eq!(stat.count_positives(), vec![2.0, 1.0, 0.0, 2.0]);
789        assert_eq!(stat.sum(), vec![4.0, 2.0, 0.0, 9.0]);
790    }
791
792    #[test]
793    fn test_sparse_column_running_stat_masked() {
794        use nalgebra_sparse::CooMatrix;
795
796        // 4 rows x 2 columns, keep only rows [0, 2]
797        // col0 = [1, 2, 3, 4] → kept: 1+3 = 4, nnz=2
798        // col1 = [0, 5, 0, 6] → kept: 0, nnz=0
799        let mut coo: CooMatrix<f32> = CooMatrix::new(4, 2);
800        coo.push(0, 0, 1.0);
801        coo.push(1, 0, 2.0);
802        coo.push(2, 0, 3.0);
803        coo.push(3, 0, 4.0);
804        coo.push(1, 1, 5.0);
805        coo.push(3, 1, 6.0);
806        let csc = CscMatrix::from(&coo);
807
808        let row_mask = vec![true, false, true, false];
809
810        let mut stat = SparseColumnRunningStatistics::<f32>::new(2, 2);
811        stat.add_csc_masked(&csc, 0, &row_mask);
812
813        assert_eq!(stat.count_positives(), vec![2.0, 0.0]);
814        assert_eq!(stat.sum(), vec![4.0, 0.0]);
815    }
816}