use crate::candle::loss::{gaussian_kl_loss, gaussian_reparameterize};
use crate::candle::nn::layers::*;
use crate::candle::traits::model::*;
use crate::candle::value_transform::anscombe_residual;
use candle_core::{Result, Tensor};
use candle_nn::{ops, BatchNorm, Linear, Module, ModuleT, VarBuilder};
pub struct LogSoftmaxJointEncoder {
n_topics: usize,
fc: Vec<StackLayers<Linear>>,
bn_z: Vec<BatchNorm>,
z_mean: Vec<Linear>,
z_lnvar: Vec<Linear>,
feature_mean: Vec<Option<Tensor>>,
}
impl JointEncoderModuleT for LogSoftmaxJointEncoder {
fn forward_t(
&self,
x_nd_vec: &[Tensor],
x0_nd_vec: &[Option<Tensor>],
train: bool,
) -> Result<(Tensor, Tensor)> {
let (z_nk, kl) = self.latent_gaussian_with_kl(x_nd_vec, x0_nd_vec, train)?;
let log_prob = ops::log_softmax(&z_nk, z_nk.rank() - 1)?;
Ok((log_prob, kl))
}
fn dim_latent(&self) -> usize {
self.n_topics
}
}
impl LogSoftmaxJointEncoder {
fn preprocess_input(
&self,
x_nd_vec: &[Tensor],
x0_nd_vec: &[Option<Tensor>],
_train: bool,
) -> Result<Vec<Tensor>> {
x_nd_vec
.iter()
.zip(x0_nd_vec)
.zip(&self.feature_mean)
.map(|((x_nd, x0_nd), mu_f)| anscombe_residual(x_nd, x0_nd.as_ref(), mu_f.as_ref()))
.collect()
}
fn latent_gaussian_with_kl(
&self,
x_nd_vec: &[Tensor],
x0_nd_vec: &[Option<Tensor>],
train: bool,
) -> Result<(Tensor, Tensor)> {
let hh_vec = self
.preprocess_input(x_nd_vec, x0_nd_vec, train)?
.into_iter()
.enumerate()
.map(|(m, xx_nd)| -> Result<Tensor> {
let fc_nl = self.fc[m].forward_t(&xx_nd, train)?;
self.bn_z[m].forward_t(&fc_nl, train)
})
.collect::<Result<Vec<_>>>()?;
let z_kl_vec = hh_vec
.into_iter()
.enumerate()
.map(|(m, hh)| -> Result<(Tensor, Tensor)> {
let z_mean_nk = self.z_mean[m].forward(&hh)?;
let z_lnvar_nk = self.z_lnvar[m].forward(&hh)?;
let z = gaussian_reparameterize(&z_mean_nk, &z_lnvar_nk, train)?;
let kl = gaussian_kl_loss(&z_mean_nk, &z_lnvar_nk)?;
Ok((z, kl))
})
.collect::<Result<Vec<_>>>()?;
let z_aux_dim = z_kl_vec[0].0.rank();
let kl_aux_dim = z_kl_vec[0].1.rank();
let z = Tensor::cat(
&z_kl_vec
.iter()
.map(|(z, _)| -> Result<_> { z.unsqueeze(z_aux_dim) })
.collect::<Result<Vec<_>>>()?,
z_aux_dim,
)?
.sum(z_aux_dim)?;
let kl = Tensor::cat(
&z_kl_vec
.iter()
.map(|(_, kl)| -> Result<_> { kl.unsqueeze(kl_aux_dim) })
.collect::<Result<Vec<_>>>()?,
kl_aux_dim,
)?
.sum(kl_aux_dim)?;
Ok((z, kl))
}
pub fn new(args: LogSoftmaxJointEncoderArgs, vb: VarBuilder) -> Result<Self> {
debug_assert!(!args.layers.is_empty());
let bn_config = candle_nn::BatchNormConfig {
eps: 1e-4,
remove_mean: true,
affine: true,
momentum: 0.1,
};
let n_features = args.n_features.clone();
let n_modalities = n_features.len();
let fc_dims = args.layers[..args.layers.len() - 1].to_vec();
let out_dim = *args.layers.last().unwrap();
let fc = n_features
.iter()
.enumerate()
.map(|(i, &in_dim)| {
stack_relu_linear(in_dim, out_dim, &fc_dims, vb.pp(format!("nn.enc.fc_{}", i)))
})
.collect::<Result<Vec<_>>>()?;
let bn_z = (0..n_modalities)
.map(|i| candle_nn::batch_norm(out_dim, bn_config, vb.pp(format!("nn.enc.bn_z_{}", i))))
.collect::<Result<Vec<_>>>()?;
let z_mean = (0..n_modalities)
.map(|i| {
candle_nn::linear(
out_dim,
args.n_topics,
vb.pp(format!("nn.enc.z.mean_{}", i)),
)
})
.collect::<Result<Vec<_>>>()?;
let z_lnvar = (0..n_modalities)
.map(|i| {
candle_nn::linear(
out_dim,
args.n_topics,
vb.pp(format!("nn.enc.z.lnvar_{}", i)),
)
})
.collect::<Result<Vec<_>>>()?;
let dev = vb.device();
let feature_mean: Vec<Option<Tensor>> = match &args.feature_mean {
Some(per_mod) => {
debug_assert_eq!(per_mod.len(), n_modalities);
per_mod
.iter()
.zip(&n_features)
.map(|(opt, &d)| match opt {
Some(s) => {
debug_assert_eq!(s.len(), d);
Tensor::from_slice(s, (1, d), dev).map(Some)
}
None => Ok(None),
})
.collect::<Result<Vec<_>>>()?
}
None => vec![None; n_modalities],
};
Ok(Self {
n_topics: args.n_topics,
fc,
bn_z,
z_mean,
z_lnvar,
feature_mean,
})
}
}
pub struct LogSoftmaxJointEncoderArgs<'a> {
pub n_features: Vec<usize>,
pub n_topics: usize,
pub layers: &'a [usize],
pub feature_mean: Option<Vec<Option<&'a [f32]>>>,
}