use crate::candle::batched_dot::batched_matvec;
use crate::candle::decoder::coarsening_map::CoarseningMap;
use crate::candle::fast_index::gather_rows;
use crate::candle::loss::nb_log_likelihood_elem;
use candle_core::{Result, Tensor};
use candle_nn::{ops, VarBuilder};
pub struct MaskedNbTarget<'a> {
pub indices: &'a Tensor,
pub residual: Option<&'a Tensor>,
pub values: &'a Tensor,
pub lib: &'a Tensor,
pub mask: &'a Tensor,
}
pub struct MaskedDenseTarget<'a> {
pub values: &'a Tensor,
pub residual: Option<&'a Tensor>,
pub lib: &'a Tensor,
pub hidden_ids: &'a Tensor,
pub hidden_weight: Option<&'a Tensor>,
}
pub struct ModuleTarget<'a> {
pub values: &'a Tensor,
pub visible_counts: &'a Tensor,
pub visible_share: &'a Tensor,
pub residual: Option<&'a Tensor>,
pub lib: &'a Tensor,
}
pub struct QueryTarget<'a> {
pub gene_ids: &'a Tensor,
pub values: &'a Tensor,
pub weight: &'a Tensor,
pub log_residual: &'a Tensor,
pub lib: &'a Tensor,
}
pub struct EmbeddedNbTopicDecoder {
n_features: usize,
n_obs: usize,
n_topics: usize,
coarsening: CoarseningMap,
topic_embeddings: Tensor,
features: std::sync::Arc<crate::candle::feature_embedding::FeatureEmbedding>,
log_phi_1d: Tensor,
log_pi_1d: Tensor,
}
pub const BACKGROUND_VAR: &str = "log_pi";
pub fn pin_background(varmap: &candle_nn::VarMap, prefix: &str, log_pi_1d: &Tensor) -> Result<()> {
let name = format!("{prefix}.{BACKGROUND_VAR}");
let tbl = varmap.data().lock().unwrap();
let var = tbl
.get(&name)
.ok_or_else(|| candle_core::Error::Msg(format!("no var `{name}` to pin")))?;
var.set(&log_pi_1d.to_device(var.device())?.to_dtype(var.dtype())?)
}
pub fn log_background_from_mean(mean_d: &[f32], device: &candle_core::Device) -> Result<Tensor> {
let total: f64 = mean_d.iter().map(|&m| f64::from(m.max(0.0))).sum();
let d = mean_d.len().max(1) as f64;
let floor = 1e-3 / d;
let log_pi: Vec<f32> = mean_d
.iter()
.map(|&m| {
let p = if total > 0.0 {
f64::from(m.max(0.0)) / total
} else {
1.0 / d
};
p.max(floor).ln() as f32
})
.collect();
Tensor::from_vec(log_pi, (1, mean_d.len()), device)
}
impl EmbeddedNbTopicDecoder {
pub fn new(
n_topics: usize,
features: std::sync::Arc<crate::candle::feature_embedding::FeatureEmbedding>,
vs: VarBuilder,
) -> Result<Self> {
let identity = CoarseningMap::identity(features.n_features(), features.device())?;
Self::new_with_coarsening(n_topics, features, identity, vs)
}
pub fn new_with_coarsening(
n_topics: usize,
features: std::sync::Arc<crate::candle::feature_embedding::FeatureEmbedding>,
coarsening: CoarseningMap,
vs: VarBuilder,
) -> Result<Self> {
let n_features = features.n_features();
let embedding_dim = features.embedding_dim();
if coarsening.n_fine() != n_features {
candle_core::bail!(
"EmbeddedNbTopicDecoder: the coarsening covers {} features but ρ has {n_features}",
coarsening.n_fine()
);
}
let n_obs = coarsening.n_coarse();
let init_ws = candle_nn::init::DEFAULT_KAIMING_NORMAL;
let topic_embeddings =
vs.get_with_hints((n_topics, embedding_dim), "topic.embeddings", init_ws)?;
let log_phi_1d = vs.get_with_hints((1, n_obs), "log_phi", candle_nn::Init::Const(0.693))?;
let log_pi_1d = vs.get_with_hints(
(1, n_obs),
BACKGROUND_VAR,
candle_nn::Init::Const(-(n_obs as f64).ln()),
)?;
Ok(Self {
n_features,
n_obs,
n_topics,
coarsening,
topic_embeddings,
features,
log_phi_1d,
log_pi_1d,
})
}
pub fn coarsening(&self) -> &CoarseningMap {
&self.coarsening
}
pub fn log_phi(&self) -> &Tensor {
&self.log_phi_1d
}
pub fn log_background(&self) -> &Tensor {
&self.log_pi_1d
}
pub fn phi(&self) -> Result<Tensor> {
self.log_phi_1d.exp()
}
pub fn dim_obs(&self) -> usize {
self.n_obs
}
pub fn n_features(&self) -> usize {
self.n_features
}
pub fn dim_latent(&self) -> usize {
self.n_topics
}
pub fn full_logits_kd(&self) -> Result<Tensor> {
let table = self
.features
.map_rows_linear(|rows| self.coarsening.coarsen_mean_dh(rows))?;
self.centered_topic_embeddings()?
.matmul(&table.t()?)?
.broadcast_add(&self.log_pi_1d)
}
fn centered_topic_embeddings(&self) -> Result<Tensor> {
let alpha_mean_1h = self.topic_embeddings.mean_keepdim(0)?;
self.topic_embeddings.broadcast_sub(&alpha_mean_1h)
}
pub fn get_dictionary(&self) -> Result<Tensor> {
let logits_kd = self.full_logits_kd()?;
let log_beta_kd = ops::log_softmax(&logits_kd, logits_kd.rank() - 1)?;
log_beta_kd.transpose(0, 1)?.contiguous()
}
pub fn log_partition_from_logits(full_kd: &Tensor) -> Result<Tensor> {
let k = full_kd.dim(0)?;
Self::log_partition_k1(full_kd)?.reshape((1, 1, k))
}
pub fn log_partition_k1(full_kd: &Tensor) -> Result<Tensor> {
let m = full_kd.max_keepdim(1)?; let lse = (full_kd.broadcast_sub(&m)?.exp()?.sum_keepdim(1)? + 1e-20)?.log()?; lse + m
}
pub fn mixture_rate_nd(&self, log_theta_nk: &Tensor, full_kd: &Tensor) -> Result<Tensor> {
let logz_k1 = Self::log_partition_k1(full_kd)?; let beta_kd = full_kd.broadcast_sub(&logz_k1)?.exp()?; log_theta_nk.exp()?.matmul(&beta_kd) }
pub(crate) fn mixture_rate_nk(
&self,
log_theta_nk: &Tensor,
indices: &Tensor,
full_kd: &Tensor,
) -> Result<Tensor> {
let n = indices.dim(0)?;
let k = indices.dim(1)?;
let t = self.n_topics;
let theta_nt = log_theta_nk.exp()?; let flat = self.coarsening.groups_of(indices)?.flatten_all()?;
let logz_11k = Self::log_partition_from_logits(full_kd)?; let logits = gather_rows(&full_kd.t()?.contiguous()?, &flat)? .reshape((n, k, t))?; let beta_nkt = logits.broadcast_sub(&logz_11k)?.exp()?;
let rate_nk = batched_matvec(&beta_nkt, &theta_nt)?; if self.coarsening.is_identity() {
return Ok(rate_nk);
}
rate_nk.mul(&self.coarsening.log_share_at(indices)?.exp()?)
}
pub fn impute_masked_nb(
&self,
log_theta_nk: &Tensor,
target: &MaskedNbTarget<'_>,
full_kd: &Tensor,
) -> Result<Tensor> {
let MaskedNbTarget {
indices,
residual: residual_nk,
values: values_nk,
lib: lib_n1,
mask: mask_nk,
} = *target;
let (n, k) = (indices.dim(0)?, indices.dim(1)?);
let theta_beta_nk = self.mixture_rate_nk(log_theta_nk, indices, full_kd)?;
let flat = self.coarsening.groups_of(indices)?.flatten_all()?; let log_phi_nk = gather_rows(&self.log_phi_1d.squeeze(0)?, &flat)?.reshape((n, k))?;
nb_score(
values_nk,
&theta_beta_nk,
residual_nk,
lib_n1,
&log_phi_nk,
Some(mask_nk),
)
}
pub fn impute_masked_multinomial(
&self,
log_theta_nk: &Tensor,
target: &MaskedNbTarget<'_>,
full_kd: &Tensor,
) -> Result<Tensor> {
let p_nk = self.mixture_rate_nk(log_theta_nk, target.indices, full_kd)?; multinomial_score(target.values, &p_nk, Some(target.mask))
}
pub fn impute_dense_nb(
&self,
log_theta_nk: &Tensor,
target: &MaskedDenseTarget<'_>,
full_kd: &Tensor,
) -> Result<Tensor> {
let ids = target.hidden_ids;
let (n, dh) = ids.dims2()?;
let rate_h = self
.mixture_rate_nd(log_theta_nk, full_kd)?
.gather(ids, 1)?;
let values_h = target.values.contiguous()?.gather(ids, 1)?;
let residual_h = target
.residual
.map(|r| r.contiguous()?.gather(ids, 1))
.transpose()?;
let log_phi_h =
gather_rows(&self.log_phi_1d.squeeze(0)?, &ids.flatten_all()?)?.reshape((n, dh))?;
nb_score(
&values_h,
&rate_h,
residual_h.as_ref(),
target.lib,
&log_phi_h,
target.hidden_weight,
)
}
pub fn impute_dense_multinomial(
&self,
log_theta_nk: &Tensor,
target: &MaskedDenseTarget<'_>,
full_kd: &Tensor,
) -> Result<Tensor> {
let ids = target.hidden_ids;
let rate_h = self
.mixture_rate_nd(log_theta_nk, full_kd)?
.gather(ids, 1)?;
let values_h = target.values.contiguous()?.gather(ids, 1)?;
multinomial_score(&values_h, &rate_h, target.hidden_weight)
}
pub fn score_unseen_modules_nb(
&self,
log_theta_nk: &Tensor,
target: &ModuleTarget<'_>,
full_km: &Tensor,
) -> Result<(Tensor, Tensor)> {
let (rate_nm, unseen, share, scored) = self.unseen_parts(log_theta_nk, target, full_km)?;
let mu = rate_nm.mul(&share)?.broadcast_mul(target.lib)?;
let mu = match target.residual {
Some(r) => mu.mul(r)?,
None => mu,
};
let log_phi = self.log_phi_1d.broadcast_as(mu.shape())?;
let elem = nb_log_likelihood_elem(&unseen, &mu, &log_phi)?;
Ok((elem.mul(&scored)?.sum(1)?, scored.sum(1)?))
}
pub fn score_unseen_modules_multinomial(
&self,
log_theta_nk: &Tensor,
target: &ModuleTarget<'_>,
full_km: &Tensor,
) -> Result<(Tensor, Tensor)> {
let (rate_nm, unseen, share, scored) = self.unseen_parts(log_theta_nk, target, full_km)?;
let p = rate_nm.mul(&share)?;
let ll = (unseen * (p + 1e-20)?.log()?)?;
Ok((ll.mul(&scored)?.sum(1)?, scored.sum(1)?))
}
fn unseen_parts(
&self,
log_theta_nk: &Tensor,
target: &ModuleTarget<'_>,
full_km: &Tensor,
) -> Result<(Tensor, Tensor, Tensor, Tensor)> {
let rate_nm = self.mixture_rate_nd(log_theta_nk, full_km)?;
let unseen = (target.values - target.visible_counts)?.clamp(0.0, f64::INFINITY)?;
let share = target.visible_share.affine(-1.0, 1.0)?;
let scored = share.gt(1e-6)?.to_dtype(rate_nm.dtype())?;
Ok((rate_nm, unseen, share, scored))
}
pub fn score_queries_nb(
&self,
log_theta_nk: &Tensor,
q: &QueryTarget<'_>,
full_km: &Tensor,
) -> Result<Tensor> {
let rate_nm = self.mixture_rate_nd(log_theta_nk, full_km)?; let ids_m = self.coarsening.groups_of(q.gene_ids)?; let rate_nq = rate_nm.gather(&ids_m, 1)?; let log_factor = (self.coarsening.log_share_at(q.gene_ids)? + q.log_residual)?;
let mu = rate_nq.mul(&log_factor.exp()?)?.broadcast_mul(q.lib)?;
let (n, qn) = ids_m.dims2()?;
let log_phi =
gather_rows(&self.log_phi_1d.squeeze(0)?, &ids_m.flatten_all()?)?.reshape((n, qn))?;
let elem = nb_log_likelihood_elem(q.values, &mu, &log_phi)?;
elem.mul(q.weight)?.sum(1)
}
}
fn nb_score(
values: &Tensor,
rate: &Tensor,
residual: Option<&Tensor>,
lib_n1: &Tensor,
log_phi: &Tensor,
weight: Option<&Tensor>,
) -> Result<Tensor> {
let mu = match residual {
Some(r) => rate.mul(r)?.broadcast_mul(lib_n1)?,
None => rate.broadcast_mul(lib_n1)?,
};
let elem = nb_log_likelihood_elem(values, &mu, log_phi)?;
weighted_row_sum(elem, weight)
}
fn multinomial_score(values: &Tensor, rate: &Tensor, weight: Option<&Tensor>) -> Result<Tensor> {
let ll = (values * (rate + 1e-20)?.log()?)?;
weighted_row_sum(ll, weight)
}
fn weighted_row_sum(elem: Tensor, weight: Option<&Tensor>) -> Result<Tensor> {
let last = elem.rank() - 1;
match weight {
Some(w) => elem.mul(w)?.sum(last),
None => elem.sum(last),
}
}
#[cfg(test)]
#[path = "masked_etm_tests.rs"]
mod masked_etm_tests;