use crate::error::InfoError;
use std::collections::HashMap;
use std::hash::Hash;
pub fn js_divergence<T>(p: &[T], q: &[T]) -> Result<f64, InfoError>
where
T: Eq + Hash,
{
if p.is_empty() || q.is_empty() {
return Err(InfoError::EmptyInput);
}
let mut p_counts: HashMap<&T, usize> = HashMap::new();
for x in p {
*p_counts.entry(x).or_insert(0) += 1;
}
let mut q_counts: HashMap<&T, usize> = HashMap::new();
for x in q {
*q_counts.entry(x).or_insert(0) += 1;
}
let p_total = p.len() as f64;
let q_total = q.len() as f64;
let mut kl_p_m = 0.0f64;
for (x, &pc) in &p_counts {
let p_x = pc as f64 / p_total;
let q_x = q_counts.get(x).map_or(0.0, |&qc| qc as f64 / q_total);
let m_x = 0.5 * (p_x + q_x); kl_p_m += p_x * (p_x / m_x).log2();
}
let mut kl_q_m = 0.0f64;
for (x, &qc) in &q_counts {
let q_x = qc as f64 / q_total;
let p_x = p_counts.get(x).map_or(0.0, |&pc| pc as f64 / p_total);
let m_x = 0.5 * (p_x + q_x); kl_q_m += q_x * (q_x / m_x).log2();
}
Ok(0.5 * (kl_p_m + kl_q_m))
}
pub fn js_divergence_unchecked<T>(p: &[T], q: &[T]) -> f64
where
T: Eq + Hash,
{
js_divergence(p, q).expect("js_divergence failed")
}