use super::*;
pub struct PbSampleLayout {
pub centroids: DMatrix<f32>,
pub cell_counts: Vec<f32>,
pub pb_sample_to_batch: Vec<usize>,
pub pb_sample_to_group: Vec<usize>,
pub bg_to_pbsamp: HashMap<(usize, usize), usize>,
pub cell_to_pbsamp: Vec<usize>,
pub singleton_col: Vec<Option<usize>>,
pub bulk_batches: Vec<usize>,
pub pb_sample_to_stratum: Option<Vec<usize>>,
}
impl PbSampleLayout {
#[must_use]
pub fn is_bulk(&self, pbsamp: usize) -> bool {
self.is_bulk_batch(self.pb_sample_to_batch[pbsamp])
}
#[must_use]
pub fn is_bulk_batch(&self, batch: usize) -> bool {
self.bulk_batches.contains(&batch)
}
}
pub struct PbSampleCollection {
pub layout: PbSampleLayout,
pub gene_sums: Vec<Vec<(usize, f32)>>,
pub num_genes: usize,
}
pub(super) struct BatchAccumulator {
centroid_sum: Vec<f32>,
gene_sum: HashMap<usize, f32>,
count: usize,
}
pub(super) struct PbSampleData {
centroid: Vec<f32>,
gene_sums: Vec<(usize, f32)>,
cell_count: f32,
batch: usize,
group: usize,
}
pub(super) fn build_pb_sample_layout(
group_to_cols: &[Vec<usize>],
col_to_batch: &[usize],
proj_kn: &DMatrix<f32>,
col_weight: Option<&[f32]>,
anchor_batches: &[usize],
bulk_batches: &[usize],
cell_to_stratum: Option<&[usize]>,
) -> anyhow::Result<PbSampleLayout> {
let proj_dim = proj_kn.nrows();
struct CentroidAccum {
centroid_sum: Vec<f32>,
count: f32,
}
let weight_of = |c: usize| col_weight.map_or(1.0, |w| w[c]);
type CentroidTuple = (usize, usize, Vec<f32>, f32, Option<usize>);
let per_group_results: Vec<Vec<CentroidTuple>> = group_to_cols
.par_iter()
.enumerate()
.map(|(group, cells)| {
let mut batch_data: HashMap<usize, CentroidAccum> = HashMap::default();
let mut singletons: Vec<CentroidTuple> = Vec::new();
for &glob_idx in cells {
let batch = col_to_batch[glob_idx];
let w = weight_of(glob_idx);
if anchor_batches.contains(&batch) || bulk_batches.contains(&batch) {
let centroid: Vec<f32> =
(0..proj_dim).map(|d| proj_kn[(d, glob_idx)]).collect();
singletons.push((batch, group, centroid, w, Some(glob_idx)));
continue;
}
let acc = batch_data.entry(batch).or_insert_with(|| CentroidAccum {
centroid_sum: vec![0f32; proj_dim],
count: 0.0,
});
for d in 0..proj_dim {
acc.centroid_sum[d] += proj_kn[(d, glob_idx)] * w;
}
acc.count += w;
}
let mut out: Vec<CentroidTuple> = batch_data
.into_iter()
.filter(|(_, acc)| acc.count > 0.0)
.map(|(batch, acc)| {
let inv_count = 1.0 / acc.count;
let centroid: Vec<f32> =
acc.centroid_sum.iter().map(|v| v * inv_count).collect();
(batch, group, centroid, acc.count, None)
})
.collect::<Vec<_>>();
out.extend(singletons);
out
})
.collect();
let all_pbsamp: Vec<_> = per_group_results.into_iter().flatten().collect();
let num_pb = all_pbsamp.len();
if num_pb == 0 {
return Err(anyhow::anyhow!("no pb-samples built"));
}
let mut centroids = DMatrix::<f32>::zeros(proj_dim, num_pb);
let mut cell_counts = Vec::with_capacity(num_pb);
let mut pbsamp_to_batch = Vec::with_capacity(num_pb);
let mut pbsamp_to_group = Vec::with_capacity(num_pb);
let mut singleton_col = Vec::with_capacity(num_pb);
let mut bg_to_pbsamp = HashMap::default();
let ncols = col_to_batch.len();
let mut cell_to_pbsamp = vec![usize::MAX; ncols];
for (i, (batch, group, centroid, count, sc)) in all_pbsamp.into_iter().enumerate() {
for (d, &v) in centroid.iter().enumerate() {
centroids[(d, i)] = v;
}
cell_counts.push(count);
pbsamp_to_batch.push(batch);
pbsamp_to_group.push(group);
singleton_col.push(sc);
match sc {
Some(col) => cell_to_pbsamp[col] = i,
None => {
bg_to_pbsamp.insert((batch, group), i);
}
}
}
for (group, cells) in group_to_cols.iter().enumerate() {
for &c in cells {
let b = col_to_batch[c];
if cell_to_pbsamp[c] == usize::MAX {
if let Some(&pbsamp) = bg_to_pbsamp.get(&(b, group)) {
cell_to_pbsamp[c] = pbsamp;
}
}
}
}
let pb_sample_to_stratum = if let Some(strata) = cell_to_stratum {
anyhow::ensure!(
strata.len() == ncols,
"cell_to_stratum has {} entries, layout has {} columns",
strata.len(),
ncols
);
let mut out = vec![0usize; num_pb];
let mut seen = vec![false; num_pb];
for (c, &pbsamp) in cell_to_pbsamp.iter().enumerate() {
if pbsamp == usize::MAX {
continue;
}
if !seen[pbsamp] {
out[pbsamp] = strata[c];
seen[pbsamp] = true;
}
}
Some(out)
} else {
None
};
Ok(PbSampleLayout {
centroids,
cell_counts,
pb_sample_to_batch: pbsamp_to_batch,
pb_sample_to_group: pbsamp_to_group,
bg_to_pbsamp,
cell_to_pbsamp,
singleton_col,
bulk_batches: bulk_batches.to_vec(),
pb_sample_to_stratum,
})
}
pub(super) fn collect_pb_sample_gene_sums(
data_vec: &SparseIoVec,
group_to_cols: &[Vec<usize>],
cell_to_pbsamp: &[usize],
num_pb: usize,
) -> anyhow::Result<Vec<Vec<(usize, f32)>>> {
use indicatif::ParallelProgressIterator;
let prog_bar = styled_progress_bar(group_to_cols.len() as u64, "groups (pb-sample gene sums)");
let gene_sum_maps: Vec<(usize, HashMap<usize, f32>)> = group_to_cols
.par_iter()
.progress_with(prog_bar.clone())
.flat_map(|cells| {
let yy = data_vec
.read_columns_csc(cells.iter().cloned())
.expect("read_columns_csc");
let mut per_pbsamp: HashMap<usize, HashMap<usize, f32>> = HashMap::default();
for (local_idx, y_j) in yy.col_iter().enumerate() {
let col = cells[local_idx];
let pbsamp = cell_to_pbsamp[col];
if pbsamp == usize::MAX {
continue;
}
let w = data_vec.column_multiplicity(col);
let gene_map = per_pbsamp.entry(pbsamp).or_default();
for (&gene, &val) in y_j.row_indices().iter().zip(y_j.values().iter()) {
*gene_map.entry(gene).or_default() += val * w;
}
}
per_pbsamp.into_iter().collect::<Vec<_>>()
})
.collect();
let mut gene_sums: Vec<Vec<(usize, f32)>> = vec![vec![]; num_pb];
for (pbsamp_idx, gene_map) in gene_sum_maps {
let mut sorted: Vec<(usize, f32)> = gene_map.into_iter().collect();
sorted.sort_unstable_by_key(|&(g, _)| g);
gene_sums[pbsamp_idx] = sorted;
}
Ok(gene_sums)
}
pub(super) fn build_pb_samples(
data_vec: &SparseIoVec,
proj_kn: &DMatrix<f32>,
num_genes: usize,
anchor_batches: &[usize],
bulk_batches: &[usize],
cell_to_stratum: Option<&[usize]>,
) -> anyhow::Result<PbSampleCollection> {
let group_to_cols = data_vec
.take_grouped_columns()
.ok_or(anyhow::anyhow!("columns not assigned to groups"))?;
let col_to_batch: Vec<usize> = (0..proj_kn.ncols())
.map(|c| data_vec.get_batch_membership(std::iter::once(c))[0])
.collect();
let weights: Option<Vec<f32>> = data_vec.has_column_multiplicity().then(|| {
(0..proj_kn.ncols())
.map(|c| data_vec.column_multiplicity(c))
.collect()
});
let layout = build_pb_sample_layout(
group_to_cols,
&col_to_batch,
proj_kn,
weights.as_deref(),
anchor_batches,
bulk_batches,
cell_to_stratum,
)?;
let num_pb = layout.cell_counts.len();
let gene_sums =
collect_pb_sample_gene_sums(data_vec, group_to_cols, &layout.cell_to_pbsamp, num_pb)?;
Ok(PbSampleCollection {
layout,
gene_sums,
num_genes,
})
}
pub(crate) fn knn_distinct_pbsamples_in_batch(
bknn: &ColumnDict<usize>,
query: &[f32],
knn: usize,
cell_to_pbsamp: &[usize],
own_pbsamp: usize,
pb_sample_to_stratum: Option<&[usize]>,
own_stratum: Option<usize>,
) -> anyhow::Result<Vec<(usize, f32)>> {
let n = bknn.num_points();
if n == 0 || knn == 0 {
return Ok(Vec::new());
}
let mut query_k = (knn * 4 + 1).min(n);
let mut best: HashMap<usize, f32> = HashMap::default();
loop {
let (cell_ids, dists) = bknn.search_by_query_data(query, query_k)?;
best.clear();
for (&c, &d) in cell_ids.iter().zip(dists.iter()) {
let other_pbsamp = cell_to_pbsamp[c];
if other_pbsamp == usize::MAX || other_pbsamp == own_pbsamp {
continue;
}
if let (Some(st), Some(own_s)) = (pb_sample_to_stratum, own_stratum) {
if st[other_pbsamp] != own_s {
continue;
}
}
best.entry(other_pbsamp)
.and_modify(|old| {
if d < *old {
*old = d;
}
})
.or_insert(d);
}
if best.len() >= knn || query_k >= n {
break;
}
query_k = query_k.saturating_mul(4).min(n);
}
let mut per_batch: Vec<(usize, f32)> = best.drain().collect();
per_batch.sort_by(|a, b| a.1.partial_cmp(&b.1).unwrap_or(std::cmp::Ordering::Equal));
per_batch.truncate(knn);
Ok(per_batch)
}
pub(crate) fn bbknn_match_one_pbsamp(
layout: &PbSampleLayout,
batch_knn_lookup: &[ColumnDict<usize>],
knn: usize,
pbsamp: usize,
anchor_batches: Option<&[usize]>,
) -> anyhow::Result<Vec<(usize, f32)>> {
let pbsamp_batch = layout.pb_sample_to_batch[pbsamp];
let centroid: Vec<f32> = layout.centroids.column(pbsamp).iter().copied().collect();
let pb_stratum = layout.pb_sample_to_stratum.as_deref();
let own_stratum = pb_stratum.map(|s| s[pbsamp]);
let mut all_hits: Vec<(usize, f32)> = Vec::new();
match anchor_batches {
None => {
for (b, bknn) in batch_knn_lookup.iter().enumerate() {
if b == pbsamp_batch || layout.is_bulk_batch(b) {
continue;
}
let per_batch = knn_distinct_pbsamples_in_batch(
bknn,
¢roid,
knn,
&layout.cell_to_pbsamp,
pbsamp,
pb_stratum,
own_stratum,
)?;
all_hits.extend(per_batch);
}
}
Some(anchors) => {
for &b in anchors {
let Some(bknn) = batch_knn_lookup.get(b) else {
continue;
};
let per_batch = knn_distinct_pbsamples_in_batch(
bknn,
¢roid,
knn,
&layout.cell_to_pbsamp,
usize::MAX, pb_stratum,
own_stratum,
)?;
all_hits.extend(per_batch);
}
}
}
Ok(all_hits)
}
pub(super) fn per_batch_sc_neighbors(
layout: &PbSampleLayout,
batch_knn_lookup: &[ColumnDict<usize>],
knn: usize,
anchor_batches: Option<&[usize]>,
) -> anyhow::Result<Vec<Vec<(usize, f32)>>> {
use indicatif::ParallelProgressIterator;
let num_pb = layout.cell_counts.len();
let prog_bar = styled_progress_bar(num_pb as u64, "pb-samples (BBKNN match)");
let result = (0..num_pb)
.into_par_iter()
.progress_with(prog_bar.clone())
.map(|pbsamp| bbknn_match_one_pbsamp(layout, batch_knn_lookup, knn, pbsamp, anchor_batches))
.collect();
prog_bar.finish_and_clear();
result
}
pub(super) fn build_pb_sample_to_cells(layout: &PbSampleLayout) -> Vec<Vec<usize>> {
let num_pb = layout.cell_counts.len();
let mut out: Vec<Vec<usize>> = vec![vec![]; num_pb];
for (c, &pbsamp) in layout.cell_to_pbsamp.iter().enumerate() {
if pbsamp != usize::MAX {
out[pbsamp].push(c);
}
}
out
}
#[cfg(test)]
#[path = "pb_samples_tests.rs"]
mod tests;