use crate::alg::union_find::UnionFind;
use special::Gamma as SpecialGamma;
pub struct BhcInput<'a> {
pub k: usize,
pub m: usize,
pub gene_sum: &'a [f64],
pub size_sum: &'a [f64],
pub effective_size: &'a [usize],
}
#[derive(Debug, Clone)]
pub struct BhcMerge {
pub id: i32,
pub left: i32,
pub right: i32,
pub log_bf: f64,
pub n_samples: i32,
}
struct NodeStats {
t_gene: Vec<f64>,
s_size: f64,
n_samples: i32,
f_cache: f64,
id: i32,
}
#[inline]
fn order_pair(a: i32, b: i32) -> (i32, i32) {
if a < b {
(a, b)
} else {
(b, a)
}
}
struct Prior {
gamma: f64,
gamma_bg: Vec<f64>,
sum_lgamma_gamma_bg: f64,
lgamma_gamma: f64,
}
#[inline]
fn f_node(t_gene: &[f64], s_size: f64, prior: &Prior) -> f64 {
let mut sum_new = 0.0f64;
for (&t, &gb) in t_gene.iter().zip(prior.gamma_bg.iter()) {
sum_new += SpecialGamma::ln_gamma(gb + t).0;
}
prior.lgamma_gamma - SpecialGamma::ln_gamma(prior.gamma + s_size).0 + sum_new
- prior.sum_lgamma_gamma_bg
}
#[inline]
fn pairwise_log_bf(a: &NodeStats, b: &NodeStats, prior: &Prior) -> f64 {
let mut merged_sum = 0.0f64;
for ((&ta, &tb), &gb) in a
.t_gene
.iter()
.zip(b.t_gene.iter())
.zip(prior.gamma_bg.iter())
{
merged_sum += SpecialGamma::ln_gamma(gb + ta + tb).0;
}
let s_merged = a.s_size + b.s_size;
let f_merged = prior.lgamma_gamma - SpecialGamma::ln_gamma(prior.gamma + s_merged).0
+ merged_sum
- prior.sum_lgamma_gamma_bg;
f_merged - a.f_cache - b.f_cache
}
pub fn bhc_merge(input: BhcInput<'_>, gamma: f64) -> Vec<BhcMerge> {
const BG_EPS: f64 = 1e-9;
let k = input.k;
let m = input.m;
let gamma = gamma.max(1e-6);
assert_eq!(input.gene_sum.len(), k * m, "gene_sum must be k × m");
assert_eq!(input.size_sum.len(), k);
assert_eq!(input.effective_size.len(), k);
let mut bg = vec![0.0f64; m];
let mut total_mass = 0.0f64;
for c in 0..k {
if input.effective_size[c] == 0 {
continue;
}
let row = &input.gene_sum[c * m..(c + 1) * m];
for (dst, &src) in bg.iter_mut().zip(row.iter()) {
*dst += src;
}
total_mass += input.size_sum[c];
}
let total_mass = total_mass.max(BG_EPS);
for p in bg.iter_mut() {
*p = (*p / total_mass).max(BG_EPS);
}
let gamma_bg: Vec<f64> = bg.iter().map(|&p| gamma * p).collect();
let sum_lgamma_gamma_bg: f64 = gamma_bg
.iter()
.map(|&gb| SpecialGamma::ln_gamma(gb).0)
.sum();
let prior = Prior {
gamma,
gamma_bg,
sum_lgamma_gamma_bg,
lgamma_gamma: SpecialGamma::ln_gamma(gamma).0,
};
let mut arena: Vec<NodeStats> = Vec::with_capacity(k);
let mut active: Vec<usize> = Vec::with_capacity(k);
for c in 0..k {
if input.effective_size[c] == 0 {
continue;
}
let s_raw = input.size_sum[c];
let n_eff = input.effective_size[c] as f64;
let scale = if s_raw > 0.0 { n_eff / s_raw } else { 0.0 };
let t: Vec<f64> = input.gene_sum[c * m..(c + 1) * m]
.iter()
.map(|&x| x * scale)
.collect();
let f = f_node(&t, n_eff, &prior);
arena.push(NodeStats {
t_gene: t,
s_size: n_eff,
n_samples: input.effective_size[c] as i32,
f_cache: f,
id: c as i32,
});
active.push(arena.len() - 1);
}
let k_eff = active.len();
if k_eff < 2 {
return Vec::new();
}
let mut merges: Vec<BhcMerge> = Vec::with_capacity(k_eff - 1);
let mut next_id: i32 = k as i32;
while active.len() >= 2 {
let mut best: Option<(usize, usize, f64, (i32, i32))> = None;
for i in 0..active.len() {
for j in (i + 1)..active.len() {
let a = &arena[active[i]];
let b = &arena[active[j]];
let bf = pairwise_log_bf(a, b, &prior);
let key = order_pair(a.id, b.id);
let pick = match best {
None => true,
Some((_, _, best_bf, best_key)) => {
bf > best_bf || (bf == best_bf && key < best_key)
}
};
if pick {
best = Some((i, j, bf, key));
}
}
}
let (i, j, log_bf, _) = best.expect("at least one pair exists");
let ai = active[i];
let aj = active[j];
let (left_id, right_id) = order_pair(arena[ai].id, arena[aj].id);
let new_t: Vec<f64> = arena[ai]
.t_gene
.iter()
.zip(arena[aj].t_gene.iter())
.map(|(x, y)| x + y)
.collect();
let new_s = arena[ai].s_size + arena[aj].s_size;
let new_n = arena[ai].n_samples + arena[aj].n_samples;
let new_f = f_node(&new_t, new_s, &prior);
arena.push(NodeStats {
t_gene: new_t,
s_size: new_s,
n_samples: new_n,
f_cache: new_f,
id: next_id,
});
let new_arena_idx = arena.len() - 1;
merges.push(BhcMerge {
id: next_id,
left: left_id,
right: right_id,
log_bf,
n_samples: new_n,
});
next_id += 1;
active.remove(j);
active.remove(i);
active.push(new_arena_idx);
}
merges
}
pub fn bhc_cut(merges: &[BhcMerge], k: usize, cutoff: f64) -> Vec<i32> {
let mut uf = UnionFind::new(k);
let mut rep: Vec<usize> = Vec::with_capacity(k + merges.len());
rep.extend(0..k);
rep.resize(k + merges.len(), 0);
let mut referenced = vec![false; k];
for m in merges {
if (m.left as usize) < k {
referenced[m.left as usize] = true;
}
if (m.right as usize) < k {
referenced[m.right as usize] = true;
}
let l_rep = rep[m.left as usize];
let r_rep = rep[m.right as usize];
if m.log_bf >= cutoff {
uf.union(l_rep, r_rep);
}
rep[m.id as usize] = l_rep.min(r_rep);
}
let mut labels = vec![-1i32; k];
let mut root_to_dense: Vec<Option<i32>> = vec![None; k];
let mut next_dense: i32 = 0;
for c in 0..k {
if !referenced[c] {
continue;
}
let root = uf.find(c);
let label = match root_to_dense[root] {
Some(l) => l,
None => {
let l = next_dense;
root_to_dense[root] = Some(l);
next_dense += 1;
l
}
};
labels[c] = label;
}
labels
}