use crate::tensor::{Result, Tensor, TensorError};
pub(crate) fn divergence_triple(
pre: &[Tensor],
post: &[Tensor],
) -> Result<(f64, Option<f64>, Option<f64>)> {
if pre.is_empty() {
return Ok((0.0, None, None));
}
if pre.len() != post.len() {
return Err(TensorError::new(&format!(
"divergence_triple: pre.len() ({}) must equal post.len() ({})",
pre.len(),
post.len(),
)));
}
let pre_norm_tensors = Tensor::foreach_norm(pre, 2.0)?;
let mut pre_sq = 0.0f64;
for n in &pre_norm_tensors {
let v: f64 = n.item()?;
pre_sq += v * v;
}
let pre_norm = pre_sq.sqrt();
Tensor::foreach_add_list_(pre, post, -1.0)?;
let diff_norms = Tensor::foreach_norm(pre, 2.0)?;
let post_norms = Tensor::foreach_norm(post, 2.0)?;
let mut diff_sq = 0.0f64;
for n in &diff_norms {
let v: f64 = n.item()?;
diff_sq += v * v;
}
let mut post_sq = 0.0f64;
for n in &post_norms {
let v: f64 = n.item()?;
post_sq += v * v;
}
let post_norm = post_sq.sqrt();
let divergence = if post_norm > 1e-10 {
diff_sq.sqrt() / post_norm
} else {
0.0
};
Ok((divergence, Some(post_norm), Some(pre_norm)))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::tensor::{test_device, Tensor};
fn t(data: &[f32]) -> Tensor {
Tensor::from_f32(data, &[data.len() as i64], test_device()).unwrap()
}
#[test]
fn empty_is_zero_none_none() {
let (d, post, pre) = divergence_triple(&[], &[]).unwrap();
assert_eq!(d, 0.0);
assert!(post.is_none() && pre.is_none());
}
#[test]
fn length_mismatch_errors() {
let pre = [t(&[1.0])];
let post = [t(&[1.0]), t(&[2.0])];
assert!(divergence_triple(&pre, &post).is_err());
}
#[test]
fn matches_hand_computed_triple() {
let pre = [t(&[3.0, 4.0])]; let post = [t(&[6.0, 8.0])]; let (d, post_n, pre_n) = divergence_triple(&pre, &post).unwrap();
assert!((pre_n.unwrap() - 5.0).abs() < 1e-5);
assert!((post_n.unwrap() - 10.0).abs() < 1e-5);
assert!((d - 0.5).abs() < 1e-5, "divergence = ||pre-post||/||post|| = 5/10");
}
#[test]
fn zero_post_norm_is_guarded() {
let pre = [t(&[1.0, 1.0])];
let post = [t(&[0.0, 0.0])];
let (d, post_n, _) = divergence_triple(&pre, &post).unwrap();
assert_eq!(d, 0.0, "post_norm ~ 0 must not divide");
assert!(post_n.unwrap() < 1e-9);
}
}