use ndarray::Array2;
use ndarray_stats::CorrelationExt;
use rand::Rng;
use rayon::prelude::*;
use std::collections::HashMap;
use super::transcripts::TranscriptDataset;
const NFEATURES: usize = 10_000;
const BIN_SIZE: f32 = 20.0;
const NGENES_CANDIDATES: usize = 5000;
const NCORR_ROWS: usize = 2000;
const MIN_MEAN_TRANSCRIPTS_PER_REGION: f32 = 5e-2;
fn random_region_counts(
dataset: &TranscriptDataset,
nregions: usize,
bin_size: f32,
) -> Array2<f32> {
let mut rng = rand::rng();
let mut region_map = HashMap::new();
for i in 0..nregions {
let centroid_transcript =
&dataset.transcripts.runs[rng.random_range(0..dataset.transcripts.runs.len())].value;
let cx = centroid_transcript.x;
let cy = centroid_transcript.y;
let bin_x = (cx / bin_size).floor() as usize;
let bin_y = (cy / bin_size).floor() as usize;
region_map.insert((bin_x, bin_y), i);
}
let mut counts = Array2::zeros((nregions, dataset.ngenes()));
for transcript_run in dataset.transcripts.iter_runs() {
let bin_x = (transcript_run.value.x / bin_size).floor() as usize;
let bin_y = (transcript_run.value.y / bin_size).floor() as usize;
let region = region_map.get(&(bin_x, bin_y));
if let Some(region) = region {
counts[[*region, transcript_run.value.gene as usize]] += transcript_run.len as f32;
}
}
counts
}
fn deviance_ranking(counts: &Array2<f32>) -> Vec<usize> {
let nregions = counts.nrows();
let ngenes = counts.ncols();
let region_totals: Vec<f64> = (0..nregions)
.into_par_iter()
.map(|i| counts.row(i).iter().map(|&x| x as f64).sum())
.collect();
let gene_totals: Vec<f64> = (0..ngenes)
.into_par_iter()
.map(|j| counts.column(j).iter().map(|&x| x as f64).sum())
.collect();
let grand_total: f64 = gene_totals.iter().sum();
let counts_t = counts.t().to_owned();
let mut deviances: Vec<(usize, f64)> = (0..ngenes)
.into_par_iter()
.map(|j| {
let p_j = if grand_total > 0.0 {
gene_totals[j] / grand_total
} else {
0.0
};
let deviance: f64 = (0..nregions)
.map(|i| {
let y = counts_t[[j, i]] as f64;
let n = region_totals[i];
if n == 0.0 {
return 0.0;
}
let mu = n * p_j;
let pos = if y > 0.0 && mu > 0.0 {
y * (y / mu).ln()
} else {
0.0
};
let neg = if (n - y) > 0.0 && (n - mu) > 0.0 {
(n - y) * ((n - y) / (n - mu)).ln()
} else {
0.0
};
pos + neg
})
.sum::<f64>()
* 2.0;
(j, deviance)
})
.collect();
deviances.sort_unstable_by(|a, b| b.1.partial_cmp(&a.1).unwrap());
deviances.into_iter().map(|(j, _)| j).collect()
}
fn select_k_clusters<T: PartialOrd + Copy>(
hclust: &kodama::Dendrogram<T>,
nclusters: usize,
) -> Vec<usize> {
let n = hclust.observations();
let steps = hclust.steps();
let mut thresholds: Vec<T> = steps.iter().map(|s| s.dissimilarity).collect();
thresholds.sort_unstable_by(|a, b| a.partial_cmp(b).unwrap());
thresholds.dedup_by(|a, b| a == b);
let count_clusters = |threshold: T| -> usize {
let total = 2 * n - 1;
let mut parent: Vec<usize> = (0..total).collect();
fn find(parent: &mut Vec<usize>, mut x: usize) -> usize {
while parent[x] != x {
parent[x] = parent[parent[x]];
x = parent[x];
}
x
}
for (k, step) in steps.iter().enumerate() {
if step.dissimilarity <= threshold {
let internal = n + k;
let ra = find(&mut parent, step.cluster1);
let rb = find(&mut parent, step.cluster2);
let ri = find(&mut parent, internal);
parent[ra] = ri;
parent[rb] = ri;
}
}
let mut roots = std::collections::HashSet::new();
for i in 0..n {
roots.insert(find(&mut parent, i));
}
roots.len()
};
let mut lo = 0usize;
let mut hi = thresholds.len();
while lo < hi {
let mid = lo + (hi - lo) / 2;
if count_clusters(thresholds[mid]) <= nclusters {
hi = mid;
} else {
lo = mid + 1;
}
}
let chosen_threshold = if lo < thresholds.len() {
thresholds[lo]
} else {
thresholds[thresholds.len() - 1]
};
let total = 2 * n - 1;
let mut parent: Vec<usize> = (0..total).collect();
fn find(parent: &mut Vec<usize>, mut x: usize) -> usize {
while parent[x] != x {
parent[x] = parent[parent[x]];
x = parent[x];
}
x
}
for (k, step) in steps.iter().enumerate() {
if step.dissimilarity <= chosen_threshold {
let internal = n + k;
let ra = find(&mut parent, step.cluster1);
let rb = find(&mut parent, step.cluster2);
let ri = find(&mut parent, internal);
parent[ra] = ri;
parent[rb] = ri;
}
}
let mut root_to_cluster: HashMap<usize, usize> = HashMap::new();
let mut labels = vec![0usize; n];
for i in 0..n {
let root = find(&mut parent, i);
let next_id = root_to_cluster.len();
let cluster_id = *root_to_cluster.entry(root).or_insert(next_id);
labels[i] = cluster_id;
}
labels
}
fn select_features(dataset: &TranscriptDataset, nfeatures: usize) -> Vec<usize> {
let counts = random_region_counts(dataset, NFEATURES, BIN_SIZE);
let gene_ranking = deviance_ranking(&counts);
let ncandidates = gene_ranking.len().min(NGENES_CANDIDATES);
let col_indices: Vec<usize> = gene_ranking.into_iter().take(ncandidates).collect();
let nrows = counts.nrows();
let counts =
Array2::from_shape_fn((nrows, ncandidates), |(i, j)| counts[[i, col_indices[j]]]);
let n = nrows as f32;
let valid_cols: Vec<usize> = (0..ncandidates)
.filter(|&j| {
let col = counts.column(j);
let mean = col.sum() / n;
col.iter().any(|&v| (v - mean).abs() > 1e-9)
})
.collect();
let gene_totals: Vec<f32> = valid_cols.iter().map(|&j| counts.column(j).sum()).collect();
let col_indices: Vec<usize> = valid_cols.iter().map(|&j| col_indices[j]).collect();
let counts =
Array2::from_shape_fn((nrows, valid_cols.len()), |(i, j)| counts[[i, valid_cols[j]]]);
let nrows_f = nrows as f32;
let expr_indices: Vec<usize> = (0..col_indices.len())
.filter(|&j| gene_totals[j] / nrows_f >= MIN_MEAN_TRANSCRIPTS_PER_REGION)
.collect();
let col_indices: Vec<usize> = expr_indices.iter().map(|&j| col_indices[j]).collect();
let gene_totals: Vec<f32> = expr_indices.iter().map(|&j| gene_totals[j]).collect();
let mut counts =
Array2::from_shape_fn((nrows, expr_indices.len()), |(i, j)| counts[[i, expr_indices[j]]]);
counts.map_inplace(|v| *v = v.ln_1p());
let ngenes = col_indices.len();
let corr_counts = if nrows > NCORR_ROWS {
let mut rng = rand::rng();
let row_indices: Vec<usize> = rand::seq::index::sample(&mut rng, nrows, NCORR_ROWS).into_vec();
Array2::from_shape_fn((NCORR_ROWS, ngenes), |(i, j)| counts[[row_indices[i], j]])
} else {
counts
};
let corr = corr_counts.t().pearson_correlation().unwrap();
let mut condensed_dissim = Vec::with_capacity(ngenes * (ngenes - 1) / 2);
for i in 0..ngenes {
for j in i + 1..ngenes {
let c = corr[[i, j]];
condensed_dissim.push(if c.is_finite() {
(1.0 - c).clamp(0.0, 2.0)
} else {
1.0
});
}
}
let hclust = kodama::linkage(&mut condensed_dissim, ngenes, kodama::Method::Average);
let gene_cluster_assignments = select_k_clusters(&hclust, nfeatures);
let nclusters = gene_cluster_assignments
.iter()
.copied()
.max()
.map_or(0, |m| m + 1);
let mut best_local: Vec<Option<usize>> = vec![None; nclusters];
for (local_idx, &cluster_id) in gene_cluster_assignments.iter().enumerate() {
let current = best_local[cluster_id];
if current.is_none() || gene_totals[local_idx] > gene_totals[current.unwrap()] {
best_local[cluster_id] = Some(local_idx);
}
}
best_local
.into_iter()
.filter_map(|opt| opt.map(|local_idx| col_indices[local_idx]))
.collect()
}
impl TranscriptDataset {
pub fn select_unfactored_genes(&mut self, nunfactored: usize) {
let selected_features = select_features(self, nunfactored);
let selected_set: std::collections::HashSet<usize> =
selected_features.iter().copied().collect();
let mut ord: Vec<usize> = selected_features.clone();
for gene_idx in 0..self.gene_names.len() {
if !selected_set.contains(&gene_idx) {
ord.push(gene_idx);
}
}
let mut rev_ord = vec![0; ord.len()];
for (i, j) in ord.iter().enumerate() {
rev_ord[*j] = i;
}
self.gene_names = ord.iter().map(|&i| self.gene_names[i].clone()).collect();
for transcript_run in self.transcripts.iter_runs_mut() {
transcript_run.value.gene = rev_ord[transcript_run.value.gene as usize] as u32;
}
}
}