#![allow(dead_code)]
use crate::alg::collapse_data::PbSampleLayout;
use crate::alg::dc_poisson::{
compute_sibling_sets, intersect_with_siblings_fallback, refine_with_proposer,
CandidateProposer, ProfileSource, Profiles, RefineContext,
};
use log::{debug, info};
use rand::rngs::SmallRng;
use rand::SeedableRng;
use rustc_hash::FxHashMap;
pub use crate::alg::collapse_data::GeneSums;
pub use crate::alg::dc_poisson::{compact_labels, RefineParams};
#[derive(Debug, Clone)]
pub struct RefinedAssignment {
pub pbsamp_to_group: Vec<Vec<usize>>,
pub num_groups_per_level: Vec<usize>,
}
fn build_bbknn_neighbors(
layout: &PbSampleLayout,
batch_knn_lookup: &[legume_numeric::matrix::knn_match::ColumnDict<usize>],
knn: usize,
) -> anyhow::Result<Vec<Vec<usize>>> {
use crate::alg::collapse_data::bbknn_match_one_pbsamp;
use rayon::prelude::*;
let num_pb = layout.cell_counts.len();
(0..num_pb)
.into_par_iter()
.map(|pbsamp| -> anyhow::Result<Vec<usize>> {
let hits = bbknn_match_one_pbsamp(layout, batch_knn_lookup, knn, pbsamp, None)?;
Ok(hits.into_iter().map(|(p, _)| p).collect())
})
.collect()
}
fn build_candidate_sets(
siblings: &[Vec<usize>],
bbknn: &[Vec<usize>],
pbsamp_to_group_at_level: &[usize],
) -> Vec<Vec<usize>> {
siblings
.iter()
.enumerate()
.map(|(pbsamp, sib)| {
let mut neighbor_groups: Vec<usize> = bbknn[pbsamp]
.iter()
.map(|&j| pbsamp_to_group_at_level[j])
.collect();
neighbor_groups.sort_unstable();
neighbor_groups.dedup();
intersect_with_siblings_fallback(
sib,
&neighbor_groups,
pbsamp_to_group_at_level[pbsamp],
)
})
.collect()
}
pub struct BbknnProposer<'a> {
siblings: Vec<Vec<usize>>,
bbknn: &'a [Vec<usize>],
}
impl<'a> BbknnProposer<'a> {
pub fn new(siblings: Vec<Vec<usize>>, bbknn: &'a [Vec<usize>]) -> Self {
Self { siblings, bbknn }
}
}
impl<'a> CandidateProposer for BbknnProposer<'a> {
fn propose(&self, labels: &[usize]) -> Vec<Vec<usize>> {
build_candidate_sets(&self.siblings, self.bbknn, labels)
}
}
#[derive(Clone, Copy)]
pub struct RefineInputs<'a> {
pub layout: &'a PbSampleLayout,
pub gene_sums: &'a GeneSums,
pub num_genes: usize,
pub pb_sample_to_cells: &'a [Vec<usize>],
pub batch_knn_lookup: &'a [legume_numeric::matrix::knn_match::ColumnDict<usize>],
pub k_per_batch: usize,
pub initial_sc_to_group_per_level: &'a [Vec<usize>],
pub reproject_offsets_per_level: &'a [Vec<usize>],
}
pub fn refine_assignments(
inputs: &RefineInputs<'_>,
params: &RefineParams,
) -> anyhow::Result<RefinedAssignment> {
let RefineInputs {
layout,
gene_sums,
num_genes,
pb_sample_to_cells,
batch_knn_lookup,
k_per_batch,
initial_sc_to_group_per_level,
reproject_offsets_per_level,
} = *inputs;
let num_levels = initial_sc_to_group_per_level.len();
if num_levels == 0 {
return Err(anyhow::anyhow!("no levels"));
}
let num_pb = layout.cell_counts.len();
for (i, lvl) in initial_sc_to_group_per_level.iter().enumerate() {
if lvl.len() != num_pb {
return Err(anyhow::anyhow!(
"level {} has {} entries, expected {}",
i,
lvl.len(),
num_pb
));
}
}
let mut refined: Vec<Vec<usize>> = Vec::with_capacity(num_levels);
let mut ks: Vec<usize> = Vec::with_capacity(num_levels);
for lvl in initial_sc_to_group_per_level {
let (compact, k) = compact_labels(lvl);
refined.push(compact);
ks.push(k);
}
if params.num_gibbs == 0 && params.num_greedy == 0 {
info!("Skipping DC-Poisson refinement: --pb-refine-gibbs=0 and --pb-refine-greedy=0");
return Ok(RefinedAssignment {
pbsamp_to_group: refined,
num_groups_per_level: ks,
});
}
info!(
"Building DC-Poisson profiles ({} pb-samples × {} genes) ...",
num_pb, num_genes
);
let mut profiles = match ¶ms.profile_source {
ProfileSource::Raw => Profiles::from_gene_sums(gene_sums, num_genes),
ProfileSource::Projected { basis } => Profiles::from_projection(basis, pb_sample_to_cells),
};
if matches!(params.profile_source, ProfileSource::Raw)
&& !matches!(
params.feature_weighting,
crate::alg::dc_poisson::FeatureWeighting::None
)
{
info!("Computing NB Fisher-info feature weights ...");
profiles.apply_feature_weighting(params.feature_weighting);
}
info!(
"Building BBKNN candidate sets (knn={} per non-own batch) ...",
k_per_batch
);
let bbknn = build_bbknn_neighbors(layout, batch_knn_lookup, k_per_batch)?;
let mut rng = SmallRng::seed_from_u64(params.seed);
for level in (0..num_levels).rev() {
if level + 1 < num_levels {
let offset: Vec<usize> = match reproject_offsets_per_level.get(level) {
Some(o) if !o.is_empty() => o.clone(),
_ => child_offset_within_parent(
&initial_sc_to_group_per_level[level],
&initial_sc_to_group_per_level[level + 1],
),
};
let (reprojected, new_k) = project_to_refinement(&offset, &refined[level + 1]);
refined[level] = reprojected;
ks[level] = new_k;
}
let k = ks[level];
debug!("refining level {} (k={}, num_pb={})", level, k, num_pb);
let siblings = compute_sibling_sets(&refined, level, k);
let proposer = BbknnProposer::new(siblings, &bbknn);
let pbsamp_to_group = &mut refined[level];
let label = format!("Refine L{}/{}", num_levels - level, num_levels);
let moves = refine_with_proposer(
pbsamp_to_group,
&proposer,
&mut rng,
&RefineContext {
profiles: &profiles,
k,
params,
level_label: &label,
},
);
info!(" level {} refined: {} moves; k={} groups", level, moves, k);
let (compact, new_k) = compact_labels(pbsamp_to_group);
*pbsamp_to_group = compact;
ks[level] = new_k;
}
Ok(RefinedAssignment {
pbsamp_to_group: refined,
num_groups_per_level: ks,
})
}
fn project_to_refinement(child: &[usize], parent: &[usize]) -> (Vec<usize>, usize) {
debug_assert_eq!(child.len(), parent.len());
let pairs: Vec<(usize, usize)> = child.iter().copied().zip(parent.iter().copied()).collect();
compact_labels(&pairs)
}
fn child_offset_within_parent(child: &[usize], parent: &[usize]) -> Vec<usize> {
debug_assert_eq!(child.len(), parent.len());
let mut per_parent: FxHashMap<usize, FxHashMap<usize, usize>> = FxHashMap::default();
let mut offsets = vec![0usize; child.len()];
for (i, (&c, &p)) in child.iter().zip(parent.iter()).enumerate() {
let local = per_parent.entry(p).or_default();
let next = local.len();
offsets[i] = *local.entry(c).or_insert(next);
}
offsets
}
#[cfg(test)]
#[path = "refine_multilevel_tests.rs"]
mod tests;