use super::*;
use legume_numeric::param::traits::Inference;
use nalgebra::{DMatrix, DVector};
fn toy_stat(num_genes: usize, num_samples: usize, num_batches: usize) -> CollapsedStat {
let mut s = CollapsedStat::new(num_genes, num_samples, num_batches);
let f = |a: usize, b: usize| -> f32 { 1.0 + ((a * 7 + b * 13) % 11) as f32 };
s.observed_sum_ds = DMatrix::from_fn(num_genes, num_samples, &f);
s.imputed_sum_ds = DMatrix::from_fn(num_genes, num_samples, |g, c| 0.5 * f(g + 1, c + 2));
s.size_s = DVector::from_fn(num_samples, |c, _| 2.0 + (c % 3) as f32);
s.observed_sum_db = DMatrix::from_fn(num_genes, num_batches, |g, b| f(g, b) + 0.7);
s.n_bs = DMatrix::from_fn(num_batches, num_samples, |b, c| 1.0 + ((b + c) % 4) as f32);
s.matched_bs = DMatrix::from_fn(num_batches, num_samples, |b, c| {
0.5 + ((b + 2 * c) % 3) as f32
});
s
}
fn two_pure_pbs(a: &[f32], b: &[f32], n: f32) -> CollapsedStat {
let g = a.len();
let mut s = CollapsedStat::new(g, 2, 2);
for i in 0..g {
s.observed_sum_ds[(i, 0)] = a[i] * n;
s.imputed_sum_ds[(i, 0)] = b[i] * n;
s.observed_sum_ds[(i, 1)] = b[i] * n;
s.imputed_sum_ds[(i, 1)] = a[i] * n;
s.observed_sum_db[(i, 0)] = a[i] * n;
s.observed_sum_db[(i, 1)] = b[i] * n;
}
s.size_s = DVector::from_element(2, n);
s.n_bs = DMatrix::from_row_slice(2, 2, &[n, 0.0, 0.0, n]);
s.matched_bs = DMatrix::from_row_slice(2, 2, &[0.0, n, n, 0.0]);
s
}
fn assert_rel(got: f32, want: f32, tol: f32, tag: &str) {
assert!(
(got / want - 1.0).abs() < tol,
"{tag}: got {got}, want {want} (rel tol {tol})"
);
}
#[test]
#[allow(clippy::needless_range_loop)] fn pooled_two_pure_batches_share_one_frame() {
let a = [8.0f32, 2.0, 12.0, 1.0];
let b = [2.0f32, 8.0, 3.0, 4.0];
let stat = two_pure_pbs(&a, &b, 2000.0);
let out = optimize_block(&stat, (1.0, 1.0), 60, CalibrateTarget::All, None).unwrap();
let mu = out.mu_adjusted.as_ref().unwrap().posterior_mean();
let delta = out.delta.as_ref().unwrap().posterior_mean();
let resid = out.mu_residual.as_ref().unwrap().posterior_mean();
let gamma = out.gamma.as_ref().unwrap().posterior_mean();
for g in 0..a.len() {
let common = (a[g] * b[g]).sqrt();
assert_rel(mu[(g, 0)], common, 0.02, &format!("gene {g} mu pb0"));
assert_rel(mu[(g, 1)], common, 0.02, &format!("gene {g} mu pb1"));
assert_rel(
delta[(g, 0)] / delta[(g, 1)],
a[g] / b[g],
0.02,
&format!("gene {g} delta ratio"),
);
assert_rel(
delta[(g, 0)] * delta[(g, 1)],
1.0,
0.02,
&format!("gene {g} delta geometric mean"),
);
assert_rel(
resid[(g, 0)],
delta[(g, 0)],
0.02,
&format!("gene {g} mu_residual pb0 = own delta"),
);
assert_rel(
resid[(g, 1)],
delta[(g, 1)],
0.02,
&format!("gene {g} mu_residual pb1 = own delta"),
);
assert_rel(
gamma[(g, 0)],
delta[(g, 1)],
0.02,
&format!("gene {g} gamma pb0 = source delta"),
);
}
}
#[test]
fn anchored_frame_is_the_anchor_batch() {
let a = [8.0f32, 2.0, 12.0, 1.0];
let b = [2.0f32, 8.0, 3.0, 4.0];
let n = 2000.0;
let mut stat = two_pure_pbs(&a, &b, n);
for (g, &bg) in b.iter().enumerate() {
stat.imputed_sum_ds[(g, 1)] = bg * n;
}
stat.matched_bs = DMatrix::from_row_slice(2, 2, &[0.0, 0.0, n, n]);
stat.anchor_batches = vec![1];
let out = optimize_block(&stat, (1.0, 1.0), 60, CalibrateTarget::All, None).unwrap();
let mu = out.mu_adjusted.as_ref().unwrap().posterior_mean();
let delta = out.delta.as_ref().unwrap().posterior_mean();
for g in 0..a.len() {
assert_rel(delta[(g, 1)], 1.0, 0.02, &format!("gene {g} anchor delta"));
assert_rel(
delta[(g, 0)],
a[g] / b[g],
0.03,
&format!("gene {g} new-batch delta"),
);
assert_rel(
mu[(g, 0)],
b[g],
0.03,
&format!("gene {g} mu pb0 in the anchor frame"),
);
assert_rel(mu[(g, 1)], b[g], 0.02, &format!("gene {g} mu pb1"));
}
}
#[test]
fn merge_stat_sums_matched_mass_and_keeps_anchors() {
let mut fine = toy_stat(3, 4, 2);
fine.anchor_batches = vec![0];
let coarse = merge_stat(&fine, &[0, 1, 0, 1], 2);
for b in 0..2 {
assert_eq!(
coarse.matched_bs[(b, 0)],
fine.matched_bs[(b, 0)] + fine.matched_bs[(b, 2)]
);
assert_eq!(
coarse.matched_bs[(b, 1)],
fine.matched_bs[(b, 1)] + fine.matched_bs[(b, 3)]
);
}
assert_eq!(coarse.anchor_batches, vec![0]);
let sub = fine.select_rows(1, 2);
assert_eq!(sub.matched_bs, fine.matched_bs);
assert_eq!(sub.anchor_batches, vec![0]);
let cols = fine.select_columns(&[3, 1]);
assert_eq!(cols.matched_bs.column(0), fine.matched_bs.column(3));
assert_eq!(cols.anchor_batches, vec![0]);
}
fn assert_mat_close(a: &DMatrix<f32>, b: &DMatrix<f32>, tag: &str) {
assert_eq!(a.shape(), b.shape(), "{tag}: shape mismatch");
for (x, y) in a.iter().zip(b.iter()) {
assert!(
(x - y).abs() <= 1e-5 * (1.0 + x.abs().max(y.abs())),
"{tag}: {x} vs {y}"
);
}
}
#[test]
fn blocked_optimize_matches_single_block() {
let (g, k, b) = (10usize, 4usize, 2usize);
let stat = toy_stat(g, k, b);
let hyper = (1.0, 1.0);
let iters = 25;
let full = optimize_block(&stat, hyper, iters, CalibrateTarget::All, None).unwrap();
let ranges = [(0usize, 3usize), (3, 4), (7, 3)];
let mut mu_obs = Vec::new();
let mut mu_adj = Vec::new();
let mut mu_res = Vec::new();
let mut gam = Vec::new();
let mut del = Vec::new();
for (r0, nr) in ranges {
let sub = stat.select_rows(r0, nr);
let out = optimize_block(&sub, hyper, iters, CalibrateTarget::All, None).unwrap();
mu_obs.push(out.mu_observed);
mu_adj.push(out.mu_adjusted.unwrap());
mu_res.push(out.mu_residual.unwrap());
gam.push(out.gamma.unwrap());
del.push(out.delta.unwrap());
}
let blk_obs = GammaMatrix::vconcat(mu_obs, true);
let blk_adj = GammaMatrix::vconcat(mu_adj, true);
let blk_res = GammaMatrix::vconcat(mu_res, true);
let blk_gam = GammaMatrix::vconcat(gam, true);
let blk_del = GammaMatrix::vconcat(del, true);
assert_mat_close(
full.mu_observed.posterior_mean(),
blk_obs.posterior_mean(),
"mu_obs mean",
);
assert_mat_close(
full.mu_adjusted.as_ref().unwrap().posterior_mean(),
blk_adj.posterior_mean(),
"mu_adj mean",
);
assert_mat_close(
full.mu_residual.as_ref().unwrap().posterior_mean(),
blk_res.posterior_mean(),
"mu_resid mean",
);
assert_mat_close(
full.gamma.as_ref().unwrap().posterior_mean(),
blk_gam.posterior_mean(),
"gamma mean",
);
assert_mat_close(
full.delta.as_ref().unwrap().posterior_mean(),
blk_del.posterior_mean(),
"delta mean",
);
assert_mat_close(
full.mu_adjusted.as_ref().unwrap().posterior_log_mean(),
blk_adj.posterior_log_mean(),
"mu_adj log_mean",
);
}
#[test]
fn mean_only_sparsifies_unobserved_cells() {
let (g, k, b) = (4usize, 3usize, 2usize);
let mut stat = CollapsedStat::new(g, k, b);
stat.observed_sum_ds[(0, 0)] = 5.0; stat.observed_sum_ds[(1, 1)] = 3.0;
stat.imputed_sum_ds[(2, 2)] = 2.0; stat.size_s = DVector::from_element(k, 10.0);
stat.n_bs = DMatrix::from_element(b, k, 5.0);
stat.observed_sum_db.fill(1.0);
let out = optimize_block(&stat, (1.0, 1.0), 10, CalibrateTarget::MeanOnly, None).unwrap();
let adj = out.mu_adjusted.unwrap();
let m = adj.posterior_mean();
assert!(m[(0, 0)] > 0.0);
assert!(m[(1, 1)] > 0.0);
assert!(m[(2, 2)] > 0.0);
assert_eq!(m[(3, 0)], 0.0);
assert_eq!(m[(0, 1)], 0.0);
let out_all = optimize_block(&stat, (1.0, 1.0), 10, CalibrateTarget::All, None).unwrap();
let ma = out_all.mu_adjusted.unwrap();
assert!(
ma.posterior_mean()[(3, 0)] > 0.0,
"All path must keep the prior baseline"
);
}
#[test]
fn mean_only_vconcat_drops_stats_keeps_means() {
let (g, k, b) = (9usize, 3usize, 2usize);
let stat = toy_stat(g, k, b);
let hyper = (1.0, 1.0);
let iters = 20;
let reference = optimize_block(&stat, hyper, iters, CalibrateTarget::All, None).unwrap();
let mut blocks = Vec::new();
for (r0, nr) in [(0usize, 5usize), (5, 4)] {
let sub = stat.select_rows(r0, nr);
let mut out = optimize_block(&sub, hyper, iters, CalibrateTarget::MeanOnly, None).unwrap();
out.release_stats();
blocks.push(out.mu_observed);
}
let assembled = GammaMatrix::vconcat(blocks, false);
assert_mat_close(
reference.mu_observed.posterior_mean(),
assembled.posterior_mean(),
"mu_obs mean (mean-only)",
);
assert_eq!(
assembled.posterior_sd().nrows(),
0,
"sd should be empty under MeanOnly"
);
}
#[test]
fn observed_counts_requires_kept_stats() {
let stat = toy_stat(4, 2, 1);
let mut out = optimize(
&stat,
(1.0, 1.0),
5,
"test",
CalibrateTarget::MeanOnly,
false,
)
.unwrap();
assert!(!out.stats_kept);
let err = out
.observed_counts(&[0, 1])
.expect_err("released stats must error");
assert!(
err.to_string().contains("keep_finest_stats"),
"unexpected: {err}"
);
out = optimize(
&stat,
(1.0, 1.0),
5,
"test",
CalibrateTarget::MeanOnly,
true,
)
.unwrap();
assert!(out.stats_kept);
let err = out
.observed_counts(&[0, 99])
.expect_err("OOB membership must error");
assert!(
err.to_string().contains("out of range"),
"unexpected: {err}"
);
let (counts, sizes) = out.observed_counts(&[0, 0, 1]).unwrap();
assert_eq!(counts.nrows(), 4);
assert_eq!(sizes, vec![2.0, 1.0]);
}