legume-numeric 0.8.11

Numeric and ML foundation for the legume ecosystem (matrix, Leiden, candle, MCMC)
Documentation
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>,
    /// Per-modality `[1, D_m]` non-trainable mean rate, broadcast inside
    /// `anscombe_residual` as a multiplicative count-rate divisor.
    feature_mean: Vec<Option<Tensor>>,
}

impl JointEncoderModuleT for LogSoftmaxJointEncoder {
    /// Returns (log_prob, kl) where log_prob is log-probabilities on the simplex
    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()
    }

    ///
    /// Evaluate latent Gaussian parameters: mu and log_var
    /// * `z ~ (mu(x), exp(log_var(x)))`
    /// * `kl` divergence
    ///
    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))
    }

    /// Will create a new non-negative encoder module
    ///
    /// # Arguments
    /// * `args` - encoder arguments
    /// * `vb` - variable builder
    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();

        // data [N, D_m] -> fc stack -> final_hidden (per modality)
        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<_>>>()?;

        // fc -> K
        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],
    /// Optional per-modality mean rates `μ_d` (each length = `n_features[m]`).
    /// Outer `Option` for "no means at all"; inner `Option` per modality.
    pub feature_mean: Option<Vec<Option<&'a [f32]>>>,
}