#![allow(dead_code)]
use crate::candle::loss::nb_log_likelihood;
use crate::candle::nn::linear::*;
use crate::candle::traits::model::*;
use candle_core::{Result, Tensor};
use candle_nn::{ops, Module, VarBuilder};
pub const DECODER_NAME: &str = "nbmixture";
pub struct NbMixtureTopicDecoder {
n_features: usize,
n_topics: usize,
dictionary: SoftmaxLinear,
log_phi_1d: Tensor,
log_alpha_1d: Tensor,
rho_a: Tensor,
rho_b: Tensor,
rho_prior_weight: f32,
rho_prior_alpha: f32,
rho_prior_beta: f32,
}
impl NbMixtureTopicDecoder {
pub fn new(n_features: usize, n_topics: usize, vs: VarBuilder) -> Result<Self> {
let dictionary = log_softmax_linear(n_topics, n_features, vs.pp("dictionary"))?;
let log_phi_1d =
vs.get_with_hints((1, n_features), "log_phi", candle_nn::Init::Const(0.693))?;
let log_alpha_1d =
vs.get_with_hints((1, n_features), "log_alpha", candle_nn::Init::Const(0.0))?;
let rho_a = vs.get_with_hints((1, 1), "rho_a", candle_nn::Init::Const(-0.5))?;
let rho_b = vs.get_with_hints((1, 1), "rho_b", candle_nn::Init::Const(0.0))?;
Ok(Self {
n_features,
n_topics,
dictionary,
log_phi_1d,
log_alpha_1d,
rho_a,
rho_b,
rho_prior_weight: 0.0,
rho_prior_alpha: 2.0,
rho_prior_beta: 18.0,
})
}
pub fn set_rho_prior(&mut self, weight: f32, alpha: f32, beta: f32) {
self.rho_prior_weight = weight;
self.rho_prior_alpha = alpha;
self.rho_prior_beta = beta;
}
pub fn rho_prior_weight(&self) -> f32 {
self.rho_prior_weight
}
pub fn rho_prior_alpha(&self) -> f32 {
self.rho_prior_alpha
}
pub fn rho_prior_beta(&self) -> f32 {
self.rho_prior_beta
}
pub fn phi(&self) -> Result<Tensor> {
self.log_phi_1d.exp()
}
pub fn log_phi(&self) -> &Tensor {
&self.log_phi_1d
}
pub fn log_alpha(&self) -> &Tensor {
&self.log_alpha_1d
}
pub fn alpha(&self) -> Result<Tensor> {
let r = self.log_alpha_1d.rank();
ops::log_softmax(&self.log_alpha_1d, r - 1)?.exp()
}
pub fn rho_a(&self) -> &Tensor {
&self.rho_a
}
pub fn rho_b(&self) -> &Tensor {
&self.rho_b
}
fn rho_prior_llik(&self, rho_n1: &Tensor, one_minus_rho: &Tensor) -> Result<Option<Tensor>> {
if self.rho_prior_weight <= 0.0 {
return Ok(None);
}
let eps = 1e-6f64;
let log_rho = (rho_n1 + eps)?.log()?;
let log_1m_rho = (one_minus_rho + eps)?.log()?;
let term = ((log_rho * f64::from(self.rho_prior_alpha - 1.0))?
+ (log_1m_rho * f64::from(self.rho_prior_beta - 1.0))?)?;
Ok(Some((term.squeeze(1)? * f64::from(self.rho_prior_weight))?))
}
pub fn rho_from_lib(&self, lib_n1: &Tensor) -> Result<Tensor> {
let log_lib = (lib_n1 + 1e-8)?.log()?;
let z = log_lib
.broadcast_mul(&self.rho_a)?
.broadcast_add(&self.rho_b)?;
ops::sigmoid(&z)
}
}
impl NewDecoder for NbMixtureTopicDecoder {
fn new(n_features: usize, n_topics: usize, vs: VarBuilder) -> Result<Self> {
NbMixtureTopicDecoder::new(n_features, n_topics, vs)
}
}
impl DecoderModuleT for NbMixtureTopicDecoder {
fn forward(&self, z_nk: &Tensor) -> Result<Tensor> {
self.dictionary.forward(z_nk)
}
fn get_dictionary(&self) -> Result<Tensor> {
self.dictionary.weight_dk()
}
fn forward_with_llik<LlikFn>(
&self,
z_nk: &Tensor,
x_nd: &Tensor,
_llik: &LlikFn,
) -> Result<(Tensor, Tensor)>
where
LlikFn: Fn(&Tensor, &Tensor) -> Result<Tensor>,
{
let last_dim = x_nd.rank() - 1;
let topic_recon_nd = self.dictionary.forward(z_nk)?; let alpha_1d = self.alpha()?;
let lib_n1 = x_nd.sum(last_dim)?.unsqueeze(1)?; let rho_n1 = self.rho_from_lib(&lib_n1)?;
let one_minus_rho = rho_n1.affine(-1.0, 1.0)?;
let pi_nd = topic_recon_nd
.broadcast_mul(&one_minus_rho)?
.broadcast_add(&alpha_1d.broadcast_mul(&rho_n1)?)?;
let mu_nd = pi_nd.broadcast_mul(&lib_n1)?;
let data_llik = nb_log_likelihood(x_nd, &mu_nd, &self.log_phi_1d)?;
let llik = match self.rho_prior_llik(&rho_n1, &one_minus_rho)? {
Some(prior) => (data_llik + prior)?,
None => data_llik,
};
Ok((pi_nd, llik))
}
fn llik_is_gene_chunked(&self) -> bool {
true
}
fn llik_gene_chunked(&self, z_nk: &Tensor, x_nd: &Tensor, gene_chunk: usize) -> Result<Tensor> {
let last = x_nd.rank() - 1;
let chunk = gene_chunk.max(1).min(self.n_features);
let lib_n1 = x_nd.sum(last)?.unsqueeze(1)?; let rho_n1 = self.rho_from_lib(&lib_n1)?;
let one_minus_rho = rho_n1.affine(-1.0, 1.0)?;
let alpha_1d = self.alpha()?;
let log_w_kd = self.dictionary.log_weight_kd()?;
let mut data_llik: Option<Tensor> = None;
for (start, len) in super::gene_slices(self.n_features, chunk) {
let topic = self
.dictionary
.forward_log_slice(z_nk, Some(&log_w_kd), start, len)?
.exp()?;
let ambient = alpha_1d.narrow(last, start, len)?.broadcast_mul(&rho_n1)?;
let pi = topic
.broadcast_mul(&one_minus_rho)?
.broadcast_add(&ambient)?;
let mu = pi.broadcast_mul(&lib_n1)?;
let x = x_nd.narrow(last, start, len)?;
let log_phi = self
.log_phi_1d
.narrow(1, start, len)?
.broadcast_as(x.shape())?;
let part = crate::candle::loss::nb_log_likelihood_elem(&x, &mu, &log_phi)?.sum(last)?;
data_llik = Some(match data_llik.take() {
Some(acc) => acc.add(&part)?,
None => part,
});
}
let data_llik =
data_llik.ok_or_else(|| candle_core::Error::Msg("no gene slices to score".into()))?;
match self.rho_prior_llik(&rho_n1, &one_minus_rho)? {
Some(prior) => data_llik + prior,
None => Ok(data_llik),
}
}
fn dim_obs(&self) -> usize {
self.n_features
}
fn dim_latent(&self) -> usize {
self.n_topics
}
fn build_ess_llik<'a>(
&'a self,
x_nd: &'a Tensor,
topic_smoothing: f64,
) -> Result<EssLlikFn<'a>> {
let last_dim = x_nd.rank() - 1;
let log_dict_dk = self.get_dictionary()?.detach();
let beta_kd = log_dict_dk.t()?.exp()?.contiguous()?;
let log_phi = self.log_phi_1d.detach();
let log_alpha_det = self.log_alpha_1d.detach();
let alpha_1d = {
let r = log_alpha_det.rank();
ops::log_softmax(&log_alpha_det, r - 1)?.exp()?
};
let rho_a = self.rho_a.detach();
let rho_b = self.rho_b.detach();
let lib_n1 = x_nd.sum(last_dim)?.unsqueeze(1)?;
let log_lib = (&lib_n1 + 1e-8)?.log()?;
let rho_n1 = ops::sigmoid(&log_lib.broadcast_mul(&rho_a)?.broadcast_add(&rho_b)?)?;
let one_minus_rho = rho_n1.affine(-1.0, 1.0)?;
let k = self.dim_latent() as f64;
Ok(Box::new(move |z_nk: &Tensor| {
let mut z = ops::softmax(z_nk, 1)?;
if topic_smoothing > 0.0 {
z = ((z * (1.0 - topic_smoothing))? + topic_smoothing / k)?;
}
let topic_recon = z.matmul(&beta_kd)?;
let pi = topic_recon
.broadcast_mul(&one_minus_rho)?
.broadcast_add(&alpha_1d.broadcast_mul(&rho_n1)?)?;
let mu = pi.broadcast_mul(&lib_n1)?;
nb_log_likelihood(x_nd, &mu, &log_phi)
}))
}
}
#[cfg(test)]
#[path = "nb_mixture_tests.rs"]
mod chunked_llik_tests;