use crate::candle::loss::nb_log_likelihood;
use crate::candle::traits::model::*;
use candle_core::{Result, Tensor};
use candle_nn::{ops, Linear, Module, VarBuilder};
pub struct GaussianNbDecoder {
n_features: usize,
n_latent: usize,
decoder: Linear,
log_phi_1d: Tensor,
}
impl GaussianNbDecoder {
pub fn new(n_features: usize, n_latent: usize, vs: VarBuilder) -> Result<Self> {
let decoder = candle_nn::linear(n_latent, n_features, vs.pp("gauss_decoder"))?;
let log_phi_1d =
vs.get_with_hints((1, n_features), "log_phi", candle_nn::Init::Const(0.693))?;
Ok(Self {
n_features,
n_latent,
decoder,
log_phi_1d,
})
}
#[must_use]
pub fn feature_bias(&self) -> Option<Tensor> {
self.decoder.bias().cloned()
}
fn log_pi(&self, z_nk: &Tensor) -> Result<Tensor> {
let logits_nd = self.decoder.forward(z_nk)?; ops::log_softmax(&logits_nd, logits_nd.rank() - 1)
}
}
impl NewDecoder for GaussianNbDecoder {
fn new(n_features: usize, n_latent: usize, vs: VarBuilder) -> Result<Self> {
GaussianNbDecoder::new(n_features, n_latent, vs)
}
}
impl DecoderModuleT for GaussianNbDecoder {
fn forward(&self, z_nk: &Tensor) -> Result<Tensor> {
self.log_pi(z_nk)?.exp()
}
fn get_dictionary(&self) -> Result<Tensor> {
Ok(self.decoder.weight().clone())
}
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 = x_nd.rank() - 1;
let logits_nd = self.decoder.forward(z_nk)?; let pi_nd = ops::softmax(&logits_nd, last)?;
let lib_n1 = x_nd.sum_keepdim(last)?; let mu_nd = pi_nd.broadcast_mul(&lib_n1)?;
let llik = nb_log_likelihood(x_nd, &mu_nd, &self.log_phi_1d)?;
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 chunk = gene_chunk.max(1).min(self.n_features);
let last = x_nd.rank() - 1;
let lib_n1 = x_nd.sum_keepdim(last)?; let w_dk = self.decoder.weight();
let bias_d = self.decoder.bias();
let logits_of = |start: usize, len: usize| -> Result<Tensor> {
let w = w_dk.narrow(0, start, len)?; let l = z_nk.matmul(&w.t()?)?; match bias_d {
Some(b) => l.broadcast_add(&b.narrow(0, start, len)?.unsqueeze(0)?),
None => Ok(l),
}
};
let mut running_max: Option<Tensor> = None;
let mut running_sum: Option<Tensor> = None;
for (start, len) in super::gene_slices(self.n_features, chunk) {
let logits = logits_of(start, len)?;
let m = logits.max_keepdim(last)?; let (new_max, sum) = match (running_max.take(), running_sum.take()) {
(Some(pm), Some(ps)) => {
let new_max = pm.maximum(&m)?;
let rescaled = ps.mul(&pm.sub(&new_max)?.exp()?)?;
let add = logits.broadcast_sub(&new_max)?.exp()?.sum_keepdim(last)?;
(new_max, rescaled.add(&add)?)
}
_ => {
let sum = logits.broadcast_sub(&m)?.exp()?.sum_keepdim(last)?;
(m, sum)
}
};
running_max = Some(new_max);
running_sum = Some(sum);
}
let max_n1 = running_max.expect("at least one gene slice");
let log_denom = running_sum
.expect("at least one gene slice")
.log()?
.add(&max_n1)?;
let mut llik: Option<Tensor> = None;
for (start, len) in super::gene_slices(self.n_features, chunk) {
let logits = logits_of(start, len)?;
let pi = logits.broadcast_sub(&log_denom)?.exp()?; 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)?;
llik = Some(match llik.take() {
Some(acc) => acc.add(&part)?,
None => part,
});
}
llik.ok_or_else(|| candle_core::Error::Msg("no gene slices to score".into()))
}
fn dim_obs(&self) -> usize {
self.n_features
}
fn dim_latent(&self) -> usize {
self.n_latent
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::candle::encoder::{GaussianEncoder, GaussianEncoderArgs};
use candle_core::{DType, Device};
use candle_nn::{VarBuilder, VarMap};
#[test]
fn test_gaussian_encoder_decoder_smoke() {
let dev = Device::Cpu;
let varmap = VarMap::new();
let vb = VarBuilder::from_varmap(&varmap, DType::F32, &dev);
let (n, d, h, k) = (4usize, 6usize, 8usize, 3usize);
let enc = GaussianEncoder::new(
GaussianEncoderArgs {
n_features: d,
n_latent: k,
layers: &[h],
feature_mean: None,
},
&varmap,
vb.pp("enc"),
)
.unwrap();
let dec = GaussianNbDecoder::new(d, k, vb.pp("dec")).unwrap();
let x = Tensor::rand(0f32, 5f32, (n, d), &dev).unwrap();
let (z, kl) = enc.forward_t(&x, None, true).unwrap();
assert_eq!(z.dims(), &[n, k]);
assert_eq!(kl.dims(), &[n]);
let noop = |_a: &Tensor, _b: &Tensor| Ok(_a.clone());
let (pi, llik) = dec.forward_with_llik(&z, &x, &noop).unwrap();
assert_eq!(pi.dims(), &[n, d]);
assert_eq!(llik.dims(), &[n]);
for s in pi.sum(1).unwrap().to_vec1::<f32>().unwrap() {
assert!((s - 1.0).abs() < 1e-4, "pi row sum {s} != 1");
}
for v in llik.to_vec1::<f32>().unwrap() {
assert!(v.is_finite(), "llik {v} not finite");
}
}
}
#[cfg(test)]
#[path = "gaussian_nb_tests.rs"]
mod chunked_llik_tests;