use candle_core::{Result, Tensor};
use candle_nn::VarBuilder;
pub trait EncoderModuleT {
fn forward_t(
&self,
x_nd: &Tensor,
x0_nd: Option<&Tensor>,
train: bool,
) -> Result<(Tensor, Tensor)>;
fn dim_latent(&self) -> usize;
}
pub trait JointEncoderModuleT {
fn forward_t(
&self,
x_nd_vec: &[Tensor],
x0_nd_vec: &[Option<Tensor>],
train: bool,
) -> Result<(Tensor, Tensor)>;
fn dim_latent(&self) -> usize;
}
pub trait DecoderModuleT {
fn forward(&self, z_nk: &Tensor) -> Result<Tensor>;
fn get_dictionary(&self) -> Result<Tensor>;
fn forward_with_llik<LlikFn>(
&self,
z_nk: &Tensor,
x_nd: &Tensor,
llik: &LlikFn,
) -> Result<(Tensor, Tensor)>
where
LlikFn: Fn(&Tensor, &Tensor) -> Result<Tensor>;
fn llik_is_gene_chunked(&self) -> bool {
false
}
fn llik_gene_chunked(&self, z_nk: &Tensor, x_nd: &Tensor, gene_chunk: usize) -> Result<Tensor> {
let _ = gene_chunk;
let dense = |_: &Tensor, _: &Tensor| -> Result<Tensor> {
candle_core::bail!("the dense fallback needs no likelihood closure")
};
let (_, llik) = self.forward_with_llik(z_nk, x_nd, &dense)?;
Ok(llik)
}
fn dim_obs(&self) -> usize;
fn dim_latent(&self) -> usize;
fn attach_feature_weights(
&mut self,
_weights: &[f32],
_dev: &candle_core::Device,
) -> Result<()> {
Ok(())
}
}
pub trait NewDecoder: Sized {
fn new(n_features: usize, n_topics: usize, vs: VarBuilder) -> Result<Self>;
}
pub fn joint_multinomial_llik(
log_recon_vec: Vec<Tensor>,
x_nd_vec: &[Tensor],
) -> Result<(Vec<Tensor>, Tensor)> {
let recon_vec: Vec<Tensor> = log_recon_vec
.iter()
.map(|x| x.exp())
.collect::<Result<Vec<_>>>()?;
let llik_vec = x_nd_vec
.iter()
.zip(&log_recon_vec)
.map(|(x, log_recon)| -> Result<Tensor> {
let ret = x
.clamp(0.0, f64::INFINITY)?
.mul(log_recon)?
.sum(x.rank() - 1)?;
ret.unsqueeze(ret.rank())
})
.collect::<Result<Vec<Tensor>>>()?;
let k = llik_vec[0].rank();
let llik = Tensor::cat(&llik_vec, k - 1)?.sum(k - 1)?;
Ok((recon_vec, llik))
}
pub trait JointDecoderModuleT {
fn forward(&self, z_nk: &Tensor) -> Result<Vec<Tensor>>;
fn get_dictionary(&self) -> Result<Vec<Tensor>>;
fn forward_with_llik<LlikFn>(
&self,
z_nk: &Tensor,
x_nd: &[Tensor],
llik: &LlikFn,
) -> Result<(Vec<Tensor>, Tensor)>
where
LlikFn: Fn(&Tensor, &Tensor) -> Result<Tensor>;
fn dim_obs(&self) -> &[usize];
fn dim_latent(&self) -> usize;
}