use super::stats::DEFAULT_COARSEST_SORT_DIM;
use super::*;
pub(super) fn pad_numeric_labels(cell_to_group: &[usize], k: usize) -> Vec<String> {
let width = {
let mut w = 1usize;
let mut n = k.max(1) - 1;
while n >= 10 {
w += 1;
n /= 10;
}
w
};
cell_to_group
.iter()
.map(|g| format!("{:0width$}", g, width = width))
.collect()
}
pub(super) fn fine_to_coarse_from_refined(
pbsamp_to_fine: &[usize],
pbsamp_to_coarse: &[usize],
num_fine: usize,
) -> Vec<usize> {
let mut mapping = vec![usize::MAX; num_fine];
for pbsamp in 0..pbsamp_to_fine.len() {
let f = pbsamp_to_fine[pbsamp];
if mapping[f] == usize::MAX {
mapping[f] = pbsamp_to_coarse[pbsamp];
} else {
debug_assert_eq!(
mapping[f], pbsamp_to_coarse[pbsamp],
"refinement broke hierarchy at fine group {}",
f
);
}
}
mapping
}
pub(super) fn initial_per_level_from_hash(
fine_codes: &[usize],
pb_sample_to_cells: &[Vec<usize>],
level_dims: &[usize],
strata_bits: usize,
) -> Vec<Vec<usize>> {
let num_pb = pb_sample_to_cells.len();
level_dims
.iter()
.map(|&d| {
let width = d.saturating_add(strata_bits);
let mask = if width >= usize::BITS as usize {
usize::MAX
} else {
(1_usize << width).wrapping_sub(1)
};
let codes: Vec<usize> = (0..num_pb)
.map(|pbsamp| fine_codes[pb_sample_to_cells[pbsamp][0]] & mask)
.collect();
crate::alg::refine_multilevel::compact_labels(&codes).0
})
.collect()
}
pub(super) fn build_reproject_offsets(
fine_codes: &[usize],
pb_sample_to_cells: &[Vec<usize>],
level_dims: &[usize],
strata_bits: usize,
) -> Vec<Vec<usize>> {
let raw: Vec<usize> = pb_sample_to_cells
.iter()
.map(|cells| fine_codes[cells[0]])
.collect();
(0..level_dims.len())
.map(|level| {
if level + 1 < level_dims.len() {
let parent_dim = level_dims[level + 1].saturating_add(strata_bits);
let nbits = level_dims[level].saturating_sub(level_dims[level + 1]);
let mask = if nbits >= usize::BITS as usize {
usize::MAX
} else {
(1_usize << nbits).wrapping_sub(1)
};
raw.iter().map(|&c| (c >> parent_dim) & mask).collect()
} else {
Vec::new()
}
})
.collect()
}
pub(super) fn refine_or_identity(
allow_refine: bool,
inputs: &crate::alg::refine_multilevel::RefineInputs<'_>,
refine_params: &crate::alg::refine_multilevel::RefineParams,
) -> anyhow::Result<crate::alg::refine_multilevel::RefinedAssignment> {
if allow_refine {
crate::alg::refine_multilevel::refine_assignments(inputs, refine_params)
} else {
let mut pbsamp_to_group: Vec<Vec<usize>> =
Vec::with_capacity(inputs.initial_sc_to_group_per_level.len());
let mut num_groups_per_level =
Vec::with_capacity(inputs.initial_sc_to_group_per_level.len());
for lvl in inputs.initial_sc_to_group_per_level {
let (compact, k) = crate::alg::refine_multilevel::compact_labels(lvl);
num_groups_per_level.push(k);
pbsamp_to_group.push(compact);
}
Ok(crate::alg::refine_multilevel::RefinedAssignment {
pbsamp_to_group,
num_groups_per_level,
})
}
}
pub(super) fn split_anchored_finest_groups(
mut refined: crate::alg::refine_multilevel::RefinedAssignment,
layout: &PbSampleLayout,
) -> crate::alg::refine_multilevel::RefinedAssignment {
let mut anchored: Vec<(usize, usize)> = layout
.singleton_col
.iter()
.enumerate()
.filter_map(|(p, col)| col.map(|c| (c, p)))
.collect();
if anchored.is_empty() {
return refined;
}
anchored.sort_unstable_by_key(|&(col, _)| col);
let finest = &mut refined.pbsamp_to_group[0];
let ordinary: Vec<usize> = layout
.singleton_col
.iter()
.zip(finest.iter())
.filter_map(|(a, &g)| a.is_none().then_some(g))
.collect();
let (compact, k_new) = crate::alg::refine_multilevel::compact_labels(&ordinary);
let mut compact = compact.into_iter();
for (p, a) in layout.singleton_col.iter().enumerate() {
if a.is_none() {
finest[p] = compact
.next()
.expect("one compacted label per ordinary pb-sample");
}
}
for (j, &(_, p)) in anchored.iter().enumerate() {
finest[p] = k_new + j;
}
refined.num_groups_per_level[0] = k_new + anchored.len();
info!(
"Append-only finest partition: {} new-data groups + {} carried + {} bulk singletons",
k_new,
anchored
.iter()
.filter(|&&(_, p)| !layout.is_bulk(p))
.count(),
anchored.iter().filter(|&&(_, p)| layout.is_bulk(p)).count(),
);
refined
}
#[derive(Clone, Copy)]
pub(super) struct RefineCollectCtx<'a> {
pub(super) fine_codes: &'a [usize],
pub(super) group_to_cols_finest: &'a [Vec<usize>],
pub(super) level_dims: &'a [usize],
pub(super) num_features: usize,
pub(super) num_batches: usize,
pub(super) knn: usize,
pub(super) opt_iter: usize,
pub(super) refine_params: &'a crate::alg::refine_multilevel::RefineParams,
pub(super) output_calibration: legume_numeric::param::traits::CalibrateTarget,
pub(super) anchor_batches: Option<&'a [usize]>,
pub(super) summary_batches: Option<&'a [usize]>,
pub(super) bulk_batches: Option<&'a [usize]>,
pub(super) observe_panels: bool,
pub(super) keep_finest_stats: bool,
pub(super) pb_tree: Option<&'a PbTree>,
pub(super) cell_to_stratum: Option<&'a [usize]>,
pub(super) exclude_unmatched_from_delta: bool,
pub(super) strata_bits: usize,
}
pub(super) fn refine_and_collect_single_layer(
data_vec: &mut SparseIoVec,
proj_kn: &DMatrix<f32>,
ctx: &RefineCollectCtx<'_>,
) -> anyhow::Result<MultilevelCollapseOut> {
let RefineCollectCtx {
fine_codes,
group_to_cols_finest: _,
level_dims,
num_features,
num_batches,
knn,
opt_iter,
refine_params,
output_calibration,
anchor_batches: _,
summary_batches: _,
bulk_batches: _,
observe_panels: _,
keep_finest_stats: _,
pb_tree: _,
cell_to_stratum: _,
exclude_unmatched_from_delta: _,
strata_bits,
} = *ctx;
info!(
"Multi-level refinement path (BBKNN + DC-SBM): {} levels",
level_dims.len()
);
let pb_samples = build_pb_samples(
data_vec,
proj_kn,
num_features,
ctx.summary_batches.unwrap_or(&[]),
ctx.bulk_batches.unwrap_or(&[]),
ctx.cell_to_stratum,
)?;
let num_pb = pb_samples.layout.cell_counts.len();
let ncells_dbg = proj_kn.ncols();
info!(
"Built {} pb-samples from {} cells (ratio {:.2}; knn={})",
num_pb,
ncells_dbg,
num_pb as f32 / ncells_dbg.max(1) as f32,
knn
);
if num_pb as f32 > 0.8 * ncells_dbg as f32 {
warn!(
"pb-sample count ({}) is close to cell count ({}) — hash partition is too fine \
(many 1-cell pb-samples). Consider lowering --sort-dim.",
num_pb, ncells_dbg
);
}
let ncols = proj_kn.ncols();
let pb_sample_to_cells = build_pb_sample_to_cells(&pb_samples.layout);
let initial_per_level =
initial_per_level_from_hash(fine_codes, &pb_sample_to_cells, level_dims, strata_bits);
let empty: [ColumnDict<usize>; 0] = [];
let batch_knn: &[ColumnDict<usize>] = if num_batches >= 2 {
data_vec
.batch_knn_lookup()
.ok_or_else(|| anyhow::anyhow!("batch_knn_lookup not built"))?
.as_slice()
} else {
&empty
};
let reproject_offsets =
build_reproject_offsets(fine_codes, &pb_sample_to_cells, level_dims, strata_bits);
let inputs = crate::alg::refine_multilevel::RefineInputs {
layout: &pb_samples.layout,
gene_sums: &pb_samples.gene_sums,
num_genes: num_features,
pb_sample_to_cells: &pb_sample_to_cells,
batch_knn_lookup: batch_knn,
k_per_batch: knn,
initial_sc_to_group_per_level: &initial_per_level,
reproject_offsets_per_level: &reproject_offsets,
};
let refined = split_anchored_finest_groups(
refine_or_identity(num_batches >= 2, &inputs, refine_params)?,
&pb_samples.layout,
);
{
let n_leaves = initial_per_level
.first()
.map(|lvl| lvl.iter().copied().max().map_or(0, |m| m + 1))
.unwrap_or(0);
let finest = &refined.pbsamp_to_group[0];
let k_fin = refined.num_groups_per_level[0];
let mut batch_mask = vec![0u128; k_fin];
let b2g = &pb_samples.layout.pb_sample_to_batch;
for (pb, &g) in finest.iter().enumerate() {
let b = b2g[pb];
if b < 128 {
batch_mask[g] |= 1u128 << b;
}
}
let spans: Vec<u32> = batch_mask.iter().map(|m| m.count_ones()).collect();
let multi = spans.iter().filter(|&&c| c > 1).count();
let max_b = spans.iter().copied().max().unwrap_or(0);
let mean_b = spans.iter().map(|&c| c as f64).sum::<f64>() / k_fin.max(1) as f64;
info!(
"collapse structure: {} pb-samples, {} batches, {} leaf codes (finest init), \
{} refined finest groups; finest groups spanning >1 batch: {}/{} \
(max {} batches/group, mean {:.2})",
num_pb, num_batches, n_leaves, k_fin, multi, k_fin, max_b, mean_b
);
}
let num_levels = level_dims.len();
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);
let nthreads = rayon::current_num_threads();
info!(
"Assigning {} cells to {} finest pb-sample groups ({} rayon threads) ...",
ncols, k_finest, nthreads
);
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 = ctx.exclude_unmatched_from_delta;
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,
ctx.anchor_batches,
&mut fine_stat,
)?;
}
info!(
"Level 1/{}: refined k={} (finest; {} cells read)",
num_levels, k_finest, ncols
);
if ctx.observe_panels {
attach_observability(&mut fine_stat, data_vec)?;
}
let mut results: Vec<CollapsedOut> = Vec::with_capacity(num_levels);
let finest_out = optimize(
&fine_stat,
(1.0, 1.0),
opt_iter,
&format!("Fit L1/{}", num_levels),
output_calibration,
ctx.keep_finest_stats,
)?;
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 {}/{}: refined 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!("Fit L{}/{}", level + 1, num_levels),
output_calibration,
false,
)?;
results.push(out);
prev_stat = coarse_stat;
}
info!(
"Fitted pseudobulk posteriors for {} refined levels: k = {:?} (finest first)",
num_levels, refined.num_groups_per_level
);
let cell_to_pb_per_level = cell_to_pb_per_level(&refined, &pb_sample_to_cells, ncols);
let pb_tree = ctx.pb_tree.map(|tree| {
let mut tree = tree.clone();
let mut leaf_pbs: HashMap<usize, std::collections::BTreeSet<usize>> = HashMap::default();
if let Some(finest) = cell_to_pb_per_level.first() {
for (c, &code) in fine_codes.iter().enumerate() {
leaf_pbs.entry(code).or_default().insert(finest[c]);
}
}
let mut leaves: Vec<(usize, Vec<usize>)> = leaf_pbs
.into_iter()
.map(|(code, pbs)| (code, pbs.into_iter().collect()))
.collect();
leaves.sort_by_key(|x| x.0);
tree.leaf_to_finest_pb = leaves;
tree
});
Ok(MultilevelCollapseOut {
levels: results,
cell_to_pb_per_level,
pb_tree,
})
}
fn cell_to_pb_per_level(
refined: &crate::alg::refine_multilevel::RefinedAssignment,
pb_sample_to_cells: &[Vec<usize>],
ncols: usize,
) -> Vec<Vec<usize>> {
refined
.pbsamp_to_group
.iter()
.map(|groups| {
let mut c2g = vec![0usize; ncols];
for (pbsamp, cells) in pb_sample_to_cells.iter().enumerate() {
for &c in cells {
c2g[c] = groups[pbsamp];
}
}
c2g
})
.collect()
}
pub struct StackCollapseOut {
pub levels: Vec<Vec<CollapsedOut>>,
pub cell_to_pb_per_level: Vec<Vec<usize>>,
}
pub(super) fn refine_and_collect_stack(
stack: &mut SparseIoStack,
proj_kn: &DMatrix<f32>,
ctx: &RefineCollectCtx<'_>,
) -> anyhow::Result<StackCollapseOut> {
let RefineCollectCtx {
fine_codes,
group_to_cols_finest,
level_dims,
num_features: _,
num_batches,
knn,
opt_iter,
refine_params,
output_calibration,
anchor_batches: _,
summary_batches: _,
bulk_batches: _,
observe_panels: _,
keep_finest_stats,
pb_tree: _,
cell_to_stratum,
exclude_unmatched_from_delta,
strata_bits,
} = *ctx;
let num_layers = stack.num_types();
info!(
"Multi-level stack refinement (BBKNN + DC-SBM): {} layers × {} levels",
num_layers,
level_dims.len()
);
let ncols = proj_kn.ncols();
let col_to_batch: Vec<usize> = stack.stack[0].get_batch_membership(0..ncols);
let layout = build_pb_sample_layout(
group_to_cols_finest,
&col_to_batch,
proj_kn,
None,
&[],
&[],
cell_to_stratum,
)?;
let num_pb = layout.cell_counts.len();
let owner_num_features = stack.stack[0].num_rows();
let gene_sums_owner = collect_pb_sample_gene_sums(
&stack.stack[0],
group_to_cols_finest,
&layout.cell_to_pbsamp,
num_pb,
)?;
let pb_sample_to_cells = build_pb_sample_to_cells(&layout);
let initial_per_level =
initial_per_level_from_hash(fine_codes, &pb_sample_to_cells, level_dims, strata_bits);
let empty: [ColumnDict<usize>; 0] = [];
let batch_knn: &[ColumnDict<usize>] = if num_batches >= 2 {
stack.stack[0]
.batch_knn_lookup()
.ok_or_else(|| anyhow::anyhow!("batch_knn_lookup not built"))?
.as_slice()
} else {
&empty
};
let reproject_offsets =
build_reproject_offsets(fine_codes, &pb_sample_to_cells, level_dims, strata_bits);
let inputs = crate::alg::refine_multilevel::RefineInputs {
layout: &layout,
gene_sums: &gene_sums_owner,
num_genes: owner_num_features,
pb_sample_to_cells: &pb_sample_to_cells,
batch_knn_lookup: batch_knn,
k_per_batch: knn,
initial_sc_to_group_per_level: &initial_per_level,
reproject_offsets_per_level: &reproject_offsets,
};
let refined = refine_or_identity(num_batches >= 2, &inputs, refine_params)?;
let mut per_layer_gene_sums: Vec<GeneSums> = Vec::with_capacity(num_layers);
for (d, layer) in stack.stack.iter().enumerate() {
if d == 0 {
per_layer_gene_sums.push(gene_sums_owner.clone());
} else {
per_layer_gene_sums.push(collect_pb_sample_gene_sums(
layer,
group_to_cols_finest,
&layout.cell_to_pbsamp,
num_pb,
)?);
}
}
let num_levels = level_dims.len();
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);
let nthreads = rayon::current_num_threads();
info!(
"Assigning {} cells to {} finest pb-sample groups across {} layers ({} rayon threads) ...",
ncols, k_finest, num_layers, nthreads
);
for layer in stack.stack.iter_mut() {
layer.assign_groups(&finest_str, None);
}
let mut fine_stats: Vec<CollapsedStat> = Vec::with_capacity(num_layers);
let mut finest_layer_results = Vec::with_capacity(num_layers);
for (d, layer) in stack.stack.iter().enumerate() {
let num_features = layer.num_rows();
let mut stat = CollapsedStat::new(num_features, k_finest, num_batches);
stat.exclude_unmatched_from_delta = exclude_unmatched_from_delta;
info!(
"Layer {}/{}: collecting basic stats over {} groups ...",
d + 1,
num_layers,
k_finest
);
layer.collect_basic_stat(&mut stat)?;
if num_batches >= 2 {
info!(
"Layer {}/{}: collecting per-batch stats ({} batches) ...",
d + 1,
num_layers,
num_batches
);
layer.collect_batch_stat(&mut stat)?;
let batch_knn = layer
.batch_knn_lookup()
.ok_or_else(|| anyhow::anyhow!("batch_knn_lookup not built"))?;
info!(
"Layer {}/{}: collecting cross-batch matched stats (knn={}) over {} pb-samples ...",
d + 1,
num_layers,
knn,
num_pb
);
collect_matched_stat_coarse(
&layout,
&per_layer_gene_sums[d],
&refined.pbsamp_to_group[0],
batch_knn.as_slice(),
knn,
ctx.anchor_batches,
&mut stat,
)?;
}
let out = optimize(
&stat,
(1.0, 1.0),
opt_iter,
&format!("Fit L1/{} layer {}/{}", num_levels, d + 1, num_layers),
output_calibration,
keep_finest_stats,
)?;
finest_layer_results.push(out);
fine_stats.push(stat);
}
info!(
"Level 1/{}: refined k={} (finest; {} layers × {} cells)",
num_levels, k_finest, num_layers, ncols
);
let mut results: Vec<Vec<CollapsedOut>> = Vec::with_capacity(num_levels);
results.push(finest_layer_results);
let mut prev_stats = fine_stats;
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 level_opt_iter = (opt_iter / 2).max(10);
let mut layer_results = Vec::with_capacity(num_layers);
let mut coarse_stats = Vec::with_capacity(num_layers);
for (d, prev_stat) in prev_stats.iter().enumerate() {
let coarse_stat = merge_stat(prev_stat, &fine_to_coarse, k_level);
let out = optimize(
&coarse_stat,
(1.0, 1.0),
level_opt_iter,
&format!(
"Fit L{}/{} layer {}/{}",
level + 1,
num_levels,
d + 1,
num_layers
),
output_calibration,
false,
)?;
layer_results.push(out);
coarse_stats.push(coarse_stat);
}
info!(
"Level {}/{}: refined k={} (merged from {}, {} layers)",
level + 1,
num_levels,
k_level,
k_prev,
num_layers
);
results.push(layer_results);
prev_stats = coarse_stats;
}
Ok(StackCollapseOut {
levels: results,
cell_to_pb_per_level: cell_to_pb_per_level(&refined, &pb_sample_to_cells, ncols),
})
}
pub(super) fn compute_level_sort_dims(finest_sort_dim: usize, num_levels: usize) -> Vec<usize> {
if num_levels <= 1 {
return vec![finest_sort_dim];
}
let coarsest = DEFAULT_COARSEST_SORT_DIM.min(finest_sort_dim);
let mut dims = Vec::with_capacity(num_levels);
for level in 0..num_levels {
let t = level as f32 / (num_levels - 1) as f32;
let dim = finest_sort_dim as f32 - t * (finest_sort_dim - coarsest) as f32;
let dim = dim.round() as usize;
if dims.last() != Some(&dim) {
dims.push(dim);
}
}
dims
}
#[cfg(test)]
#[path = "refine_tests.rs"]
mod reproject_tests;