data_beans/
sparse_data_visitors.rs1#![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 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 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
45pub 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
122pub 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}