use candle_core::{Result, Tensor};
use crate::candle::sgvb::BlackBoxLikelihood;
pub struct PoissonLikelihood {
y: Tensor,
}
impl PoissonLikelihood {
pub fn new(y: Tensor) -> Self {
Self { y }
}
}
impl BlackBoxLikelihood for PoissonLikelihood {
fn log_likelihood(&self, etas: &[&Tensor]) -> Result<Tensor> {
let eta = etas[0];
let y_eta = eta.broadcast_mul(&self.y)?;
let exp_eta = eta.exp()?;
let log_prob = (y_eta - exp_eta)?;
log_prob.sum(2)?.sum(1)
}
}
pub struct OffsetPoissonLikelihood {
y: Tensor,
offset: Tensor,
}
impl OffsetPoissonLikelihood {
pub fn new(y: Tensor, offset: f32) -> candle_core::Result<Self> {
let device = y.device().clone();
let (n, k) = (y.dim(0)?, y.dim(1)?);
let offset = Tensor::full(offset, (n, k), &device)?;
Ok(Self { y, offset })
}
}
impl BlackBoxLikelihood for OffsetPoissonLikelihood {
fn log_likelihood(&self, etas: &[&Tensor]) -> Result<Tensor> {
let eta = etas[0];
let eta_full = eta.broadcast_add(&self.offset)?;
let y_eta = eta_full.broadcast_mul(&self.y)?;
let exp_eta = eta_full.exp()?;
let log_prob = (y_eta - exp_eta)?;
log_prob.sum(2)?.sum(1)
}
}
#[cfg(test)]
mod tests {
use super::*;
use candle_core::Device;
#[test]
fn test_poisson_likelihood() -> Result<()> {
let device = Device::Cpu;
let y = Tensor::from_vec(vec![1.0f32, 2.0, 3.0], (3, 1), &device)?;
let eta = Tensor::from_vec(vec![0.0f32, 0.5, 1.0], (1, 3, 1), &device)?;
let likelihood = PoissonLikelihood::new(y);
let log_lik = likelihood.log_likelihood(&[&eta])?;
let val: f32 = log_lik.get(0)?.to_scalar()?;
assert!(val.is_finite());
println!("Poisson log_lik: {}", val);
Ok(())
}
}