use nalgebra::DMatrix;
use super::latents::Latents;
use super::MODALITIES;
use crate::sim::multiome::sample_poisson_from_logits;
const EPS_LOG: f32 = 1e-12;
const NEG_BIG: f32 = -50.0;
pub type Triplet = (u64, u64, f32);
pub type RowKey = (usize, usize);
pub struct RateContext<'a> {
pub lats: &'a Latents,
pub log_topic: &'a DMatrix<f32>,
pub log_topic_future: Option<&'a DMatrix<f32>>,
pub batch_membership: &'a [usize],
}
pub fn precompute_log_topic(lats: &Latents) -> DMatrix<f32> {
precompute_log_topic_from(&lats.beta_topic_gk, &lats.theta_kn)
}
pub fn precompute_log_topic_from(
beta_topic_gk: &DMatrix<f32>,
theta_kn: &DMatrix<f32>,
) -> DMatrix<f32> {
let lambda = beta_topic_gk * theta_kn;
lambda.map(|v| v.max(EPS_LOG).ln())
}
pub fn sample_count_modality(
ctx: &RateContext,
ln_delta_count: &DMatrix<f32>,
depth_count: usize,
rseed: u64,
) -> Vec<Triplet> {
let lats = ctx.lats;
let g = lats.beta_g.len();
let n = lats.theta_kn.ncols();
let depths = vec![depth_count as f32; n];
let log_mu_spliced = assemble_log_mu(ctx, ln_delta_count);
let log_mu_unspliced = match ctx.log_topic_future {
Some(future) => assemble_log_mu_from(future, lats, ctx.batch_membership, ln_delta_count),
None => log_mu_spliced.clone(),
};
let mut log_rate = DMatrix::<f32>::zeros(2 * g, n);
for c in 0..2 {
let log_mu = if c == 0 {
&log_mu_spliced
} else {
&log_mu_unspliced
};
for gi in 0..g {
let row = c * g + gi;
let log_alpha = (lats.alpha_per_mod[0][(c, gi)] + EPS_LOG).ln();
for j in 0..n {
log_rate[(row, j)] = log_alpha + log_mu[(gi, j)];
}
}
}
sample_poisson_from_logits(&log_rate, &depths, rseed)
}
pub fn sample_modifier_modality(
ctx: &RateContext,
m_idx: usize,
held_out: &[Vec<bool>],
ln_delta_m: &DMatrix<f32>,
depth_m: usize,
rseed: u64,
) -> (Vec<Triplet>, Vec<RowKey>) {
let lats = ctx.lats;
let g = lats.beta_g.len();
let n = lats.theta_kn.ncols();
let c_m = lats.alpha_per_mod[m_idx].nrows();
let emit_genes: Vec<usize> = (0..g)
.filter(|&gi| lats.phi[m_idx][gi] && !held_out[m_idx][gi])
.collect();
let d = emit_genes.len() * c_m;
let row_keys: Vec<RowKey> = emit_genes
.iter()
.flat_map(|&gi| (0..c_m).map(move |c| (gi, c)))
.collect();
if d == 0 {
log::warn!(
"modality '{}' has 0 emitted rows (all substrate-negative or held-out)",
MODALITIES[m_idx]
);
return (Vec::new(), row_keys);
}
let log_mu = assemble_log_mu(ctx, ln_delta_m);
let log_r = build_log_modifier_rate(ctx, m_idx, ln_delta_m);
let mut log_rate = DMatrix::<f32>::zeros(d, n);
let alpha_m = &lats.alpha_per_mod[m_idx];
for (row_id, &(gi, ci)) in row_keys.iter().enumerate() {
let log_alpha = (alpha_m[(ci, gi)] + EPS_LOG).ln();
for j in 0..n {
let v = log_alpha + log_mu[(gi, j)] + log_r[(gi, j)];
log_rate[(row_id, j)] = if v.is_finite() { v } else { NEG_BIG };
}
}
let depths = vec![depth_m as f32; n];
let triplets = sample_poisson_from_logits(&log_rate, &depths, rseed);
(triplets, row_keys)
}
fn assemble_log_mu(ctx: &RateContext, ln_delta: &DMatrix<f32>) -> DMatrix<f32> {
assemble_log_mu_from(ctx.log_topic, ctx.lats, ctx.batch_membership, ln_delta)
}
fn assemble_log_mu_from(
log_topic: &DMatrix<f32>,
lats: &Latents,
batch_membership: &[usize],
ln_delta: &DMatrix<f32>,
) -> DMatrix<f32> {
let g = lats.beta_g.len();
let n = log_topic.ncols();
let bb = ln_delta.ncols();
let mut log_mu = log_topic.clone();
for j in 0..n {
let b = batch_membership[j];
for gi in 0..g {
let delta = if bb > 1 { ln_delta[(gi, b)] } else { 0.0 };
log_mu[(gi, j)] += lats.beta_g[gi] + delta;
}
}
log_mu
}
fn build_log_modifier_rate(
ctx: &RateContext,
m_idx: usize,
ln_delta: &DMatrix<f32>,
) -> DMatrix<f32> {
let lats = ctx.lats;
let g = lats.beta_g.len();
let n = lats.theta_kn.ncols();
let k_prog = lats.a_mk.ncols();
let bb = ln_delta.ncols();
let mut z_scaled = lats.z_gk.clone();
for gi in 0..g {
for k in 0..k_prog {
z_scaled[(gi, k)] *= lats.a_mk[(m_idx, k)];
}
}
let prog = z_scaled * &lats.theta_kn;
let mut log_r = DMatrix::<f32>::zeros(g, n);
for j in 0..n {
let b = ctx.batch_membership[j];
for gi in 0..g {
let phi = if lats.phi[m_idx][gi] { 1.0 } else { 0.0 };
let delta = if bb > 1 { ln_delta[(gi, b)] } else { 0.0 };
log_r[(gi, j)] = lats.base_gm[(gi, m_idx)] + phi * prog[(gi, j)] + delta;
}
}
log_r
}