use super::topic::LevelData;
use super::{clip_grads_and_step, smooth_topics};
use crate::candle::convert::to_1d;
use crate::candle::data::loader::{InMemoryArgs, InMemoryData};
use crate::candle::traits::model::EncoderModuleT;
use candle_core::{Device, Result, Tensor, Var};
use candle_nn::{AdamW, Optimizer, ParamsAdamW, VarMap};
use log::{debug, info};
use rand::rngs::SmallRng;
use rand::{RngExt, SeedableRng};
use std::sync::atomic::{AtomicBool, Ordering};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum PairMetric {
Hellinger,
Euclidean,
}
#[derive(Clone, Debug, Default)]
pub struct LevelPairs {
pub pairs: Vec<(u32, u32, f32)>,
pub margin: f32,
}
pub fn check_pairs(level: &LevelPairs, n_rows: usize) -> anyhow::Result<()> {
let pairs = &level.pairs;
if let Some(&(a, b, _)) = pairs
.iter()
.find(|&&(a, b, _)| a as usize >= n_rows || b as usize >= n_rows)
{
anyhow::bail!("pair ({a}, {b}) is out of range for a level of {n_rows} rows");
}
if let Some(&(_, _, w)) = pairs.iter().find(|&&(_, _, w)| !(w.is_finite() && w > 0.0)) {
anyhow::bail!("pair weight {w} is not finite and positive");
}
anyhow::ensure!(
level.margin.is_finite() && level.margin >= 0.0,
"margin {} is not finite and non-negative",
level.margin
);
Ok(())
}
pub fn latent_distance(a: &Tensor, b: &Tensor, metric: PairMetric) -> Result<Tensor> {
let diff = match metric {
PairMetric::Hellinger => (a.exp()?.sqrt()? - b.exp()?.sqrt()?)?,
PairMetric::Euclidean => (a - b)?,
};
let d = (diff.sqr()?.sum(1)? + 1e-12)?.sqrt()?;
match metric {
PairMetric::Hellinger => d.affine(std::f64::consts::FRAC_1_SQRT_2, 0.0),
PairMetric::Euclidean => Ok(d),
}
}
const QUANTILE_PAIRS: usize = 200_000;
pub fn quantile_distance(z: &Tensor, metric: PairMetric, q: f32) -> anyhow::Result<f32> {
let n = z.dim(0)?;
anyhow::ensure!(
n >= 2,
"a quantile of pair distances needs two rows, not {n}"
);
anyhow::ensure!((0.0..=1.0).contains(&q), "quantile {q} is outside [0, 1]");
let (a, b): (Vec<u32>, Vec<u32>) = if n * (n - 1) / 2 <= QUANTILE_PAIRS {
(0..n as u32)
.flat_map(|i| (i + 1..n as u32).map(move |j| (i, j)))
.unzip()
} else {
let mut rng = SmallRng::seed_from_u64(0x5EED);
(0..QUANTILE_PAIRS)
.map(|_| {
let i = rng.random_range(0..n);
let j = (i + rng.random_range(1..n)) % n;
(i as u32, j as u32)
})
.unzip()
};
let rows = |idx: Vec<u32>| -> Result<Tensor> {
let len = idx.len();
z.index_select(&Tensor::from_vec(idx, len, z.device())?, 0)
};
let mut d = latent_distance(&rows(a)?, &rows(b)?, metric)?.to_vec1::<f32>()?;
let at = (q * (d.len() - 1) as f32).floor() as usize;
let (_, &mut v, _) = d.select_nth_unstable_by(at, f32::total_cmp);
Ok(v)
}
pub fn pair_hinge(d: &Tensor, w: &Tensor, margin: f32) -> Result<Tensor> {
if margin <= 0.0 {
return d.zeros_like()?.sum_all();
}
let gap = d.affine(-1.0 / f64::from(margin), 1.0)?.relu()?;
let num = (gap.sqr()? * w)?.sum_all()?;
let den = w.sum_all()?;
num / den
}
pub fn pair_batch(n: usize, batch: usize, step: usize) -> Vec<usize> {
if n == 0 || batch == 0 {
return Vec::new();
}
let take = batch.min(n);
let start = (step * take) % n;
(0..take).map(|i| (start + i) % n).collect()
}
pub struct ReviseConfig<'a> {
pub dev: &'a Device,
pub metric: PairMetric,
pub topic_smoothing: f64,
pub learning_rate: f32,
pub max_epochs: usize,
pub batch: usize,
pub grad_clip: f32,
pub stop: &'a AtomicBool,
pub verbose: bool,
}
#[derive(Clone, Debug, Default, PartialEq)]
pub struct ReviseTrace {
pub hinge: Vec<Vec<f32>>,
pub satisfied: Vec<Vec<f32>>,
pub steps: usize,
}
pub fn encoder_trainable_vars(varmap: &VarMap, prefix: &str) -> Vec<Var> {
let is_running_stat = |name: &str| {
["running_mean", "running_var"]
.iter()
.any(|stat| name == *stat || name.ends_with(&format!(".{stat}")))
};
let data = varmap.data().lock().expect("VarMap lock poisoned");
let mut vars: Vec<(&String, &Var)> = data
.iter()
.filter(|(name, _)| name.starts_with(prefix) && !is_running_stat(name))
.collect();
vars.sort_by(|a, b| a.0.cmp(b.0));
vars.into_iter().map(|(_, var)| var.clone()).collect()
}
pub fn revise_encoder<Enc: EncoderModuleT>(
level_data: &[LevelData],
encoder: &Enc,
varmap: &VarMap,
encoder_prefix: &str,
per_level: &[LevelPairs],
config: &ReviseConfig,
) -> anyhow::Result<ReviseTrace> {
anyhow::ensure!(config.batch > 0, "a revise needs a pair batch above zero");
let vars = encoder_trainable_vars(varmap, encoder_prefix);
anyhow::ensure!(
!vars.is_empty(),
"no trainable variables under prefix {encoder_prefix:?}"
);
anyhow::ensure!(
per_level.len() <= level_data.len(),
"pairs for {} levels, data for {}",
per_level.len(),
level_data.len()
);
for (level, (lp, &(input, _, _))) in per_level.iter().zip(level_data).enumerate() {
check_pairs(lp, input.nrows()).map_err(|e| anyhow::anyhow!("level {level}: {e}"))?;
}
let none = LevelPairs::default();
let levels: Vec<(&LevelPairs, Option<InMemoryData>)> = level_data
.iter()
.enumerate()
.map(|(level, &(input, null, _))| {
let lp = per_level.get(level).unwrap_or(&none);
let data = if lp.pairs.is_empty() || lp.margin <= 0.0 {
None
} else {
let args = InMemoryArgs {
input,
input_null: null,
output: None,
output_null: None,
};
Some(InMemoryData::from_device(args, config.dev)?)
};
Ok((lp, data))
})
.collect::<anyhow::Result<_>>()?;
let measure = |trace: &mut ReviseTrace| -> anyhow::Result<bool> {
let mut hinge = Vec::with_capacity(levels.len());
let mut satisfied = Vec::with_capacity(levels.len());
for (lp, data) in &levels {
let (h, s) = match data {
Some(data) => level_hinge(encoder, data, lp, config)?,
None => (0.0, 1.0),
};
hinge.push(h);
satisfied.push(s);
}
let done = hinge.iter().all(|&h| h == 0.0);
trace.hinge.push(hinge);
trace.satisfied.push(satisfied);
Ok(done)
};
let mut trace = ReviseTrace::default();
if measure(&mut trace)? {
return Ok(trace);
}
let steps_per_epoch = levels
.iter()
.filter(|(_, data)| data.is_some())
.map(|(lp, _)| lp.pairs.len().div_ceil(config.batch))
.max()
.unwrap_or(0);
let params = ParamsAdamW {
lr: f64::from(config.learning_rate),
weight_decay: 0.0,
..Default::default()
};
let mut adam = AdamW::new(vars, params)?;
for epoch in 0..config.max_epochs {
for _ in 0..steps_per_epoch {
let mut loss: Option<Tensor> = None;
for (lp, data) in &levels {
let Some(data) = data else { continue };
let take = pair_batch(lp.pairs.len(), config.batch, trace.steps);
let (d, w) = batch_distance(encoder, data, lp, &take, config)?;
let h = pair_hinge(&d, &w, lp.margin)?;
loss = Some(match loss {
Some(l) => (l + h)?,
None => h,
});
}
if let Some(loss) = loss {
clip_grads_and_step(&mut adam, &loss, f64::from(config.grad_clip))?;
}
trace.steps += 1;
if config.stop.load(Ordering::Relaxed) {
break;
}
}
let done = measure(&mut trace)?;
let msg = format!(
"[revise epoch {epoch}] hinge={:?} satisfied={:?}",
trace.hinge.last().unwrap(),
trace.satisfied.last().unwrap()
);
if config.verbose {
info!("{msg}");
} else {
debug!("{msg}");
}
if done || config.stop.load(Ordering::SeqCst) {
break;
}
}
Ok(trace)
}
fn batch_distance<Enc: EncoderModuleT>(
encoder: &Enc,
data: &InMemoryData,
lp: &LevelPairs,
take: &[usize],
config: &ReviseConfig,
) -> anyhow::Result<(Tensor, Tensor)> {
let n = take.len();
let both: Vec<u32> = take
.iter()
.map(|&i| lp.pairs[i].0)
.chain(take.iter().map(|&i| lp.pairs[i].1))
.collect();
let w: Vec<f32> = take.iter().map(|&i| lp.pairs[i].2).collect();
let (x, null) = data.device_rows(&both)?;
let (z, _) = encoder.forward_t(&x, null.as_ref(), false)?;
let z = smooth_topics(z, config.topic_smoothing)?;
let d = latent_distance(&z.narrow(0, 0, n)?, &z.narrow(0, n, n)?, config.metric)?;
let w = to_1d(&w, d.device())?;
Ok((d, w))
}
fn level_hinge<Enc: EncoderModuleT>(
encoder: &Enc,
data: &InMemoryData,
lp: &LevelPairs,
config: &ReviseConfig,
) -> anyhow::Result<(f32, f32)> {
let (mut num, mut den, mut met) = (0f64, 0f64, 0usize);
let idx: Vec<usize> = (0..lp.pairs.len()).collect();
for take in idx.chunks(config.batch) {
let (d, _) = batch_distance(encoder, data, lp, take, config)?;
let d: Vec<f32> = d.detach().to_vec1()?;
for (&i, &d) in take.iter().zip(&d) {
let w = f64::from(lp.pairs[i].2);
let gap = (1.0 - f64::from(d) / f64::from(lp.margin)).max(0.0);
num += w * gap * gap;
den += w;
met += usize::from(d >= lp.margin);
}
}
Ok(((num / den) as f32, met as f32 / lp.pairs.len() as f32))
}
#[cfg(test)]
#[path = "pairs_tests.rs"]
mod pairs_tests;