pub mod gaussian;
pub mod marginals;
pub mod reference;
use gaussian::CopulaCovariance;
use log::info;
use reference::SparseRef;
const MU_HAT_ACTIVE_THRESHOLD: f32 = 1e-5;
pub struct GlobalCopulaFit {
pub gene_names: Vec<Box<str>>,
pub n_genes: usize,
pub hvg_indices: Vec<usize>,
pub mu_hat: Vec<f32>,
pub r_hat: Vec<f32>,
pub active_genes: Vec<usize>,
pub hvg_pos: Vec<Option<u32>>,
pub copula: CopulaCovariance,
}
pub struct GlobalCopulaArgs<'a> {
pub sc: &'a SparseRef,
pub n_hvg: usize,
pub copula_rank: usize,
pub regularization: f32,
pub r_floor: f32,
}
pub fn fit_global_copula(args: &GlobalCopulaArgs) -> anyhow::Result<GlobalCopulaFit> {
let n_genes = args
.sc
.num_rows()
.ok_or_else(|| anyhow::anyhow!("reference has no num_rows"))?;
let n_cells = args
.sc
.num_columns()
.ok_or_else(|| anyhow::anyhow!("reference has no num_columns"))?;
if n_cells < 2 {
anyhow::bail!(
"reference has only {} cells; need ≥2 to fit a copula",
n_cells
);
}
let gene_names = args.sc.row_names()?;
info!("reference: {} genes × {} cells", n_genes, n_cells);
let cells: Vec<usize> = (0..n_cells).collect();
let (stats, marginals) =
reference::per_gene_stats_and_marginals(args.sc, &cells, n_genes, args.r_floor)?;
let hvg_indices = reference::select_hvg(&stats, args.n_hvg);
let z = reference::build_z_matrix(args.sc, &cells, &hvg_indices, &marginals)?;
let copula = CopulaCovariance::fit(&z, args.copula_rank, args.regularization)?;
let mean_ridge = copula.ridge_sd.mean();
info!(
"fit global copula: {} HVGs, rank {} (cap {}) + per-row ridge sd (mean {:.3})",
hvg_indices.len(),
copula.rank(),
args.copula_rank,
mean_ridge
);
let r_hat: Vec<f32> = marginals.iter().map(|f| f.r).collect();
let mu_hat: Vec<f32> = stats.iter().map(|s| s.mu as f32).collect();
let active_genes: Vec<usize> = mu_hat
.iter()
.enumerate()
.filter_map(|(g, &m)| (m >= MU_HAT_ACTIVE_THRESHOLD).then_some(g))
.collect();
info!(
"active genes (μ̂ ≥ {:.0e}): {}/{} ({:.1}% — undetectable genes skipped from sampling)",
MU_HAT_ACTIVE_THRESHOLD,
active_genes.len(),
n_genes,
100.0 * active_genes.len() as f32 / n_genes as f32
);
let mut hvg_pos = vec![None; n_genes];
for (h, &g) in hvg_indices.iter().enumerate() {
hvg_pos[g] = Some(h as u32);
}
Ok(GlobalCopulaFit {
gene_names,
n_genes,
hvg_indices,
mu_hat,
r_hat,
active_genes,
hvg_pos,
copula,
})
}