use super::{clip_and_step_dense_all, smooth_topics, TrainScores};
use crate::candle::data::indexed::labeled_bar;
use crate::candle::data::masked_dense::{DenseMaskedLevel, DenseMaskedMinibatch, MaskedDraw};
use crate::candle::decoder::coarsening_map::CoarseningMap;
use crate::candle::decoder::masked_etm::{EmbeddedNbTopicDecoder, MaskedDenseTarget, ModuleTarget};
use crate::candle::encoder::indexed::IndexedEmbeddingEncoder;
pub use crate::candle::lora::LoraPlus;
use crate::matrix::rand_util::mix_seed;
use candle_core::{DType, Device, Tensor, Var};
use candle_nn::{AdamW, Optimizer};
use log::{info, warn};
use nalgebra::DMatrix;
use std::sync::atomic::{AtomicBool, Ordering};
pub use crate::candle::data::masked_dense::MaskSchedule;
type Mat = DMatrix<f32>;
pub struct IndexedTrainConfig<'a> {
pub parameters: &'a candle_nn::VarMap,
pub dev: &'a Device,
pub epochs: usize,
pub gpu_mem_fraction: Option<f32>,
pub minibatch_size: usize,
pub learning_rate: f32,
pub topic_smoothing: f64,
pub stop: &'a AtomicBool,
pub feature_mean: &'a [f32],
pub grad_clip: f32,
pub feature_embedding_l2: f32,
pub weight_decay: f32,
pub feature_anchor: Option<FeatureAnchor<'a>>,
}
#[derive(Clone, Copy, Debug)]
pub struct FeatureAnchor<'a> {
pub base_var: &'a str,
pub lora: Option<LoraPlus<'a>>,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum MaskedLikelihood {
Nb,
Multinomial,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum LatentHead {
Softmax,
StickBreaking,
Gaussian,
}
pub struct MaskedTrainOpts {
pub mask_schedule: MaskSchedule,
pub likelihood: MaskedLikelihood,
pub latent: LatentHead,
pub poisson_thin: bool,
pub seed: u64,
}
impl Default for MaskedTrainOpts {
fn default() -> Self {
Self {
mask_schedule: MaskSchedule::Fixed,
likelihood: MaskedLikelihood::Nb,
latent: LatentHead::Softmax,
poisson_thin: false,
seed: 42,
}
}
}
fn poisson_draw(rates: &Mat, seed: u64) -> Mat {
use rand::{rngs::SmallRng, SeedableRng};
use rand_distr::{Distribution, Poisson};
use rayon::prelude::*;
let nrows = rates.nrows();
let mut out = rates.clone();
out.as_mut_slice()
.par_chunks_mut(nrows.max(1))
.enumerate()
.for_each(|(col_idx, col): (usize, &mut [f32])| {
let mut rng =
SmallRng::seed_from_u64(crate::matrix::rand_util::mix_seed(seed, col_idx as u64));
for v in col.iter_mut() {
*v = if *v > 0.0 && v.is_finite() {
Poisson::new(f64::from(*v)).map_or(0.0, |p| p.sample(&mut rng) as f32)
} else {
0.0
};
}
});
out
}
#[must_use]
pub fn epoch_seed(seed: u64, epoch: usize, level: usize) -> u64 {
let salt = ((epoch as u64) << 32) | (level as u64);
mix_seed(crate::matrix::rand_util::name_seed(seed, "mask"), salt)
}
pub struct DenseModuleTargets {
pub values_nm: Tensor,
pub visible_counts_nm: Tensor,
pub visible_share_nm: Tensor,
pub lib_n1: Tensor,
}
pub fn dense_module_targets(
modules: &CoarseningMap,
target_nd: &Tensor,
visible_nd: &Tensor,
) -> candle_core::Result<DenseModuleTargets> {
let values_nm = modules.aggregate_columns(target_nd)?;
let visible_counts_nm = modules.aggregate_columns(&(target_nd * visible_nd)?)?;
let share_1d = modules.log_share_1d().exp()?; let visible_share_nm = modules.aggregate_columns(&visible_nd.broadcast_mul(&share_1d)?)?;
let lib_n1 = (values_nm.sum_keepdim(1)? + 1.0)?;
Ok(DenseModuleTargets {
values_nm,
visible_counts_nm,
visible_share_nm,
lib_n1,
})
}
pub struct MaskedEncoderInput<'a> {
pub indices: &'a Tensor,
pub values: &'a Tensor,
pub values_null: Option<&'a Tensor>,
pub values_mean: Option<&'a Tensor>,
pub visible_mask: &'a Tensor,
}
pub fn masked_encode(
encoder: &IndexedEmbeddingEncoder,
head: LatentHead,
input: &MaskedEncoderInput,
train: bool,
) -> candle_core::Result<Tensor> {
match head {
LatentHead::Gaussian => encoder.forward_indexed_masked_gaussian(
input.indices,
input.values,
input.values_null,
input.values_mean,
input.visible_mask,
train,
),
LatentHead::StickBreaking => encoder.forward_indexed_masked_stick(
input.indices,
input.values,
input.values_null,
input.values_mean,
input.visible_mask,
train,
),
LatentHead::Softmax => encoder.forward_indexed_masked(
input.indices,
input.values,
input.values_null,
input.values_mean,
input.visible_mask,
train,
),
}
}
pub struct MaskedDenseInput<'a> {
pub x_nd: &'a Tensor,
pub x0_nd: Option<&'a Tensor>,
pub mean_1d: Option<&'a Tensor>,
pub visible_nd: &'a Tensor,
}
pub fn masked_encode_dense(
encoder: &IndexedEmbeddingEncoder,
head: LatentHead,
input: &MaskedDenseInput,
train: bool,
) -> candle_core::Result<Tensor> {
match head {
LatentHead::Gaussian => encoder.forward_dense_masked_gaussian(
input.x_nd,
input.x0_nd,
input.mean_1d,
input.visible_nd,
train,
),
LatentHead::StickBreaking => encoder.forward_dense_masked_stick(
input.x_nd,
input.x0_nd,
input.mean_1d,
input.visible_nd,
train,
),
LatentHead::Softmax => encoder.forward_dense_masked(
input.x_nd,
input.x0_nd,
input.mean_1d,
input.visible_nd,
train,
),
}
}
pub fn decoder_log_theta(
raw_z: Tensor,
head: LatentHead,
topic_smoothing: f64,
) -> candle_core::Result<Tensor> {
let log_theta = match head {
LatentHead::Gaussian => candle_nn::ops::log_softmax(&raw_z, 1)?,
LatentHead::Softmax | LatentHead::StickBreaking => raw_z,
};
smooth_topics(log_theta, topic_smoothing)
}
pub type LevelData<'a> = (&'a Mat, Option<&'a Mat>, &'a Mat);
fn resident_dense_levels(
level_data: &[LevelData],
config: &IndexedTrainConfig,
dev: &Device,
) -> anyhow::Result<Vec<DenseMaskedLevel>> {
level_data
.iter()
.map(|&(mixed, batch, target)| {
DenseMaskedLevel::from_mats(mixed, batch, target, config.feature_mean, dev)
})
.collect()
}
struct StepLoss {
loss: Tensor,
llik_sum: Tensor,
units_sum: Tensor,
}
pub(crate) struct EpochAccum {
llik: Tensor,
scored: Tensor,
}
impl EpochAccum {
pub(crate) fn new(dev: &Device) -> candle_core::Result<Self> {
let zero = || Tensor::zeros((), DType::F32, dev);
Ok(Self {
llik: zero()?,
scored: zero()?,
})
}
pub(crate) fn add(&mut self, llik_sum: &Tensor, units_sum: &Tensor) -> candle_core::Result<()> {
self.llik = (&self.llik + llik_sum.detach())?;
self.scored = (&self.scored + units_sum.detach())?;
Ok(())
}
pub(crate) fn read(&self) -> candle_core::Result<f32> {
let scored = self.scored.to_scalar::<f32>()?;
Ok(if scored > 0.0 {
self.llik.to_scalar::<f32>()? / scored
} else {
0.0
})
}
}
fn score_hidden_genes(
decoder: &EmbeddedNbTopicDecoder,
log_z: &Tensor,
likelihood: MaskedLikelihood,
mb: &DenseMaskedMinibatch,
full_kd: &Tensor,
) -> candle_core::Result<(Tensor, Tensor)> {
let lib_n1 = (mb.target_nd.sum_keepdim(1)? + 1.0)?;
let target = MaskedDenseTarget {
values: &mb.target_nd,
residual: None,
lib: &lib_n1,
hidden_ids: &mb.hidden_ids,
hidden_weight: mb.hidden_weight.as_ref(),
};
let llik = match likelihood {
MaskedLikelihood::Nb => decoder.impute_dense_nb(log_z, &target, full_kd)?,
MaskedLikelihood::Multinomial => {
decoder.impute_dense_multinomial(log_z, &target, full_kd)?
}
};
let (n, dh) = mb.hidden_ids.dims2()?;
let units = match mb.hidden_weight.as_ref() {
Some(w) => w.sum(1)?,
None => Tensor::full(dh as f32, n, llik.device())?,
};
Ok((llik, units))
}
fn masked_minibatch_loss(
encoder: &IndexedEmbeddingEncoder,
decoder: &EmbeddedNbTopicDecoder,
config: &IndexedTrainConfig,
opts: &MaskedTrainOpts,
mb: &DenseMaskedMinibatch,
mean_1d: &Tensor,
ridge_step: f64,
) -> anyhow::Result<StepLoss> {
let raw_z = masked_encode_dense(
encoder,
opts.latent,
&MaskedDenseInput {
x_nd: &mb.x_nd,
x0_nd: mb.x0_nd.as_ref(),
mean_1d: Some(mean_1d),
visible_nd: &mb.visible_nd,
},
true,
)?;
let log_z = decoder_log_theta(raw_z, opts.latent, config.topic_smoothing)?;
let full_kd = decoder.full_logits_kd()?;
let (llik, units) = if decoder.coarsening().is_identity() {
score_hidden_genes(decoder, &log_z, opts.likelihood, mb, &full_kd)?
} else {
let t = dense_module_targets(decoder.coarsening(), &mb.target_nd, &mb.visible_nd)?;
let module_target = ModuleTarget {
values: &t.values_nm,
visible_counts: &t.visible_counts_nm,
visible_share: &t.visible_share_nm,
residual: None,
lib: &t.lib_n1,
};
match opts.likelihood {
MaskedLikelihood::Nb => {
decoder.score_unseen_modules_nb(&log_z, &module_target, &full_kd)?
}
MaskedLikelihood::Multinomial => {
decoder.score_unseen_modules_multinomial(&log_z, &module_target, &full_kd)?
}
}
};
let llik_sum = llik.sum_all()?;
let units_sum = units.sum_all()?;
let mut loss = llik_sum.neg()?.div(&units_sum.clamp(1.0, f64::INFINITY)?)?;
if ridge_step > 0.0 {
if let Some(r) = encoder.features().lora_ridge()? {
loss = (loss + r.affine(ridge_step, 0.0)?)?;
}
}
if config.feature_embedding_l2 > 0.0 && config.feature_anchor.is_none() {
let rho_l2 = encoder
.features()
.ridge_table()
.sqr()?
.mean_all()?
.affine(f64::from(config.feature_embedding_l2), 0.0)?;
loss = (loss + rho_l2)?;
}
Ok(StepLoss {
loss,
llik_sum,
units_sum,
})
}
pub fn train_masked(
level_data: &[LevelData],
encoder: &IndexedEmbeddingEncoder,
decoders: &[EmbeddedNbTopicDecoder],
config: &IndexedTrainConfig,
mask_fraction: f64,
opts: &MaskedTrainOpts,
) -> anyhow::Result<TrainScores> {
let num_levels = level_data.len();
let total_epochs = config.epochs;
for (level, (&(mixed, _, _), decoder)) in level_data.iter().zip(decoders.iter()).enumerate() {
info!(
"Level {}/{}: {} samples, decoder dim {} over {} genes (masked-imputation ETM)",
level + 1,
num_levels,
mixed.ncols(),
decoder.dim_obs(),
decoder.n_features(),
);
}
info!(
"Masked-imputation training: {num_levels} levels, {total_epochs} epochs, mask={mask_fraction}"
);
let pinned: Vec<String> = {
let suffix = format!(".{}", crate::candle::decoder::masked_etm::BACKGROUND_VAR);
let tbl = config.parameters.data().lock().unwrap();
tbl.keys()
.filter(|name| name.ends_with(&suffix))
.cloned()
.collect()
};
let frozen: Vec<&str> = pinned
.iter()
.map(String::as_str)
.chain(config.feature_anchor.map(|a| a.base_var))
.chain(config.feature_anchor.and_then(|a| a.lora).map(|l| l.v_var))
.collect();
let adam_vars: Vec<Var> =
crate::candle::frozen_features::trainable_vars(config.parameters, &frozen);
let mut adams = vec![AdamW::new(
adam_vars,
candle_nn::ParamsAdamW {
lr: f64::from(config.learning_rate),
weight_decay: f64::from(config.weight_decay),
..Default::default()
},
)?];
if let Some(l) = config.feature_anchor.and_then(|a| a.lora) {
adams.push(l.optimizer(config.parameters, config.learning_rate)?);
}
let prog_bar = labeled_bar("Epochs", total_epochs as u64);
let mut llik_trace = Vec::with_capacity(total_epochs);
let draw = MaskedDraw {
schedule: opts.mask_schedule,
mask_fraction,
};
let mut levels = resident_dense_levels(level_data, config, config.dev)?;
let minibatch_size = match (config.gpu_mem_fraction, levels.first()) {
(Some(frac), Some(level0)) => {
let cap = config.minibatch_size;
crate::candle::device::auto_chunk_size(config.dev, cap, 16.min(cap), frac, |n| {
let mb = level0
.probe_minibatch(n, epoch_seed(opts.seed, 0, 0), &draw)
.map_err(|e| candle_core::Error::Msg(e.to_string()))?;
let fwd = masked_minibatch_loss(
encoder,
&decoders[0],
config,
opts,
&mb,
level0.feature_mean_1d(),
0.0,
)
.map_err(|e| candle_core::Error::Msg(e.to_string()))?;
Ok(fwd.loss)
})
.unwrap_or(cap)
}
_ => config.minibatch_size,
};
for epoch in 0..total_epochs {
if opts.poisson_thin {
let thinned: Vec<(Mat, Option<&Mat>, Mat)> = level_data
.iter()
.enumerate()
.map(|(level, &(mixed, batch, target))| {
let epoch_salt = (epoch as u64) << 32 | level as u64;
let x = poisson_draw(mixed, mix_seed(opts.seed, epoch_salt));
let y = if std::ptr::eq(mixed, target) {
x.clone()
} else {
poisson_draw(target, mix_seed(opts.seed, !epoch_salt))
};
(x, batch, y)
})
.collect();
let refs: Vec<LevelData> = thinned.iter().map(|(x, b, y)| (x, *b, y)).collect();
levels = resident_dense_levels(&refs, config, config.dev)?;
}
let mut acc = EpochAccum::new(config.dev)?;
let mut skipped_steps = 0usize;
for (level, lv) in levels.iter().enumerate() {
let decoder = &decoders[level];
let ep = lv.begin_epoch(epoch_seed(opts.seed, epoch, level), &draw, minibatch_size)?;
let ridge_step = config.feature_anchor.and_then(|a| a.lora).map_or(0.0, |l| {
f64::from(l.ridge) / (levels.len() * ep.n_batches().max(1)) as f64
});
for b in 0..ep.n_batches() {
let mb = ep.batch(b)?;
let fwd = masked_minibatch_loss(
encoder,
decoder,
config,
opts,
&mb,
lv.feature_mean_1d(),
ridge_step,
)?;
acc.add(&fwd.llik_sum, &fwd.units_sum)?;
let grads = fwd.loss.backward()?;
if !clip_and_step_dense_all(&mut adams, grads, f64::from(config.grad_clip))? {
skipped_steps += 1;
}
if config.stop.load(Ordering::Relaxed) {
break;
}
}
}
let per_metric = acc.read()?;
llik_trace.push(per_metric);
prog_bar.set_message(format!("llik={per_metric:.3}"));
prog_bar.inc(1);
if skipped_steps > 0 {
warn!(
"[epoch {epoch}] skipped {skipped_steps} optimizer step(s): \
non-finite gradient norm. Lower --learning-rate or --grad-clip \
if this persists."
);
}
info!("[epoch {epoch}] masked llik/unit={per_metric:.4}");
if config.stop.load(Ordering::SeqCst) {
prog_bar.finish_and_clear();
info!("Stopping early at epoch {epoch}");
return Ok(TrainScores {
kl: vec![0.0; llik_trace.len()],
llik: llik_trace,
});
}
}
prog_bar.finish_and_clear();
info!("done masked-imputation training");
Ok(TrainScores {
kl: vec![0.0; llik_trace.len()],
llik: llik_trace,
})
}
#[cfg(test)]
#[path = "masked_topic_tests.rs"]
mod masked_topic_tests;