use crate::sparse_data_visitors::styled_progress_bar;
use crate::sparse_io_vector::SparseIoVec;
use indicatif::ParallelProgressIterator;
use legume_numeric::matrix::knn_graph::KnnGraph;
use legume_numeric::matrix::parquet::{write_named_table, Column};
use legume_numeric::matrix::utils::generate_minibatch_intervals;
use rayon::prelude::*;
use rustc_hash::FxHashMap as HashMap;
use std::sync::{Arc, Mutex};
#[derive(Clone, Debug, Default)]
pub struct PairSet {
pub pairs: Vec<(usize, usize)>,
pub weights: Option<Vec<f32>>,
}
impl PairSet {
pub fn new(pairs: Vec<(usize, usize)>) -> Self {
Self {
pairs,
weights: None,
}
}
pub fn len(&self) -> usize {
self.pairs.len()
}
pub fn is_empty(&self) -> bool {
self.pairs.is_empty()
}
pub fn collapse_by(&self, node_labels: &[usize]) -> (PairSet, Vec<usize>) {
collapse_pairs(&self.pairs, node_labels)
}
}
pub fn collapse_pairs(pairs: &[(usize, usize)], node_labels: &[usize]) -> (PairSet, Vec<usize>) {
let mut key_to_coarse: HashMap<(usize, usize), usize> = HashMap::default();
let mut coarse: Vec<(usize, usize)> = Vec::new();
let mut fine_to_coarse = Vec::with_capacity(pairs.len());
for &(i, j) in pairs {
let (li, lj) = (node_labels[i], node_labels[j]);
let key = (li.min(lj), li.max(lj));
let c = *key_to_coarse.entry(key).or_insert_with(|| {
let next = coarse.len();
coarse.push(key);
next
});
fine_to_coarse.push(c);
}
(PairSet::new(coarse), fine_to_coarse)
}
pub struct CellPairs<'a> {
pub data: &'a SparseIoVec,
pairs: &'a [(usize, usize)],
weights: Option<&'a [f32]>,
pub weight_column: &'static str,
}
impl<'a> CellPairs<'a> {
pub fn new(data: &'a SparseIoVec, pairs: &'a [(usize, usize)]) -> Self {
Self {
data,
pairs,
weights: None,
weight_column: "distance",
}
}
pub fn with_weights(
data: &'a SparseIoVec,
pairs: &'a [(usize, usize)],
weights: &'a [f32],
) -> anyhow::Result<Self> {
anyhow::ensure!(
weights.len() == pairs.len(),
"{} weights for {} pairs",
weights.len(),
pairs.len()
);
Ok(Self {
data,
pairs,
weights: Some(weights),
weight_column: "distance",
})
}
pub fn from_graph(data: &'a SparseIoVec, graph: &'a KnnGraph) -> Self {
Self {
data,
pairs: &graph.edges,
weights: Some(&graph.distances),
weight_column: "distance",
}
}
pub fn pairs(&self) -> &[(usize, usize)] {
self.pairs
}
pub fn weights(&self) -> Option<&[f32]> {
self.weights
}
pub fn num_pairs(&self) -> usize {
self.pairs.len()
}
pub fn num_features(&self) -> usize {
self.data.num_rows()
}
pub fn visit_pairs_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),
&CellPairs,
&SharedIn,
Arc<Mutex<&mut SharedOut>>,
) -> anyhow::Result<()>
+ Sync
+ Send,
SharedIn: Sync + Send + ?Sized,
SharedOut: Sync + Send,
{
let ntot = self.num_pairs();
let jobs = generate_minibatch_intervals(ntot, self.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)| -> anyhow::Result<()> {
visitor((lb, ub), self, shared_in, arc_shared_out.clone())
})
.collect::<anyhow::Result<()>>();
prog_bar.finish_and_clear();
result
}
pub fn to_parquet(
&self,
file_path: &str,
extra: &[(Box<str>, Column<'_>)],
) -> anyhow::Result<()> {
let num_pairs = self.num_pairs();
let cell_names = self.data.column_names()?;
let left: Vec<Box<str>> = self
.pairs()
.iter()
.map(|&(left, _)| cell_names[left].clone())
.collect();
let right: Vec<Box<str>> = self
.pairs()
.iter()
.map(|&(_, right)| cell_names[right].clone())
.collect();
let mut columns: Vec<(Box<str>, Column<'_>)> = vec![
("left_cell".into(), Column::Str(&left)),
("right_cell".into(), Column::Str(&right)),
];
for (name, col) in extra {
let len = col.len();
if len != num_pairs {
return Err(anyhow::anyhow!(
"column `{}` carries {} values for {} pairs",
name,
len,
num_pairs
));
}
columns.push((
name.clone(),
match col {
Column::Str(d) => Column::Str(d),
Column::F32(d) => Column::F32(d),
Column::I32(d) => Column::I32(d),
Column::I64(d) => Column::I64(d),
},
));
}
if let Some(w) = self.weights {
columns.push((self.weight_column.into(), Column::F32(w)));
}
let row_names: Vec<Box<str>> = (0..num_pairs)
.map(|i| i.to_string().into_boxed_str())
.collect();
write_named_table(file_path, "cell_pair", &row_names, &columns)
}
}