use super::{DensityError, gaussian_kde};
#[test]
fn density_matches_hand_value() -> Result<(), DensityError> {
let kde = gaussian_kde(&[2.0, 4.0, 6.0])?;
let f4 = kde.density_at(4.0);
assert!(
(f4 - 0.159_078_110_463_200_72).abs() < 1e-12,
"f(4) was {f4}"
);
Ok(())
}
#[test]
fn bandwidth_follows_scotts_rule() -> Result<(), DensityError> {
let kde = gaussian_kde(&[2.0, 4.0, 6.0])?;
let h = kde.bandwidth();
assert!((h - 1.605_483_123_520_461_4).abs() < 1e-12, "h was {h}");
let var = kde.variance();
assert!(
(var - h * h).abs() < 1e-12,
"variance/bandwidth mismatch: {var}"
);
Ok(())
}
#[test]
fn empty_sample_is_rejected() {
assert_eq!(gaussian_kde(&[]), Err(DensityError::EmptyInput));
}
#[test]
fn single_point_is_rejected() {
assert_eq!(gaussian_kde(&[3.0]), Err(DensityError::SinglePoint));
}
#[test]
fn zero_variance_is_rejected() {
assert_eq!(
gaussian_kde(&[5.0, 5.0, 5.0]),
Err(DensityError::ZeroVariance)
);
}
#[test]
fn non_finite_is_rejected() {
assert_eq!(
gaussian_kde(&[1.0, f64::NAN, 3.0]),
Err(DensityError::NonFinite)
);
assert_eq!(
gaussian_kde(&[1.0, f64::INFINITY, 3.0]),
Err(DensityError::NonFinite)
);
}
#[test]
fn density_slice_matches_pointwise() -> Result<(), DensityError> {
let kde = gaussian_kde(&[1.0, 2.0, 2.5, 3.0, 7.0, 8.0])?;
let xs = [2.0, 5.0, 7.5];
let batch = kde.density(&xs);
assert_eq!(batch.len(), xs.len(), "batch length");
for (i, &x) in xs.iter().enumerate() {
let one = kde.density_at(x);
let many = batch.get(i).copied().unwrap_or(f64::NAN);
assert!((one - many).abs() < 1e-15, "index {i}: {one} vs {many}");
}
Ok(())
}
#[test]
fn density_integrates_to_one() -> Result<(), DensityError> {
let kde = gaussian_kde(&[1.0, 2.0, 2.5, 3.0, 7.0, 8.0])?;
let lo = -20.0_f64;
let hi = 30.0_f64;
let steps = 50_000_usize;
let step = (hi - lo) / f64::from(u32::try_from(steps).unwrap_or(0));
let mut area = 0.0_f64;
let mut prev = kde.density_at(lo);
for i in 1..=steps {
let x = step.mul_add(f64::from(u32::try_from(i).unwrap_or(0)), lo);
let cur = kde.density_at(x);
area = (0.5 * (prev + cur)).mul_add(step, area);
prev = cur;
}
assert!((area - 1.0).abs() < 1e-4, "integral was {area}");
Ok(())
}