pub mod masked_topic;
pub mod topic;
use candle_core::Tensor;
use candle_nn::{AdamW, Optimizer};
use log::info;
use std::time::Duration;
pub struct TrainScores {
pub llik: Vec<f32>,
pub kl: Vec<f32>,
}
pub fn smooth_topics(log_z_nk: Tensor, alpha: f64) -> candle_core::Result<Tensor> {
if alpha > 0.0 {
let kk = log_z_nk.dim(1)? as f64;
((log_z_nk.exp()? * (1.0 - alpha))? + alpha / kk)?.log()
} else {
Ok(log_z_nk)
}
}
pub fn stick_breaking_log_simplex(logits_nk: &Tensor) -> candle_core::Result<Tensor> {
let k = logits_nk.dim(1)?;
if k == 1 {
return logits_nk.zeros_like();
}
let eta = logits_nk.narrow(1, 0, k - 1)?; let log_1mv = crate::candle::loss::log_sigmoid(&eta.neg()?)?; let incl = log_1mv.cumsum(1)?; let head = (&eta + &incl)?; let tail = incl.narrow(1, k - 2, 1)?; Tensor::cat(&[&head, &tail], 1) }
#[derive(Default)]
pub struct PhaseTimers {
pub precompute: Duration,
pub encoder_fwd: Duration,
pub decoder_fwd: Duration,
pub backward: Duration,
pub optimize: Duration,
}
impl PhaseTimers {
pub fn log_summary(&self) {
let total =
self.precompute + self.encoder_fwd + self.decoder_fwd + self.backward + self.optimize;
let total_s = total.as_secs_f64().max(1e-9);
let pct = |d: Duration| 100.0 * d.as_secs_f64() / total_s;
info!(
"phase timing — precompute {:.1}s ({:.0}%), encoder_fwd {:.1}s ({:.0}%), \
decoder_fwd {:.1}s ({:.0}%), backward {:.1}s ({:.0}%), opt_step {:.1}s ({:.0}%)",
self.precompute.as_secs_f64(),
pct(self.precompute),
self.encoder_fwd.as_secs_f64(),
pct(self.encoder_fwd),
self.decoder_fwd.as_secs_f64(),
pct(self.decoder_fwd),
self.backward.as_secs_f64(),
pct(self.backward),
self.optimize.as_secs_f64(),
pct(self.optimize),
);
}
}
fn apply_global_l2_clip(
grads: &mut candle_core::backprop::GradStore,
max_norm: f64,
) -> anyhow::Result<bool> {
let ids: Vec<_> = grads.get_ids().copied().collect();
let mut sumsq: Option<Tensor> = None;
for id in &ids {
if let Some(g) = grads.get_id(*id) {
let s = g.sqr()?.sum_all()?;
sumsq = Some(match sumsq {
None => s,
Some(prev) => (prev + s)?,
});
}
}
let Some(sumsq) = sumsq else {
return Ok(true);
};
if !sumsq.to_scalar::<f32>()?.is_finite() {
return Ok(false);
}
let inv_norm = sumsq.sqrt()?.affine(1.0, 1e-6)?.powf(-1.0)?;
let scale = inv_norm.affine(max_norm, 0.0)?.clamp(0.0_f64, 1.0_f64)?;
for id in &ids {
if let Some(g) = grads.get_id(*id) {
let scaled = g.broadcast_mul(&scale)?;
grads.insert_id(*id, scaled);
}
}
Ok(true)
}
pub fn clip_grads_and_step<O: Optimizer>(
opt: &mut O,
loss: &Tensor,
max_norm: f64,
) -> anyhow::Result<()> {
if max_norm <= 0.0 {
opt.backward_step(loss)?;
return Ok(());
}
let mut grads = loss.backward()?;
if apply_global_l2_clip(&mut grads, max_norm)? {
opt.step(&grads)?;
}
Ok(())
}
pub fn clip_and_step_dense(
adam: &mut AdamW,
grads: candle_core::backprop::GradStore,
max_norm: f64,
) -> anyhow::Result<bool> {
clip_and_step_dense_all(std::slice::from_mut(adam), grads, max_norm)
}
pub fn clip_and_step_dense_all(
adams: &mut [AdamW],
mut grads: candle_core::backprop::GradStore,
max_norm: f64,
) -> anyhow::Result<bool> {
if max_norm > 0.0 && !apply_global_l2_clip(&mut grads, max_norm)? {
return Ok(false);
}
for adam in adams.iter_mut() {
adam.step(&grads)?;
}
Ok(true)
}
pub type LevelLossHook<'a> = dyn Fn(Tensor, usize) -> anyhow::Result<Tensor> + 'a;
#[cfg(test)]
mod tests {
use super::*;
use candle_core::Device;
#[test]
fn stick_breaking_matches_reference_and_normalizes() {
let dev = Device::Cpu;
let ln2 = std::f32::consts::LN_2;
let logits =
Tensor::from_vec(vec![0.0f32, ln2, 999.0, -1.5, 2.0, -999.0], (2, 3), &dev).unwrap();
let log_theta = stick_breaking_log_simplex(&logits).unwrap();
assert_eq!(log_theta.dims(), &[2, 3]);
let theta = log_theta.exp().unwrap().to_vec2::<f32>().unwrap();
let expect0 = [0.5f32, 1.0 / 3.0, 1.0 / 6.0];
for (got, want) in theta[0].iter().zip(expect0.iter()) {
assert!((got - want).abs() < 1e-5, "θ₀ {got} vs {want}");
}
for row in &theta {
let sum: f32 = row.iter().sum();
assert!((sum - 1.0).abs() < 1e-5, "row sum {sum} ≠ 1");
assert!(row.iter().all(|&p| p > 0.0), "non-positive θ entry");
}
}
#[test]
fn stick_breaking_k1_is_degenerate() {
let dev = Device::Cpu;
let logits = Tensor::from_vec(vec![3.7f32, -2.0], (2, 1), &dev).unwrap();
let log_theta = stick_breaking_log_simplex(&logits).unwrap();
for v in log_theta.flatten_all().unwrap().to_vec1::<f32>().unwrap() {
assert_eq!(v, 0.0, "K=1 log θ must be 0");
}
}
}