use core::f64;
use candle_core::{Result, Tensor};
use candle_nn::ops;
pub fn gaussian_kl_loss(z_mean: &Tensor, z_lnvar: &Tensor) -> Result<Tensor> {
let z_var = z_lnvar.exp()?;
(z_var - 1. + z_mean.powf(2.)? - z_lnvar)?.sum(z_mean.rank() - 1)? * 0.5
}
pub fn gaussian_reparameterize(z_mean: &Tensor, z_lnvar: &Tensor, train: bool) -> Result<Tensor> {
if train {
let eps = Tensor::randn_like(z_mean, 0., 1.)?;
z_mean + (z_lnvar * 0.5)?.exp()? * eps
} else {
Ok(z_mean.clone())
}
}
pub fn gaussian_neg_log_prob(z: &Tensor, mean: &Tensor, lnvar: &Tensor) -> Result<Tensor> {
let var = lnvar.exp()?;
let diff = (z - mean)?;
((&diff * &diff)?.div(&var)? + lnvar)?.sum(z.rank() - 1)? * 0.5
}
pub fn topic_likelihood(x_nd: &Tensor, recon_nd: &Tensor) -> Result<Tensor> {
let eps = 1e-8;
let log_recon_nd = (recon_nd + eps)?.log()?;
x_nd.clamp(0.0, f64::INFINITY)?
.mul(&log_recon_nd)?
.sum(x_nd.rank() - 1)
}
pub fn multinomial_nll_profiled(x: &Tensor, s: &Tensor, totals: &Tensor) -> Result<Tensor> {
let lse = s.log_sum_exp(1)?;
let data = (x * s)?.sum(1)?;
(totals * lse)? - data
}
pub fn topic_log_likelihood(x_nd: &Tensor, log_recon_nd: &Tensor) -> Result<Tensor> {
x_nd.clamp(0.0, f64::INFINITY)?
.mul(log_recon_nd)?
.sum(x_nd.rank() - 1)
}
pub fn dirichlet_likelihood(x_nd: &Tensor, mass_nd: &Tensor) -> Result<Tensor> {
let a_nd = x_nd.add(mass_nd)?;
let term1 = approx_lgamma(&a_nd)?
.sub(&approx_lgamma(mass_nd)?)?
.sum(a_nd.rank() - 1)?;
let term2 = approx_lgamma(&mass_nd.sum(mass_nd.rank() - 1)?)?
.sub(&approx_lgamma(&a_nd.sum(a_nd.rank() - 1)?)?)?;
term1.add(&term2)
}
pub fn approx_lgamma(x: &Tensor) -> Result<Tensor> {
let term1 = (x.neg()? - 0.0810614667)?;
let term2 = x.log()?.neg()?;
let term3 = (x + 0.5)?.mul(&(x + 1.0)?.log()?)?;
term1.add(&term2)?.add(&term3)
}
pub fn poisson_likelihood(x_nd: &Tensor, rate_nd: &Tensor) -> Result<Tensor> {
x_nd.mul(&rate_nd.log()?)?
.sub(rate_nd)?
.sum(x_nd.rank() - 1)
}
pub fn zi_topic_log_likelihood(
x_nd: &Tensor,
log_recon_nd: &Tensor,
dropout_logit_1d: &Tensor,
) -> Result<Tensor> {
let pi = ops::sigmoid(dropout_logit_1d)?;
let one_minus_pi = (1.0 - &pi)?;
let is_zero = x_nd.eq(0.0f64)?;
let eps = 1e-20;
let log_pi = (&pi + eps)?.log()?;
let log_one_minus_pi = (&one_minus_pi + eps)?.log()?;
let log_term2 = log_one_minus_pi.broadcast_add(log_recon_nd)?; let max_val = log_pi.broadcast_maximum(&log_term2)?;
let sum_exp = log_pi
.broadcast_sub(&max_val)?
.exp()?
.add(&log_term2.broadcast_sub(&max_val)?.exp()?)?;
let zero_llik = sum_exp.log()?.add(&max_val)?;
let nonzero_llik = log_one_minus_pi.broadcast_add(&x_nd.mul(log_recon_nd)?)?;
is_zero
.where_cond(&zero_llik, &nonzero_llik)?
.sum(x_nd.rank() - 1)
}
pub fn zi_topic_likelihood(
x_nd: &Tensor,
recon_nd: &Tensor,
dropout_logit_1d: &Tensor,
) -> Result<Tensor> {
let eps = 1e-20;
let log_recon_nd = (recon_nd + eps)?.log()?;
zi_topic_log_likelihood(x_nd, &log_recon_nd, dropout_logit_1d)
}
pub fn nb_log_likelihood(x_nd: &Tensor, mu_nd: &Tensor, log_phi_1d: &Tensor) -> Result<Tensor> {
let log_phi = log_phi_1d.broadcast_as(x_nd.shape())?;
nb_log_likelihood_elem(x_nd, mu_nd, &log_phi)?.sum(x_nd.rank() - 1)
}
pub fn nb_log_likelihood_elem(x: &Tensor, mu: &Tensor, log_phi: &Tensor) -> Result<Tensor> {
let phi = log_phi.clamp(-10.0, 10.0)?.exp()?;
let mu = mu.clamp(1e-6, 1e6)?;
let eps = 1e-8;
let phi_plus_mu = (&phi + &mu)?;
let log_phi = (&phi + eps)?.log()?;
let log_phi_plus_mu = (&phi_plus_mu + eps)?.log()?;
let log_mu = (&mu + eps)?.log()?;
let term_phi = phi.mul(&(&log_phi - &log_phi_plus_mu)?)?;
let term_x = x.mul(&(&log_mu - &log_phi_plus_mu)?)?;
let x_plus_phi = (x + &phi)?;
let lgamma_term = approx_lgamma(&x_plus_phi)?
.sub(&approx_lgamma(&phi)?)?
.sub(&approx_lgamma(&(x + 1.0)?)?)?;
lgamma_term + term_phi + term_x
}
pub fn log_sigmoid(x: &Tensor) -> Result<Tensor> {
let min_x0 = x.neg()?.relu()?.neg()?; let softplus_tail = (x.abs()?.neg()?.exp()? + 1.0)?.log()?; min_x0.sub(&softplus_tail)
}
pub fn gaussian_likelihood(x_nd: &Tensor, hat_nd: &Tensor) -> Result<Tensor> {
x_nd.sub(hat_nd)?.powf(2.)?.sum(1)? * (-0.5)
}