Skip to main content

data_beans/
sparse_data_visitors.rs

1#![allow(dead_code)]
2
3use crate::sparse_io_vector::SparseIoVec;
4use indicatif::ParallelProgressIterator;
5use legume_numeric::matrix::utils::generate_minibatch_intervals;
6use rayon::prelude::*;
7use std::sync::{Arc, Mutex};
8
9pub trait VisitColumnsOps {
10    /// visit all the columns by sequential blocks.  The visitor
11    /// function should take (a) `(lb, ub)` (b) `&Self` (c)
12    /// `&SharedIn` (d) `Arc::new(Mutex::new(&mut SharedOut)`
13    fn visit_columns_by_block<Visitor, SharedIn, SharedOut>(
14        &self,
15        visitor: &Visitor,
16        shared_in: &SharedIn,
17        shared_out: &mut SharedOut,
18        block_size: Option<usize>,
19    ) -> anyhow::Result<()>
20    where
21        Visitor: Fn((usize, usize), &Self, &SharedIn, Arc<Mutex<&mut SharedOut>>) -> anyhow::Result<()>
22            + Sync
23            + Send,
24        SharedIn: Sync + Send + ?Sized,
25        SharedOut: Sync + Send;
26
27    /// visit all the columns by predefined groups assigned by
28    /// `self.assign_groups`. The visitor function should take (a)
29    /// `group_index` (b) `&[columns_in_the_group]` (c) `&Self` (d)
30    /// `&SharedIn` (e) `Arc::new(Mutex::new(&mut SharedOut)`
31    fn visit_columns_by_group<Visitor, SharedIn, SharedOut>(
32        &self,
33        visitor: &Visitor,
34        shared_in: &SharedIn,
35        shared_out: &mut SharedOut,
36    ) -> anyhow::Result<()>
37    where
38        Visitor: Fn(usize, &[usize], &Self, &SharedIn, Arc<Mutex<&mut SharedOut>>) -> anyhow::Result<()>
39            + Sync
40            + Send,
41        SharedIn: Sync + Send + ?Sized,
42        SharedOut: Sync + Send;
43}
44
45/// Shared bar used by every par-over-column visitor so progress shows the
46/// workspace-standard `[elapsed] bar pos/len (eta) <unit_label>` style across
47/// random projection, collapsing, and downstream column scans. Delegates to
48/// the canonical [`legume_numeric::matrix::progress::new_progress_bar`] (single style,
49/// single `MULTI_PROGRESS`) and keeps the 500 ms steady tick so the bar
50/// animates even when a single block is slow. Exported so downstream crates
51/// (data_beans::alg, senna) can reuse the same look.
52pub fn styled_progress_bar(total: u64, unit_label: &str) -> indicatif::ProgressBar {
53    let prog_bar = legume_numeric::matrix::progress::new_progress_bar(total)
54        .with_message(unit_label.to_string());
55    prog_bar.enable_steady_tick(std::time::Duration::from_millis(500));
56    prog_bar
57}
58
59impl VisitColumnsOps for SparseIoVec {
60    fn visit_columns_by_block<Visitor, SharedIn, SharedOut>(
61        &self,
62        visitor: &Visitor,
63        shared_in: &SharedIn,
64        shared_out: &mut SharedOut,
65        block_size: Option<usize>,
66    ) -> anyhow::Result<()>
67    where
68        Visitor: Fn((usize, usize), &Self, &SharedIn, Arc<Mutex<&mut SharedOut>>) -> anyhow::Result<()>
69            + Sync
70            + Send,
71        SharedIn: Sync + Send + ?Sized,
72        SharedOut: Sync + Send,
73    {
74        let ntot = self.num_columns();
75        let num_features = self.num_rows();
76        let jobs = create_jobs(ntot, num_features, block_size);
77
78        let arc_shared_out = Arc::new(Mutex::new(shared_out));
79        let prog_bar = styled_progress_bar(jobs.len() as u64, "blocks");
80
81        let result = jobs
82            .par_iter()
83            .progress_with(prog_bar.clone())
84            .map(|&(lb, ub)| visitor((lb, ub), self, shared_in, arc_shared_out.clone()))
85            .collect();
86        prog_bar.finish_and_clear();
87        result
88    }
89
90    fn visit_columns_by_group<Visitor, SharedIn, SharedOut>(
91        &self,
92        visitor: &Visitor,
93        shared_in: &SharedIn,
94        shared_out: &mut SharedOut,
95    ) -> anyhow::Result<()>
96    where
97        Visitor: Fn(usize, &[usize], &Self, &SharedIn, Arc<Mutex<&mut SharedOut>>) -> anyhow::Result<()>
98            + Sync
99            + Send,
100        SharedIn: Sync + Send + ?Sized,
101        SharedOut: Sync + Send,
102    {
103        let group_to_cols = self.take_grouped_columns().ok_or(anyhow::anyhow!(
104            "The columns were not assigned before. Call `assign_groups`"
105        ))?;
106
107        let arc_shared_out = Arc::new(Mutex::new(shared_out));
108        let num_samples = group_to_cols.len();
109        let prog_bar = styled_progress_bar(num_samples as u64, "groups");
110
111        let result = group_to_cols
112            .par_iter()
113            .enumerate()
114            .progress_with(prog_bar.clone())
115            .map(|(sample, cells)| visitor(sample, cells, self, shared_in, arc_shared_out.clone()))
116            .collect();
117        prog_bar.finish_and_clear();
118        result
119    }
120}
121
122/// Thin wrapper around [`generate_minibatch_intervals`] kept for in-crate
123/// call sites that predate the legume_numeric::matrix split.
124pub fn create_jobs(
125    ntot: usize,
126    num_features: usize,
127    block_size: Option<usize>,
128) -> Vec<(usize, usize)> {
129    generate_minibatch_intervals(ntot, num_features, block_size)
130}