#![allow(dead_code)]
use crate::sparse_data_visitors::*;
use crate::sparse_io_stack::SparseIoStack;
use crate::sparse_io_vector::SparseIoVec;
use legume_numeric::matrix::knn_match::ColumnDict;
use legume_numeric::matrix::traits::*;
use legume_numeric::param::dmatrix_gamma::*;
use legume_numeric::param::traits::Inference;
use legume_numeric::param::traits::*;
use log::{info, warn};
use nalgebra::DMatrix;
use rayon::prelude::*;
use std::ops::AddAssign;
use std::sync::{Arc, Mutex};
use crate::alg::random_projection::binary_sort_columns;
use rustc_hash::FxHashMap as HashMap;
type CscMat = nalgebra_sparse::CscMatrix<f32>;
pub type GeneSums = Vec<Vec<(usize, f32)>>;
pub const DEFAULT_KNN: usize = 10;
pub const DEFAULT_OPT_ITER: usize = 100;
mod pb_samples;
use pb_samples::{
build_pb_sample_layout, build_pb_sample_to_cells, build_pb_samples,
collect_pb_sample_gene_sums, per_batch_sc_neighbors,
};
pub use pb_samples::{PbSampleCollection, PbSampleLayout};
pub(crate) use pb_samples::bbknn_match_one_pbsamp;
mod reassign_cells;
pub use reassign_cells::ReassignCellsParams;
mod pb_tree;
pub use pb_tree::{BelowEdge, ContrastGene, PbTree, PbTreeParams, RootRecord, SplitRecord};
mod refine;
pub use refine::StackCollapseOut;
use refine::{
compute_level_sort_dims, fine_to_coarse_from_refined, pad_numeric_labels,
refine_and_collect_single_layer, refine_and_collect_stack, split_anchored_finest_groups,
RefineCollectCtx,
};
mod stats;
use stats::{
collect_basic_stat_visitor, collect_batch_stat_visitor, collect_matched_stat_coarse,
collect_matched_stat_visitor, merge_stat, optimize, KnnParams, DEFAULT_NUM_LEVELS,
};
pub use stats::{resample_and_optimize, CollapsedOut, CollapsedStat};
mod strata;
pub use strata::collapse_columns_multilevel_with_strata;
pub struct MultilevelCollapseOut {
pub levels: Vec<CollapsedOut>,
pub cell_to_pb_per_level: Vec<Vec<usize>>,
pub pb_tree: Option<PbTree>,
}
#[derive(Clone)]
pub struct MultilevelParams {
pub knn_pb_samples: usize,
pub num_levels: usize,
pub sort_dim: usize,
pub num_opt_iter: usize,
pub refine: crate::alg::refine_multilevel::RefineParams,
pub output_calibration: legume_numeric::param::traits::CalibrateTarget,
pub anchor_batches: Option<Vec<Box<str>>>,
pub bulk_batches: Option<Vec<Box<str>>>,
pub observe_panels: bool,
pub keep_finest_stats: bool,
pub pb_tree: Option<PbTreeParams>,
pub strata: Option<Vec<usize>>,
}
impl MultilevelParams {
pub fn new(proj_dim: usize) -> Self {
Self {
knn_pb_samples: DEFAULT_KNN,
num_levels: DEFAULT_NUM_LEVELS,
sort_dim: proj_dim.min(12),
num_opt_iter: DEFAULT_OPT_ITER,
refine: crate::alg::refine_multilevel::RefineParams::default(),
output_calibration: legume_numeric::param::traits::CalibrateTarget::All,
anchor_batches: None,
bulk_batches: None,
observe_panels: true,
keep_finest_stats: false,
pb_tree: None,
strata: None,
}
}
#[must_use]
pub fn with_strata(&self, strata: Vec<usize>) -> Self {
let mut p = self.clone();
p.strata = Some(strata);
p
}
}
fn stratum_bits(strata: &[usize]) -> usize {
let mut occupied: Vec<usize> = strata.to_vec();
occupied.sort_unstable();
occupied.dedup();
if occupied.len() <= 1 {
return 0;
}
let max = *occupied.last().unwrap_or(&0);
let n = max + 1;
(usize::BITS as usize - n.saturating_sub(1).leading_zeros() as usize).max(1)
}
fn apply_strata_to_codes(
codes: &[usize],
level_dims: &[usize],
strata: &[usize],
) -> anyhow::Result<(Vec<usize>, Vec<usize>, usize)> {
anyhow::ensure!(
codes.len() == strata.len(),
"strata has {} entries, codes have {}",
strata.len(),
codes.len()
);
let s_bits = stratum_bits(strata);
if s_bits == 0 {
return Ok((codes.to_vec(), level_dims.to_vec(), 0));
}
let stratified: Vec<usize> = codes
.iter()
.zip(strata.iter())
.map(|(&c, &s)| (c << s_bits) | s)
.collect();
Ok((stratified, level_dims.to_vec(), s_bits))
}
struct FinestCodes {
codes: Vec<usize>,
widths: Vec<usize>,
tree: Option<PbTree>,
strata_bits: usize,
}
fn finest_codes(
data_vec: &SparseIoVec,
proj_kn: &DMatrix<f32>,
level_dims: &[usize],
params: &MultilevelParams,
) -> anyhow::Result<FinestCodes> {
let finest_dim = level_dims[0];
let nn = proj_kn.ncols();
let kk = proj_kn.nrows().min(finest_dim).min(nn);
let codes = binary_sort_columns(proj_kn, kk)?;
let (codes, widths, tree) = match params.pb_tree.as_ref() {
None => (codes, level_dims.to_vec(), None),
Some(rb) => {
let coarse_bits = if level_dims.len() >= 2 {
*level_dims.last().expect("non-empty level dims")
} else {
stats::DEFAULT_COARSEST_SORT_DIM.min(kk)
};
let low_mask = (1usize << coarse_bits) - 1;
let n = data_vec.num_columns();
let active: Vec<bool> = if data_vec.has_column_multiplicity() {
(0..n)
.map(|c| (data_vec.column_multiplicity(c) - 1.0).abs() <= f32::EPSILON)
.collect()
} else {
vec![true; n]
};
let low: Vec<usize> = codes.iter().map(|&c| c & low_mask).collect();
let (mut node, _) = crate::alg::dc_poisson::compact_labels(&low);
let reassigned_cells = match rb.reassign_cells.as_ref() {
Some(cr) => {
let col_to_batch = data_vec.get_batch_membership(0..n);
let csc = data_vec.read_columns_csc(0..n)?;
reassign_cells::reassign_cells_to_nodes(
&csc,
&col_to_batch,
data_vec.num_batches().max(1),
&active,
&mut node,
cr,
)
}
None => 0,
};
for (c, nd) in node.iter_mut().enumerate() {
if !active[c] {
*nd = usize::MAX;
}
}
let targets: Vec<usize> = level_dims.iter().rev().map(|&d| 1usize << d).collect();
let (codes, widths, mut tree) =
pb_tree::build_tree(data_vec, &node, &codes, &targets, rb)?;
tree.reassigned_cells = reassigned_cells;
(codes, widths, Some(tree))
}
};
maybe_stratify_codes(codes, widths, tree, params)
}
fn maybe_stratify_codes(
codes: Vec<usize>,
widths: Vec<usize>,
tree: Option<PbTree>,
params: &MultilevelParams,
) -> anyhow::Result<FinestCodes> {
let Some(strata) = params.strata.as_deref() else {
return Ok(FinestCodes {
codes,
widths,
tree,
strata_bits: 0,
});
};
anyhow::ensure!(
strata.len() == codes.len(),
"MultilevelParams.strata has {} entries, codes have {}",
strata.len(),
codes.len()
);
let (codes, widths, s_bits) = apply_strata_to_codes(&codes, &widths, strata)?;
let n_occ = {
let mut u = strata.to_vec();
u.sort_unstable();
u.dedup();
u.len()
};
info!(
"CNV strata: crossed finest codes with {} cell strata ({} occupied, {} stratum bits)",
strata.len(),
n_occ,
s_bits
);
Ok(FinestCodes {
codes,
widths,
tree,
strata_bits: s_bits,
})
}
fn resolve_named_batches(
data_vec: &SparseIoVec,
role: &str,
names: Option<&[Box<str>]>,
) -> anyhow::Result<Option<Vec<usize>>> {
let Some(names) = names else { return Ok(None) };
let map = data_vec
.batch_name_map()
.ok_or_else(|| anyhow::anyhow!("{role} batches given but no batches are registered"))?;
let mut idx = Vec::with_capacity(names.len());
for n in names {
let Some(&b) = map.get(n) else {
anyhow::bail!(
"{role} batch `{n}` is not among the registered batches ({:?})",
map.keys().collect::<Vec<_>>(),
);
};
idx.push(b);
}
Ok(Some(idx))
}
fn greedy_anchor_for_bulk(
data_vec: &SparseIoVec,
anchors: Option<Vec<usize>>,
bulk: Option<&[usize]>,
) -> Option<Vec<usize>> {
match (anchors, bulk) {
(Some(a), _) => Some(a),
(None, Some(b)) if !b.is_empty() => {
let all = data_vec.num_batches();
let frame: Vec<usize> = (0..all).filter(|i| !b.contains(i)).collect();
(!frame.is_empty()).then(|| {
info!(
"Greedy bulk correction: {} bulk batch(es) corrected toward {} cell batch(es); \
the cell frame is anchored and does not move",
b.len(),
frame.len()
);
frame
})
}
(None, _) => None,
}
}
fn ensure_disjoint_roles(anchors: Option<&[usize]>, bulk: Option<&[usize]>) -> anyhow::Result<()> {
if let (Some(a), Some(b)) = (anchors, bulk) {
if let Some(shared) = a.iter().find(|x| b.contains(x)) {
anyhow::bail!(
"batch index {shared} is named as both an anchor and a bulk batch; \
the roles are mutually exclusive"
);
}
}
Ok(())
}
fn attach_observability(stat: &mut CollapsedStat, data_vec: &SparseIoVec) -> anyhow::Result<()> {
let Some(coverage) = data_vec.row_coverage_by_backend() else {
return Ok(());
};
let ncols = data_vec.num_columns();
let mut sources = Vec::with_capacity(ncols);
for c in 0..ncols {
match data_vec.column_source(c) {
Some(b) => sources.push(b),
None => {
warn!(
"panel observability skipped: column {c} merges several backends \
(column-union alignment)"
);
return Ok(());
}
}
}
let num_genes = stat.num_genes();
let num_samples = stat.num_samples();
let num_sources = data_vec.len();
let cell_to_group = data_vec.get_group_membership(0..ncols)?;
let mut count_bs = DMatrix::<f32>::zeros(num_sources, num_samples);
let batch_of = data_vec.get_batch_membership(0..ncols);
let num_batches = stat.num_batches();
let mut source_in_batch = vec![vec![false; num_batches]; num_sources];
for c in 0..ncols {
let s = cell_to_group[c];
if s < num_samples {
count_bs[(sources[c], s)] += data_vec.column_multiplicity(c);
}
if let Some(&b) = batch_of.get(c) {
if b < num_batches {
source_in_batch[sources[c]][b] = true;
}
}
}
let mut size_ds = DMatrix::<f32>::zeros(num_genes, num_samples);
for (src, cov) in coverage.iter().enumerate() {
anyhow::ensure!(
cov.len() == num_genes,
"row coverage has {} rows but the stat has {num_genes} genes",
cov.len(),
);
for (g, &covered) in cov.iter().enumerate() {
if covered {
for s in 0..num_samples {
size_ds[(g, s)] += count_bs[(src, s)];
}
}
}
}
let mut mask_db = DMatrix::<f32>::zeros(num_genes, num_batches);
for (src, cov) in coverage.iter().enumerate() {
for (b, &used) in source_in_batch[src].iter().enumerate() {
if used {
for (g, &covered) in cov.iter().enumerate() {
if covered {
mask_db[(g, b)] = 1.0;
}
}
}
}
}
info!(
"panel observability attached: {} of {} (gene, sample) entries below full size",
size_ds
.column_iter()
.enumerate()
.map(|(s, col)| col.iter().filter(|&&v| v < stat.size_s[s]).count())
.sum::<usize>(),
num_genes * num_samples,
);
stat.size_ds = Some(size_ds);
stat.obs_mask_db = (mask_db.iter().any(|&v| v == 0.0)).then_some(mask_db);
Ok(())
}
pub struct EmptyArg {}
#[cfg(debug_assertions)]
use log::debug;
pub trait CollapsingOps {
fn collapse_columns(
&self,
knn_batches: Option<usize>,
knn_cells: Option<usize>,
reference_batch_names: Option<&[Box<str>]>,
num_opt_iter: Option<usize>,
) -> anyhow::Result<CollapsedOut>;
fn build_hnsw_per_batch<T>(
&mut self,
proj_kn: &nalgebra::DMatrix<f32>,
col_to_batch: &[T],
) -> anyhow::Result<()>
where
T: Sync + Send + std::hash::Hash + Eq + Clone + ToString;
fn collect_basic_stat(&self, stat: &mut CollapsedStat) -> anyhow::Result<()>;
fn collect_batch_stat(&self, stat: &mut CollapsedStat) -> anyhow::Result<()>;
fn collect_matched_stat(
&self,
knn_batches: usize,
knn_cols: usize,
reference_indices: Option<&[usize]>,
stat: &mut CollapsedStat,
) -> anyhow::Result<()>;
}
impl CollapsingOps for SparseIoVec {
fn build_hnsw_per_batch<T>(
&mut self,
proj_kn: &nalgebra::DMatrix<f32>,
col_to_batch: &[T],
) -> anyhow::Result<()>
where
T: Sync + Send + std::hash::Hash + Eq + Clone + ToString,
{
info!("creating batch-specific HNSW maps ...");
self.register_batches_dmatrix(proj_kn, col_to_batch)?;
info!(
"found {} columns across {} batches",
self.num_columns(),
self.num_batches()
);
Ok(())
}
fn collapse_columns(
&self,
knn_batches: Option<usize>,
knn_cells: Option<usize>,
reference_batch_names: Option<&[Box<str>]>,
num_opt_iter: Option<usize>,
) -> anyhow::Result<CollapsedOut> {
let group_to_cols = self.take_grouped_columns().ok_or(anyhow::anyhow!(
"The columns were not assigned before. Call `assign_columns_to_groups`"
))?;
let num_features = self.num_rows();
let num_groups = group_to_cols.len();
let num_batches = self.num_batches();
let mut stat = CollapsedStat::new(num_features, num_groups, num_batches);
info!("basic statistics across {} groups", num_groups);
self.collect_basic_stat(&mut stat)?;
if num_batches > 1 {
info!(
"batch-specific statistics across {} batches over {} samples",
num_batches, num_groups
);
let batch_name_map = self
.batch_name_map()
.ok_or(anyhow::anyhow!("unable to read batch names"))?;
let reference_indices = reference_batch_names.map(|x| {
x.iter()
.filter_map(|b| batch_name_map.get(b))
.copied()
.collect::<Vec<_>>()
});
if let Some(r) = reference_indices.as_ref() {
if r.is_empty() {
let ref_names = reference_batch_names
.unwrap()
.iter()
.map(|x| x.to_string())
.collect::<Vec<_>>()
.join(",");
let bat_names = self
.batch_names()
.unwrap()
.iter()
.map(|x| x.to_string())
.collect::<Vec<_>>()
.join(",");
warn!("{} vs. {}", ref_names, bat_names);
return Err(anyhow::anyhow!("no reference batch names matched!"));
}
}
self.collect_batch_stat(&mut stat)?;
info!(
"counterfactual inference across {} batches over {} samples",
num_batches, num_groups,
);
let knn_batches = knn_batches.unwrap_or(2);
let knn_cells = knn_cells.unwrap_or(DEFAULT_KNN);
self.collect_matched_stat(
knn_batches,
knn_cells,
reference_indices.as_deref(),
&mut stat,
)?;
}
info!("optimizing the collapsed parameters...");
let (a0, b0) = (1_f32, 1_f32);
optimize(
&stat,
(a0, b0),
num_opt_iter.unwrap_or(DEFAULT_OPT_ITER),
"Optimizing",
CalibrateTarget::All,
false,
)
}
fn collect_basic_stat(&self, stat: &mut CollapsedStat) -> anyhow::Result<()> {
self.visit_columns_by_group(&collect_basic_stat_visitor, &EmptyArg {}, stat)
}
fn collect_batch_stat(&self, stat: &mut CollapsedStat) -> anyhow::Result<()> {
self.visit_columns_by_group(&collect_batch_stat_visitor, &EmptyArg {}, stat)
}
fn collect_matched_stat(
&self,
knn_batches: usize,
knn_cells: usize,
reference_indices: Option<&[usize]>,
stat: &mut CollapsedStat,
) -> anyhow::Result<()> {
stat.anchor_batches = reference_indices.map(<[usize]>::to_vec).unwrap_or_default();
self.visit_columns_by_group(
&collect_matched_stat_visitor,
&KnnParams {
knn_batches,
knn_cells,
reference_indices,
},
stat,
)
}
}
pub trait MultilevelCollapsingOps {
type LevelOutput;
fn collapse_columns_multilevel<T>(
&mut self,
proj_kn: &DMatrix<f32>,
batch_membership: &[T],
params: &MultilevelParams,
) -> anyhow::Result<Self::LevelOutput>
where
T: Sync + Send + std::hash::Hash + Eq + Clone + ToString;
fn collapse_columns_multilevel_vec<T>(
&mut self,
proj_kn: &DMatrix<f32>,
batch_membership: &[T],
params: &MultilevelParams,
) -> anyhow::Result<Vec<Self::LevelOutput>>
where
T: Sync + Send + std::hash::Hash + Eq + Clone + ToString;
}
pub fn collapse_columns_multilevel_with_hierarchy<T>(
data_vec: &mut SparseIoVec,
proj_kn: &DMatrix<f32>,
batch_membership: &[T],
params: &MultilevelParams,
) -> anyhow::Result<MultilevelCollapseOut>
where
T: Sync + Send + std::hash::Hash + Eq + Clone + ToString,
{
let sort_dim = params.sort_dim;
let knn = params.knn_pb_samples;
let opt_iter = params.num_opt_iter;
data_vec.register_batch_membership(batch_membership);
let num_features = data_vec.num_rows();
let num_batches = data_vec.num_batches();
if num_batches >= 2 {
data_vec.build_hnsw_per_batch(proj_kn, batch_membership)?;
}
let level_dims = compute_level_sort_dims(sort_dim, params.num_levels);
let FinestCodes {
codes: fine_codes,
widths: level_dims,
tree: pb_tree,
strata_bits,
} = finest_codes(data_vec, proj_kn, &level_dims, params)?;
data_vec.assign_groups(&fine_codes, None);
let group_to_cols = data_vec
.take_grouped_columns()
.ok_or_else(|| anyhow::anyhow!("columns not assigned"))?
.clone();
let refine_params = ¶ms.refine;
let anchor_batches =
resolve_named_batches(data_vec, "anchor", params.anchor_batches.as_deref())?;
let bulk_batches = resolve_named_batches(data_vec, "bulk", params.bulk_batches.as_deref())?;
ensure_disjoint_roles(anchor_batches.as_deref(), bulk_batches.as_deref())?;
let summary_batches = anchor_batches.clone();
let anchor_batches = greedy_anchor_for_bulk(data_vec, anchor_batches, bulk_batches.as_deref());
let ctx = RefineCollectCtx {
fine_codes: &fine_codes,
group_to_cols_finest: &group_to_cols,
level_dims: &level_dims,
num_features,
num_batches,
knn,
opt_iter,
refine_params,
output_calibration: params.output_calibration,
anchor_batches: anchor_batches.as_deref(),
summary_batches: summary_batches.as_deref(),
bulk_batches: bulk_batches.as_deref(),
observe_panels: params.observe_panels,
keep_finest_stats: params.keep_finest_stats,
pb_tree: pb_tree.as_ref(),
cell_to_stratum: params.strata.as_deref(),
exclude_unmatched_from_delta: params.strata.is_some(),
strata_bits,
};
refine_and_collect_single_layer(data_vec, proj_kn, &ctx)
}
pub fn collapse_columns_multilevel_with_partition<T>(
data_vec: &mut SparseIoVec,
proj_kn: &DMatrix<f32>,
batch_membership: &[T],
params: &MultilevelParams,
cell_to_pb_per_level: &[Vec<usize>],
) -> anyhow::Result<MultilevelCollapseOut>
where
T: Sync + Send + std::hash::Hash + Eq + Clone + ToString,
{
let knn = params.knn_pb_samples;
let opt_iter = params.num_opt_iter;
data_vec.register_batch_membership(batch_membership);
let num_features = data_vec.num_rows();
let num_batches = data_vec.num_batches();
if num_batches >= 2 {
data_vec.build_hnsw_per_batch(proj_kn, batch_membership)?;
}
let level_dims = compute_level_sort_dims(params.sort_dim, params.num_levels);
anyhow::ensure!(
cell_to_pb_per_level.len() == level_dims.len(),
"inherited cell_to_pb has {} levels but --num-levels is {}; \
pass --num-levels to match the source run",
cell_to_pb_per_level.len(),
level_dims.len(),
);
let inherited_finest = &cell_to_pb_per_level[0];
anyhow::ensure!(
inherited_finest.len() == proj_kn.ncols(),
"inherited cell_to_pb finest level has {} cells, data has {}",
inherited_finest.len(),
proj_kn.ncols()
);
let k_inherited = inherited_finest.iter().copied().max().map_or(0, |m| m + 1);
data_vec.assign_groups(&pad_numeric_labels(inherited_finest, k_inherited), None);
let anchor_batches =
resolve_named_batches(data_vec, "anchor", params.anchor_batches.as_deref())?;
let bulk_batches = resolve_named_batches(data_vec, "bulk", params.bulk_batches.as_deref())?;
ensure_disjoint_roles(anchor_batches.as_deref(), bulk_batches.as_deref())?;
let summary_batches = anchor_batches.clone();
let anchor_batches = greedy_anchor_for_bulk(data_vec, anchor_batches, bulk_batches.as_deref());
let pb_samples = build_pb_samples(
data_vec,
proj_kn,
num_features,
summary_batches.as_deref().unwrap_or(&[]),
bulk_batches.as_deref().unwrap_or(&[]),
params.strata.as_deref(),
)?;
let num_pb = pb_samples.layout.cell_counts.len();
let ncols = proj_kn.ncols();
let pb_sample_to_cells = build_pb_sample_to_cells(&pb_samples.layout);
let num_levels = level_dims.len();
let mut pbsamp_to_group: Vec<Vec<usize>> = Vec::with_capacity(num_levels);
let mut num_groups_per_level: Vec<usize> = Vec::with_capacity(num_levels);
for (lvl_idx, lvl) in cell_to_pb_per_level.iter().enumerate() {
anyhow::ensure!(
lvl.len() == ncols,
"inherited cell_to_pb level {} has {} cells, data has {}",
lvl_idx,
lvl.len(),
ncols
);
let mut p2g: Vec<usize> = Vec::with_capacity(num_pb);
for cells in &pb_sample_to_cells {
p2g.push(modal_group(cells, lvl));
}
let (compact, k) = crate::alg::refine_multilevel::compact_labels(&p2g);
num_groups_per_level.push(k);
pbsamp_to_group.push(compact);
}
let refined = split_anchored_finest_groups(
crate::alg::refine_multilevel::RefinedAssignment {
pbsamp_to_group,
num_groups_per_level,
},
&pb_samples.layout,
);
info!(
"Inherited partition: {} cells, {} pb-samples, finest k={} (skipped BBKNN + DC-SBM refinement)",
ncols, num_pb, refined.num_groups_per_level[0]
);
let k_finest = refined.num_groups_per_level[0];
let mut cell_to_group_finest = vec![0usize; ncols];
for (pbsamp, cells) in pb_sample_to_cells.iter().enumerate() {
let g = refined.pbsamp_to_group[0][pbsamp];
for &c in cells {
cell_to_group_finest[c] = g;
}
}
let finest_str = pad_numeric_labels(&cell_to_group_finest, k_finest);
data_vec.assign_groups(&finest_str, None);
debug_assert_eq!(data_vec.num_groups(), k_finest);
let mut fine_stat = CollapsedStat::new(num_features, k_finest, num_batches);
fine_stat.exclude_unmatched_from_delta = params.strata.is_some();
info!("Collecting basic stats over {} groups ...", k_finest);
data_vec.collect_basic_stat(&mut fine_stat)?;
if num_batches >= 2 {
info!(
"Collecting per-batch stats over {} groups × {} batches ...",
k_finest, num_batches
);
data_vec.collect_batch_stat(&mut fine_stat)?;
let batch_knn = data_vec
.batch_knn_lookup()
.ok_or_else(|| anyhow::anyhow!("batch_knn_lookup not built"))?;
info!(
"Collecting cross-batch matched stats (knn={}) over {} pb-samples ...",
knn, num_pb
);
collect_matched_stat_coarse(
&pb_samples.layout,
&pb_samples.gene_sums,
&refined.pbsamp_to_group[0],
batch_knn.as_slice(),
knn,
anchor_batches.as_deref(),
&mut fine_stat,
)?;
}
if params.observe_panels {
attach_observability(&mut fine_stat, data_vec)?;
}
let mut results: Vec<CollapsedOut> = Vec::with_capacity(num_levels);
info!(
"Level 1/{}: inherited k={} (finest; {} cells)",
num_levels, k_finest, ncols
);
let finest_out = optimize(
&fine_stat,
(1.0, 1.0),
opt_iter,
&format!("Inherit L1/{}", num_levels),
CalibrateTarget::All,
false,
)?;
results.push(finest_out);
let mut prev_stat = fine_stat;
for level in 1..num_levels {
let k_prev = refined.num_groups_per_level[level - 1];
let k_level = refined.num_groups_per_level[level];
let fine_to_coarse = fine_to_coarse_from_refined(
&refined.pbsamp_to_group[level - 1],
&refined.pbsamp_to_group[level],
k_prev,
);
let coarse_stat = merge_stat(&prev_stat, &fine_to_coarse, k_level);
info!(
"Level {}/{}: inherited k={} (merged from {})",
level + 1,
num_levels,
k_level,
k_prev
);
let level_opt_iter = (opt_iter / 2).max(10);
let out = optimize(
&coarse_stat,
(1.0, 1.0),
level_opt_iter,
&format!("Inherit L{}/{}", level + 1, num_levels),
CalibrateTarget::All,
false,
)?;
results.push(out);
prev_stat = coarse_stat;
}
info!(
"Fitted pseudobulk posteriors for {} inherited levels: k = {:?} (finest first)",
num_levels, refined.num_groups_per_level
);
let mut cell_to_pb_per_level_out: Vec<Vec<usize>> = Vec::with_capacity(num_levels);
for level in 0..num_levels {
let mut c2g = vec![0usize; ncols];
for (pbsamp, cells) in pb_sample_to_cells.iter().enumerate() {
let g = refined.pbsamp_to_group[level][pbsamp];
for &c in cells {
c2g[c] = g;
}
}
cell_to_pb_per_level_out.push(c2g);
}
Ok(MultilevelCollapseOut {
levels: results,
cell_to_pb_per_level: cell_to_pb_per_level_out,
pb_tree: None,
})
}
fn modal_group(cells: &[usize], lvl: &[usize]) -> usize {
match cells {
[] => 0,
[c] => lvl[*c],
_ => {
use rustc_hash::FxHashMap;
let mut counts: FxHashMap<usize, usize> = FxHashMap::default();
for &c in cells {
*counts.entry(lvl[c]).or_insert(0) += 1;
}
counts
.into_iter()
.max_by_key(|&(_, n)| n)
.map(|(g, _)| g)
.unwrap_or(0)
}
}
}
impl MultilevelCollapsingOps for SparseIoVec {
type LevelOutput = CollapsedOut;
fn collapse_columns_multilevel<T>(
&mut self,
proj_kn: &DMatrix<f32>,
batch_membership: &[T],
params: &MultilevelParams,
) -> anyhow::Result<CollapsedOut>
where
T: Sync + Send + std::hash::Hash + Eq + Clone + ToString,
{
let mut results =
self.collapse_columns_multilevel_vec(proj_kn, batch_membership, params)?;
if results.is_empty() {
return Err(anyhow::anyhow!("no levels processed"));
}
Ok(results.remove(0))
}
fn collapse_columns_multilevel_vec<T>(
&mut self,
proj_kn: &DMatrix<f32>,
batch_membership: &[T],
params: &MultilevelParams,
) -> anyhow::Result<Vec<CollapsedOut>>
where
T: Sync + Send + std::hash::Hash + Eq + Clone + ToString,
{
let sort_dim = params.sort_dim;
let knn = params.knn_pb_samples;
let opt_iter = params.num_opt_iter;
self.register_batch_membership(batch_membership);
let num_features = self.num_rows();
let num_batches = self.num_batches();
if num_batches >= 2 {
self.build_hnsw_per_batch(proj_kn, batch_membership)?;
}
let level_dims = compute_level_sort_dims(sort_dim, params.num_levels);
info!(
"Multi-level collapsing (fine→coarse): {} levels, sort_dims={:?}, {} batches",
level_dims.len(),
level_dims,
num_batches,
);
let FinestCodes {
codes: fine_codes,
widths: level_dims,
tree: pb_tree,
strata_bits,
} = finest_codes(self, proj_kn, &level_dims, params)?;
self.assign_groups(&fine_codes, None);
let group_to_cols = self
.take_grouped_columns()
.ok_or(anyhow::anyhow!("columns not assigned"))?
.clone();
let refine_params = ¶ms.refine;
let anchor_batches =
resolve_named_batches(self, "anchor", params.anchor_batches.as_deref())?;
let bulk_batches = resolve_named_batches(self, "bulk", params.bulk_batches.as_deref())?;
ensure_disjoint_roles(anchor_batches.as_deref(), bulk_batches.as_deref())?;
let summary_batches = anchor_batches.clone();
let anchor_batches = greedy_anchor_for_bulk(self, anchor_batches, bulk_batches.as_deref());
let ctx = RefineCollectCtx {
fine_codes: &fine_codes,
group_to_cols_finest: &group_to_cols,
level_dims: &level_dims,
num_features,
num_batches,
knn,
opt_iter,
refine_params,
output_calibration: params.output_calibration,
anchor_batches: anchor_batches.as_deref(),
summary_batches: summary_batches.as_deref(),
bulk_batches: bulk_batches.as_deref(),
observe_panels: params.observe_panels,
keep_finest_stats: params.keep_finest_stats,
pb_tree: pb_tree.as_ref(),
cell_to_stratum: params.strata.as_deref(),
exclude_unmatched_from_delta: params.strata.is_some(),
strata_bits,
};
refine_and_collect_single_layer(self, proj_kn, &ctx).map(|out| out.levels)
}
}
impl MultilevelCollapsingOps for SparseIoStack {
type LevelOutput = Vec<CollapsedOut>;
fn collapse_columns_multilevel<T>(
&mut self,
proj_kn: &DMatrix<f32>,
batch_membership: &[T],
params: &MultilevelParams,
) -> anyhow::Result<Vec<CollapsedOut>>
where
T: Sync + Send + std::hash::Hash + Eq + Clone + ToString,
{
let mut results =
self.collapse_columns_multilevel_vec(proj_kn, batch_membership, params)?;
if results.is_empty() {
return Err(anyhow::anyhow!("no levels processed"));
}
Ok(results.remove(0))
}
fn collapse_columns_multilevel_vec<T>(
&mut self,
proj_kn: &DMatrix<f32>,
batch_membership: &[T],
params: &MultilevelParams,
) -> anyhow::Result<Vec<Vec<CollapsedOut>>>
where
T: Sync + Send + std::hash::Hash + Eq + Clone + ToString,
{
Ok(
collapse_stack_multilevel_with_hierarchy(self, proj_kn, batch_membership, params)?
.levels,
)
}
}
pub fn collapse_stack_multilevel_with_hierarchy<T>(
stack: &mut SparseIoStack,
proj_kn: &DMatrix<f32>,
batch_membership: &[T],
params: &MultilevelParams,
) -> anyhow::Result<StackCollapseOut>
where
T: Sync + Send + std::hash::Hash + Eq + Clone + ToString,
{
let num_layers = stack.num_types();
if num_layers == 0 {
return Err(anyhow::anyhow!("empty SparseIoStack"));
}
let sort_dim = params.sort_dim;
let knn = params.knn_pb_samples;
let opt_iter = params.num_opt_iter;
stack.register_batch_membership(batch_membership);
let num_batches = stack.stack[0].num_batches();
if num_batches >= 2 {
for layer in stack.stack.iter_mut() {
layer.build_hnsw_per_batch(proj_kn, batch_membership)?;
}
}
let ncols = proj_kn.ncols();
let refine_params = ¶ms.refine;
let level_dims = compute_level_sort_dims(sort_dim, params.num_levels);
let finest_dim = level_dims[0];
let kk = proj_kn.nrows().min(finest_dim).min(ncols);
let codes = binary_sort_columns(proj_kn, kk)?;
let (fine_codes, level_dims, strata_bits) = match params.strata.as_deref() {
Some(strata) => {
let (c, d, s) = apply_strata_to_codes(&codes, &level_dims, strata)?;
(c, d, s)
}
None => (codes, level_dims, 0),
};
for layer in stack.stack.iter_mut() {
layer.assign_groups(&fine_codes, None);
}
let group_to_cols = stack.stack[0]
.take_grouped_columns()
.ok_or(anyhow::anyhow!("columns not assigned"))?
.clone();
let num_features = stack.stack[0].num_rows();
anyhow::ensure!(
params.anchor_batches.is_none() && params.bulk_batches.is_none(),
"anchor_batches / bulk_batches are not supported on the stack path — nothing \
produces a carried reference or bulk input for stacked modalities"
);
let ctx = RefineCollectCtx {
fine_codes: &fine_codes,
group_to_cols_finest: &group_to_cols,
level_dims: &level_dims,
num_features,
num_batches,
knn,
opt_iter,
refine_params,
output_calibration: params.output_calibration,
anchor_batches: None,
summary_batches: None,
bulk_batches: None,
observe_panels: false,
keep_finest_stats: params.keep_finest_stats,
pb_tree: None,
cell_to_stratum: params.strata.as_deref(),
exclude_unmatched_from_delta: params.strata.is_some(),
strata_bits,
};
refine_and_collect_stack(stack, proj_kn, &ctx)
}