#![allow(dead_code)]
use crate::alg::batch_mixing::*;
use crate::sparse_data_visitors::*;
use crate::sparse_io_stack::SparseIoStack;
use crate::sparse_io_vector::SparseIoVec;
use legume_numeric::matrix::dmatrix_util::*;
use legume_numeric::matrix::rand_util::mix_seed;
use std::sync::{Arc, Mutex};
use legume_numeric::matrix::traits::*;
use log::{info, warn};
use nalgebra::DVector;
struct ProjVisitorIn<'a> {
basis_kd: &'a nalgebra::DMatrix<f32>,
}
pub struct RandColProjOut {
pub basis: nalgebra::DMatrix<f32>,
pub proj: nalgebra::DMatrix<f32>,
}
pub struct RandRowProjOut {
pub basis: nalgebra::DMatrix<f32>,
pub proj: nalgebra::DMatrix<f32>,
}
pub const DEFAULT_PROJECTION_SEED: u64 = 0x50524F4A_50524F4A;
pub trait RandProjOps {
fn project_columns(
&self,
target_dim: usize,
block_size: Option<usize>,
) -> anyhow::Result<RandColProjOut> {
self.project_columns_with_batch_correction::<usize>(target_dim, block_size, None)
}
fn project_columns_with_batch_correction<T>(
&self,
target_dim: usize,
block_size: Option<usize>,
batch_membership: Option<&[T]>,
) -> anyhow::Result<RandColProjOut>
where
T: Sync + Send + std::hash::Hash + Eq + Clone + ToString,
{
self.project_columns_with_batch_correction_seeded(
target_dim,
block_size,
batch_membership,
DEFAULT_PROJECTION_SEED,
)
}
fn project_columns_with_batch_correction_seeded<T>(
&self,
target_dim: usize,
block_size: Option<usize>,
batch_membership: Option<&[T]>,
seed: u64,
) -> anyhow::Result<RandColProjOut>
where
T: Sync + Send + std::hash::Hash + Eq + Clone + ToString;
fn project_columns_weighted<T>(
&self,
target_dim: usize,
block_size: Option<usize>,
batch_membership: Option<&[T]>,
row_weights: &[f32],
) -> anyhow::Result<RandColProjOut>
where
T: Sync + Send + std::hash::Hash + Eq + Clone + ToString,
{
self.project_columns_weighted_seeded(
target_dim,
block_size,
batch_membership,
row_weights,
DEFAULT_PROJECTION_SEED,
)
}
fn project_columns_weighted_seeded<T>(
&self,
target_dim: usize,
block_size: Option<usize>,
batch_membership: Option<&[T]>,
row_weights: &[f32],
seed: u64,
) -> anyhow::Result<RandColProjOut>
where
T: Sync + Send + std::hash::Hash + Eq + Clone + ToString;
fn partition_columns_to_groups(
&mut self,
proj_kn: &nalgebra::DMatrix<f32>,
num_features: Option<usize>,
ncols_per_group: Option<usize>,
) -> anyhow::Result<usize>;
fn partition_columns_to_mixed_groups<T>(
&mut self,
proj_kn: &nalgebra::DMatrix<f32>,
num_features: Option<usize>,
batch_membership: &[T],
min_batches: usize,
merge_levels: usize,
) -> anyhow::Result<usize>
where
T: std::hash::Hash + Eq + Clone,
{
let nn = proj_kn.ncols();
anyhow::ensure!(
batch_membership.len() == nn,
"batch membership size {} mismatches the number of columns {}",
batch_membership.len(),
nn
);
let kk = proj_kn
.nrows()
.min(num_features.unwrap_or(proj_kn.nrows()))
.min(nn);
let batch = batch_indices(batch_membership);
let n_batches = batch.iter().max().map_or(0, |&m| m + 1);
let codes = binary_sort_columns(proj_kn, kk)?;
let labels =
merge_poorly_mixed_bins(&codes, &batch, kk, merge_levels, min_batches.min(n_batches));
self.assign_group_labels(&labels);
let n_groups = labels
.iter()
.collect::<std::collections::HashSet<_>>()
.len();
info!(
"partitioned columns into {} groups ({} bits, bins with fewer than {} batches merged up to {} levels)",
n_groups,
kk,
min_batches.min(n_batches),
merge_levels
);
Ok(n_groups)
}
fn assign_group_labels(&mut self, labels: &[usize]);
}
fn project_columns_visitor(
job: (usize, usize),
data_vec: &SparseIoVec,
shared_in: &ProjVisitorIn<'_>,
arc_proj_kn: Arc<Mutex<&mut nalgebra::DMatrix<f32>>>,
) -> anyhow::Result<()> {
let (lb, ub) = job;
let basis_kd = shared_in.basis_kd;
let target_dim = basis_kd.nrows();
let n_block = ub - lb;
let mut xx_dm = data_vec.read_columns_csc(lb..ub)?;
for x in xx_dm.values_mut() {
*x = x.ln_1p();
}
xx_dm.normalize_columns_inplace();
let mut chunk = nalgebra::DMatrix::<f32>::zeros(target_dim, n_block);
for (j, col) in xx_dm.col_iter().enumerate() {
for (&i, &v) in col.row_indices().iter().zip(col.values()) {
chunk.column_mut(j).axpy(v, &basis_kd.column(i), 1.0);
}
}
let mut proj_kn = arc_proj_kn.lock().expect("proj_kn lock");
proj_kn.columns_range_mut(lb..ub).copy_from(&chunk);
Ok(())
}
impl RandProjOps for SparseIoStack {
fn project_columns_with_batch_correction_seeded<T>(
&self,
target_dim: usize,
block_size: Option<usize>,
batch_membership: Option<&[T]>,
seed: u64,
) -> anyhow::Result<RandColProjOut>
where
T: Sync + Send + std::hash::Hash + Eq + Clone + ToString,
{
let ncols = self.num_columns()?;
let batch_membership = batch_membership.and_then(|x| x.get(0..ncols));
let proj_vec = self
.stack
.iter()
.enumerate()
.map(|(m, data_vec)| -> anyhow::Result<_> {
data_vec.project_columns_with_batch_correction_seeded(
target_dim,
block_size,
batch_membership,
mix_seed(seed, m as u64),
)
})
.collect::<anyhow::Result<Vec<_>>>()?;
let basis = concatenate_vertical(
&proj_vec
.iter()
.map(|x| {
x.basis.to_owned()
})
.collect::<Vec<_>>(),
)?;
let proj = concatenate_vertical(
&proj_vec
.iter()
.map(|x| {
x.proj.to_owned()
})
.collect::<Vec<_>>(),
)?;
Ok(RandColProjOut { basis, proj })
}
fn project_columns_weighted_seeded<T>(
&self,
target_dim: usize,
block_size: Option<usize>,
batch_membership: Option<&[T]>,
row_weights: &[f32],
seed: u64,
) -> anyhow::Result<RandColProjOut>
where
T: Sync + Send + std::hash::Hash + Eq + Clone + ToString,
{
let ncols = self.num_columns()?;
let batch_membership = batch_membership.and_then(|x| x.get(0..ncols));
let total_rows: usize = self.stack.iter().map(|v| v.num_rows()).sum();
if row_weights.len() != total_rows {
return Err(anyhow::anyhow!(
"row_weights length {} does not match total feature count {} \
(sum of num_rows() across the stack)",
row_weights.len(),
total_rows
));
}
let mut offset = 0usize;
let mut proj_vec = Vec::with_capacity(self.stack.len());
for (m, data_vec) in self.stack.iter().enumerate() {
let nrows = data_vec.num_rows();
let slice = &row_weights[offset..offset + nrows];
offset += nrows;
proj_vec.push(data_vec.project_columns_weighted_seeded(
target_dim,
block_size,
batch_membership,
slice,
mix_seed(seed, m as u64),
)?);
}
let basis = concatenate_vertical(
&proj_vec
.iter()
.map(|x| x.basis.to_owned())
.collect::<Vec<_>>(),
)?;
let proj = concatenate_vertical(
&proj_vec
.iter()
.map(|x| x.proj.to_owned())
.collect::<Vec<_>>(),
)?;
Ok(RandColProjOut { basis, proj })
}
fn partition_columns_to_groups(
&mut self,
proj_kn: &nalgebra::DMatrix<f32>,
num_sorting_features: Option<usize>,
ncols_per_group: Option<usize>,
) -> anyhow::Result<usize> {
let nn = proj_kn.ncols();
if nn != self.num_columns()? {
return Err(anyhow::anyhow!("number of columns mismatch"));
}
let target_kk = num_sorting_features.unwrap_or(proj_kn.nrows());
let kk = proj_kn.nrows().min(target_kk).min(nn);
let binary_codes = binary_sort_columns(proj_kn, kk)?;
let max_group = *binary_codes
.iter()
.max()
.ok_or(anyhow::anyhow!("unable to determine max element"))?;
for x in self.stack.iter_mut() {
x.assign_groups(&binary_codes, ncols_per_group);
}
Ok(max_group + 1)
}
fn assign_group_labels(&mut self, labels: &[usize]) {
for x in self.stack.iter_mut() {
x.assign_groups(labels, None);
}
}
}
impl RandProjOps for SparseIoVec {
fn project_columns_with_batch_correction_seeded<T>(
&self,
target_dim: usize,
block_size: Option<usize>,
batch_membership: Option<&[T]>,
seed: u64,
) -> anyhow::Result<RandColProjOut>
where
T: Sync + Send + std::hash::Hash + Eq + Clone + ToString,
{
let nrows = self.num_rows();
let ncols = self.num_columns();
let mut proj_kn = nalgebra::DMatrix::<f32>::zeros(target_dim, ncols);
let basis_dk = nalgebra::DMatrix::<f32>::rnorm_seeded(nrows, target_dim, seed);
let basis_kd = basis_dk.transpose();
let shared_in = ProjVisitorIn {
basis_kd: &basis_kd,
};
info!(
"starting random projection loop (project_columns_with_batch_correction): \
D={nrows}, N={ncols}, K={target_dim}, block_size={:?}",
block_size
);
self.visit_columns_by_block(
&project_columns_visitor,
&shared_in,
&mut proj_kn,
block_size,
)?;
info!("finished random projection loop");
if let Some(col_to_batch) = batch_membership {
adjust_batch_biases(&mut proj_kn, col_to_batch)?;
}
let (lb, ub) = (-4., 4.);
proj_kn.scale_columns_inplace();
if proj_kn.max() > ub || proj_kn.min() < lb {
proj_kn.iter_mut().for_each(|x| {
*x = x.clamp(lb, ub);
});
proj_kn.scale_columns_inplace();
}
Ok(RandColProjOut {
basis: basis_dk,
proj: proj_kn,
})
}
fn project_columns_weighted_seeded<T>(
&self,
target_dim: usize,
block_size: Option<usize>,
batch_membership: Option<&[T]>,
row_weights: &[f32],
seed: u64,
) -> anyhow::Result<RandColProjOut>
where
T: Sync + Send + std::hash::Hash + Eq + Clone + ToString,
{
let nrows = self.num_rows();
let ncols = self.num_columns();
assert_eq!(row_weights.len(), nrows, "row_weights length mismatch");
let mut proj_kn = nalgebra::DMatrix::<f32>::zeros(target_dim, ncols);
let mut basis_dk = nalgebra::DMatrix::<f32>::rnorm_seeded(nrows, target_dim, seed);
for (r, &w) in row_weights.iter().enumerate() {
if w <= 0.0 {
basis_dk.row_mut(r).fill(0.0);
} else if (w - 1.0).abs() > 1e-6 {
basis_dk.row_mut(r).scale_mut(w);
}
}
let basis_kd = basis_dk.transpose();
let shared_in = ProjVisitorIn {
basis_kd: &basis_kd,
};
info!(
"starting random projection loop (project_columns_weighted): \
D={nrows}, N={ncols}, K={target_dim}, block_size={:?}",
block_size
);
self.visit_columns_by_block(
&project_columns_visitor,
&shared_in,
&mut proj_kn,
block_size,
)?;
info!("finished random projection loop");
if let Some(col_to_batch) = batch_membership {
adjust_batch_biases(&mut proj_kn, col_to_batch)?;
}
let (lb, ub) = (-4., 4.);
proj_kn.scale_columns_inplace();
if proj_kn.max() > ub || proj_kn.min() < lb {
proj_kn.iter_mut().for_each(|x| {
*x = x.clamp(lb, ub);
});
proj_kn.scale_columns_inplace();
}
Ok(RandColProjOut {
basis: basis_dk,
proj: proj_kn,
})
}
fn partition_columns_to_groups(
&mut self,
proj_kn: &nalgebra::DMatrix<f32>,
num_sorting_features: Option<usize>,
ncols_per_group: Option<usize>,
) -> anyhow::Result<usize> {
let nn = proj_kn.ncols();
if nn != self.num_columns() {
return Err(anyhow::anyhow!("number of columns mismatch"));
}
let target_kk = num_sorting_features.unwrap_or(proj_kn.nrows());
let kk = proj_kn.nrows().min(target_kk).min(nn);
let binary_codes = binary_sort_columns(proj_kn, kk)?;
let max_group = *binary_codes
.iter()
.max()
.ok_or(anyhow::anyhow!("unable to determine max element"))?;
self.assign_groups(&binary_codes, ncols_per_group);
Ok(max_group + 1)
}
fn assign_group_labels(&mut self, labels: &[usize]) {
self.assign_groups(labels, None);
}
}
fn adjust_batch_biases<T>(
proj_kn: &mut nalgebra::DMatrix<f32>,
col_to_batch: &[T],
) -> anyhow::Result<()>
where
T: std::hash::Hash + Eq + Clone,
{
let ncols = proj_kn.ncols();
if col_to_batch.len() != ncols {
warn!(
"The batch membership size {} mismatches with the number of columns {} ...",
col_to_batch.len(),
ncols
);
return Ok(());
}
let batch = batch_indices(col_to_batch);
let n_batches = batch.iter().max().map_or(0, |&m| m + 1);
let bits = default_state_bits(ncols, n_batches, proj_kn.nrows());
info!(
"adjusting batch biases within cell state ({} batches, {} state bits) ...",
n_batches, bits
);
*proj_kn = centre_batches_within_state(proj_kn, &batch, bits)?;
Ok(())
}
pub fn binary_sort_columns(
proj_kn: &nalgebra::DMatrix<f32>,
kk: usize,
) -> anyhow::Result<Vec<usize>> {
let nn = proj_kn.ncols();
let (_, _, mut q_nk) = proj_kn.rsvd(kk)?;
q_nk.scale_columns_inplace();
let mut binary_codes = DVector::<usize>::zeros(nn);
for k in 0..kk {
let binary_shift = |x: f32| -> usize {
if x > 0.0 {
1 << k
} else {
0
}
};
binary_codes += q_nk.column(k).map(binary_shift);
}
Ok(binary_codes.data.as_vec().clone())
}
pub fn max_binary_code(kk: usize) -> usize {
(1 << kk) - 1
}
#[cfg(test)]
mod repro_tests {
use super::*;
use crate::sparse_io::{create_sparse_from_triplets, SparseIoBackend};
use std::sync::Arc;
fn tiny_data_vec(dir: &tempfile::TempDir) -> anyhow::Result<SparseIoVec> {
let path = dir.path().join("proj.zarr");
let path_str = path.to_str().unwrap();
let (d, n) = (8usize, 12usize);
let mut triplets: Vec<(u64, u64, f32)> = Vec::new();
for i in 0..d {
for j in 0..n {
if (i * 7 + j * 3) % 5 < 3 {
triplets.push((i as u64, j as u64, 1.0 + ((i + j) % 4) as f32));
}
}
}
let nnz = triplets.len();
let mut backend = create_sparse_from_triplets(
&triplets,
(d, n, nnz),
Some(path_str),
Some(&SparseIoBackend::Zarr),
)?;
let gene_names: Vec<Box<str>> = (0..d).map(|i| format!("g{i}").into()).collect();
let cell_names: Vec<Box<str>> = (0..n).map(|j| format!("c{j}").into()).collect();
backend.register_row_names_vec(&gene_names);
backend.register_column_names_vec(&cell_names);
let mut data_vec = SparseIoVec::new();
data_vec.push(Arc::from(backend), None)?;
Ok(data_vec)
}
#[test]
fn seeded_projection_is_reproducible() -> anyhow::Result<()> {
let dir = tempfile::tempdir()?;
let data_vec = tiny_data_vec(&dir)?;
let a =
data_vec.project_columns_with_batch_correction_seeded::<usize>(4, None, None, 123)?;
let b =
data_vec.project_columns_with_batch_correction_seeded::<usize>(4, None, None, 123)?;
assert_eq!(a.basis, b.basis, "same seed → identical basis");
assert_eq!(a.proj, b.proj, "same seed → identical projection");
Ok(())
}
#[test]
fn default_projection_is_reproducible() -> anyhow::Result<()> {
let dir = tempfile::tempdir()?;
let data_vec = tiny_data_vec(&dir)?;
let a = data_vec.project_columns(4, None)?;
let b = data_vec.project_columns(4, None)?;
assert_eq!(a.basis, b.basis);
assert_eq!(a.proj, b.proj);
Ok(())
}
#[test]
fn distinct_seeds_diverge() -> anyhow::Result<()> {
let dir = tempfile::tempdir()?;
let data_vec = tiny_data_vec(&dir)?;
let a = data_vec.project_columns_with_batch_correction_seeded::<usize>(4, None, None, 1)?;
let b = data_vec.project_columns_with_batch_correction_seeded::<usize>(4, None, None, 2)?;
assert_ne!(a.basis, b.basis, "distinct seeds → distinct basis");
Ok(())
}
}