#![allow(dead_code)]
use crate::sparse_io_vector::SparseIoVec;
use indicatif::ParallelProgressIterator;
use legume_numeric::matrix::utils::generate_minibatch_intervals;
use rayon::prelude::*;
use std::sync::{Arc, Mutex};
pub trait VisitColumnsOps {
fn visit_columns_by_block<Visitor, SharedIn, SharedOut>(
&self,
visitor: &Visitor,
shared_in: &SharedIn,
shared_out: &mut SharedOut,
block_size: Option<usize>,
) -> anyhow::Result<()>
where
Visitor: Fn((usize, usize), &Self, &SharedIn, Arc<Mutex<&mut SharedOut>>) -> anyhow::Result<()>
+ Sync
+ Send,
SharedIn: Sync + Send + ?Sized,
SharedOut: Sync + Send;
fn visit_columns_by_group<Visitor, SharedIn, SharedOut>(
&self,
visitor: &Visitor,
shared_in: &SharedIn,
shared_out: &mut SharedOut,
) -> anyhow::Result<()>
where
Visitor: Fn(usize, &[usize], &Self, &SharedIn, Arc<Mutex<&mut SharedOut>>) -> anyhow::Result<()>
+ Sync
+ Send,
SharedIn: Sync + Send + ?Sized,
SharedOut: Sync + Send;
}
pub fn styled_progress_bar(total: u64, unit_label: &str) -> indicatif::ProgressBar {
let prog_bar = legume_numeric::matrix::progress::new_progress_bar(total)
.with_message(unit_label.to_string());
prog_bar.enable_steady_tick(std::time::Duration::from_millis(500));
prog_bar
}
impl VisitColumnsOps for SparseIoVec {
fn visit_columns_by_block<Visitor, SharedIn, SharedOut>(
&self,
visitor: &Visitor,
shared_in: &SharedIn,
shared_out: &mut SharedOut,
block_size: Option<usize>,
) -> anyhow::Result<()>
where
Visitor: Fn((usize, usize), &Self, &SharedIn, Arc<Mutex<&mut SharedOut>>) -> anyhow::Result<()>
+ Sync
+ Send,
SharedIn: Sync + Send + ?Sized,
SharedOut: Sync + Send,
{
let ntot = self.num_columns();
let num_features = self.num_rows();
let jobs = create_jobs(ntot, num_features, block_size);
let arc_shared_out = Arc::new(Mutex::new(shared_out));
let prog_bar = styled_progress_bar(jobs.len() as u64, "blocks");
let result = jobs
.par_iter()
.progress_with(prog_bar.clone())
.map(|&(lb, ub)| visitor((lb, ub), self, shared_in, arc_shared_out.clone()))
.collect();
prog_bar.finish_and_clear();
result
}
fn visit_columns_by_group<Visitor, SharedIn, SharedOut>(
&self,
visitor: &Visitor,
shared_in: &SharedIn,
shared_out: &mut SharedOut,
) -> anyhow::Result<()>
where
Visitor: Fn(usize, &[usize], &Self, &SharedIn, Arc<Mutex<&mut SharedOut>>) -> anyhow::Result<()>
+ Sync
+ Send,
SharedIn: Sync + Send + ?Sized,
SharedOut: Sync + Send,
{
let group_to_cols = self.take_grouped_columns().ok_or(anyhow::anyhow!(
"The columns were not assigned before. Call `assign_groups`"
))?;
let arc_shared_out = Arc::new(Mutex::new(shared_out));
let num_samples = group_to_cols.len();
let prog_bar = styled_progress_bar(num_samples as u64, "groups");
let result = group_to_cols
.par_iter()
.enumerate()
.progress_with(prog_bar.clone())
.map(|(sample, cells)| visitor(sample, cells, self, shared_in, arc_shared_out.clone()))
.collect();
prog_bar.finish_and_clear();
result
}
}
pub fn create_jobs(
ntot: usize,
num_features: usize,
block_size: Option<usize>,
) -> Vec<(usize, usize)> {
generate_minibatch_intervals(ntot, num_features, block_size)
}