use crate::candle::loss::{gaussian_kl_loss, gaussian_reparameterize};
use crate::candle::nn::batch_norm;
use crate::candle::nn::layers::*;
use crate::candle::traits::model::*;
use crate::candle::value_transform::anscombe_residual;
use candle_core::{Device, Result, Tensor};
use candle_nn::{ops, Linear, ModuleT, VarBuilder, VarMap};
pub struct LogSoftmaxEncoder {
n_topics: usize,
fc: StackLayers<Linear>,
bn_z: batch_norm::BatchNorm,
z_mean: Linear,
z_lnvar: Linear,
feature_mean: Option<Tensor>,
}
impl EncoderModuleT for LogSoftmaxEncoder {
fn forward_t(
&self,
x_nd: &Tensor,
x0_nd: Option<&Tensor>,
train: bool,
) -> Result<(Tensor, Tensor)> {
let (z_mean_nk, z_lnvar_nk) = self.latent_gaussian_params(x_nd, x0_nd, train)?;
let z_nk = gaussian_reparameterize(&z_mean_nk, &z_lnvar_nk, train)?;
let log_prob = ops::log_softmax(&z_nk, 1)?;
Ok((log_prob, gaussian_kl_loss(&z_mean_nk, &z_lnvar_nk)?))
}
fn dim_latent(&self) -> usize {
self.n_topics
}
}
impl LogSoftmaxEncoder {
fn preprocess_input(
&self,
x_nd: &Tensor,
x0_nd: Option<&Tensor>,
_train: bool,
) -> Result<Tensor> {
anscombe_residual(x_nd, x0_nd, self.feature_mean.as_ref())
}
pub fn latent_gaussian_params(
&self,
x_nd: &Tensor,
x0_nd: Option<&Tensor>,
train: bool,
) -> Result<(Tensor, Tensor)> {
let clamp_lo = -8.;
let clamp_hi = 8.;
let xx_nd = self.preprocess_input(x_nd, x0_nd, train)?;
let fc_nl = self.fc.forward_t(&xx_nd, train)?;
let bn_nl = self.bn_z.forward_t(&fc_nl, train)?;
let z_mean_nk = self
.z_mean
.forward_t(&bn_nl, train)?
.clamp(clamp_lo, clamp_hi)?;
let z_lnvar_nk = self
.z_lnvar
.forward_t(&bn_nl, train)?
.clamp(clamp_lo, clamp_hi)?;
Ok((z_mean_nk, z_lnvar_nk))
}
pub fn new(args: LogSoftmaxEncoderArgs, varmap: &VarMap, vb: VarBuilder) -> Result<Self> {
let bn_config = batch_norm::BatchNormConfig {
eps: 1e-4,
remove_mean: true,
affine: true,
momentum: 0.1,
};
debug_assert!(!args.layers.is_empty());
let fc_dims = args.layers[..args.layers.len() - 1].to_vec();
let in_dim = args.n_features;
let out_dim = *args.layers.last().unwrap();
let fc = stack_relu_linear(in_dim, out_dim, &fc_dims, vb.pp("nn.enc.fc"))?;
let bn_z = batch_norm::batch_norm(out_dim, bn_config, varmap, vb.pp("nn.enc.bn_z"))?;
let z_mean = candle_nn::linear(out_dim, args.n_topics, vb.pp("nn.enc.z.mean"))?;
let z_lnvar = candle_nn::linear(out_dim, args.n_topics, vb.pp("nn.enc.z.lnvar"))?;
let feature_mean = match args.feature_mean {
Some(slice) => {
debug_assert_eq!(
slice.len(),
args.n_features,
"feature_mean length must equal n_features"
);
Some(host_to_row_tensor(slice, vb.device())?)
}
None => None,
};
Ok(Self {
n_topics: args.n_topics,
fc,
bn_z,
z_mean,
z_lnvar,
feature_mean,
})
}
}
fn host_to_row_tensor(slice: &[f32], dev: &Device) -> Result<Tensor> {
let d = slice.len();
Tensor::from_slice(slice, (1, d), dev)
}
pub struct LogSoftmaxEncoderArgs<'a> {
pub n_features: usize,
pub n_topics: usize,
pub layers: &'a [usize],
pub feature_mean: Option<&'a [f32]>,
}