use crate::sim::copula::marginals::NbFit;
use crate::sparse_io::*;
use legume_numeric::matrix::common_io::file_ext;
use legume_numeric::matrix::sparse_stat::SparseRunningStatistics;
use nalgebra::DMatrix;
pub type SparseRef = Box<dyn SparseIo<IndexIter = Vec<usize>>>;
pub fn open_reference(path: &str) -> anyhow::Result<SparseRef> {
let ext = file_ext(path)?.to_string();
let backend = match ext.as_str() {
"h5" => SparseIoBackend::HDF5,
"zarr" | "zip" => SparseIoBackend::Zarr,
other => anyhow::bail!(
"unsupported reference extension '{}': expected h5 / zarr / zarr.zip",
other
),
};
open_sparse_matrix(path, &backend)
}
#[derive(Debug, Clone, Copy)]
pub struct GeneStats {
pub mu: f64,
pub var: f64,
pub nnz: usize,
}
pub fn per_gene_stats(
sc: &SparseRef,
cells: &[usize],
n_genes: usize,
) -> anyhow::Result<Vec<GeneStats>> {
if cells.is_empty() {
return Ok(vec![
GeneStats {
mu: 0.0,
var: 0.0,
nnz: 0,
};
n_genes
]);
}
let mat = sc.read_columns_csc(cells.to_vec())?;
let mut acc = SparseRunningStatistics::<f64>::new(n_genes);
for col in mat.col_iter() {
let vals_f64: Vec<f64> = col.values().iter().map(|&v| v as f64).collect();
acc.add_sparse_column(col.row_indices(), &vals_f64);
}
for _ in mat.ncols()..cells.len() {
acc.add_sparse_column(&[], &[]);
}
let (npos, _sum, mean, std) = acc.to_vecs();
Ok((0..n_genes)
.map(|g| GeneStats {
mu: mean[g],
var: std[g] * std[g],
nnz: npos[g] as usize,
})
.collect())
}
pub fn per_gene_stats_and_marginals(
sc: &SparseRef,
cells: &[usize],
n_genes: usize,
r_floor: f32,
) -> anyhow::Result<(Vec<GeneStats>, Vec<NbFit>)> {
let stats = per_gene_stats(sc, cells, n_genes)?;
let marginals = stats
.iter()
.map(|s| nb_fit_from_stats(s, r_floor))
.collect();
Ok((stats, marginals))
}
fn nb_fit_from_stats(s: &GeneStats, r_floor: f32) -> NbFit {
let mu = s.mu as f32;
let var = s.var as f32;
if mu <= 0.0 || var <= mu {
return NbFit {
mu,
r: f32::INFINITY,
};
}
let r = (mu * mu) / (var - mu);
NbFit {
mu,
r: r.max(r_floor),
}
}
pub fn select_hvg(stats: &[GeneStats], n_hvg: usize) -> Vec<usize> {
let means: Vec<f32> = stats.iter().map(|s| s.mu as f32).collect();
let vars: Vec<f32> = stats.iter().map(|s| s.var as f32).collect();
crate::alg::hvg::select_hvg_by_stats(&means, &vars, n_hvg)
}
pub fn build_z_matrix(
sc: &SparseRef,
cells: &[usize],
hvg: &[usize],
fits: &[NbFit],
) -> anyhow::Result<DMatrix<f32>> {
use crate::sim::copula::marginals::{inv_phi, nb_cdf_table, pit_continuity};
let n_hvg = hvg.len();
let n = cells.len();
let mat = sc.read_rows_csc(hvg.to_vec())?;
let col_off = mat.col_offsets();
let row_idx = mat.row_indices();
let vals = mat.values();
let mut max_val: Vec<u32> = vec![0; n_hvg];
for k in 0..vals.len() {
let h = row_idx[k];
let v = vals[k] as u32;
if v > max_val[h] {
max_val[h] = v;
}
}
let cdf_tables: Vec<Vec<f64>> = (0..n_hvg)
.map(|h| {
let fit = fits[hvg[h]];
nb_cdf_table(fit, (max_val[h] as usize).max(1))
})
.collect();
let u_zero: Vec<f64> = cdf_tables.iter().map(|t| pit_continuity(t, 0)).collect();
let z_zero: Vec<f32> = u_zero.iter().map(|&u| inv_phi(u) as f32).collect();
let mut z = DMatrix::<f32>::zeros(n_hvg, n);
for (col_out, &cell) in cells.iter().enumerate() {
for h in 0..n_hvg {
z[(h, col_out)] = z_zero[h];
}
let s = col_off[cell];
let e = col_off[cell + 1];
for k in s..e {
let h = row_idx[k];
let v = vals[k] as u32;
let u = pit_continuity(&cdf_tables[h], v);
z[(h, col_out)] = inv_phi(u) as f32;
}
}
Ok(z)
}