use std::cmp::Ordering;
use crate::Axis;
pub trait DistanceMetricScalar<A: Copy> {
type Output: Axis<Coord = Self::Output>;
const ORDERING: Ordering;
fn widen_coord(a: A) -> Self::Output;
#[inline(always)]
fn widen_axis(axis: &[A], out: &mut [Self::Output]) {
assert!(out.len() >= axis.len());
for (dst, &src) in out.iter_mut().zip(axis.iter()) {
*dst = Self::widen_coord(src);
}
}
fn dist1(a: Self::Output, b: Self::Output) -> Self::Output;
#[inline(always)]
fn combine_component(acc: &mut Self::Output, component: Self::Output) {
*acc += component;
}
#[inline(always)]
fn dist1_raw(a: A, b: A) -> Self::Output {
Self::dist1(Self::widen_coord(a), Self::widen_coord(b))
}
#[inline(always)]
fn dist<const K: usize>(a: &[Self::Output; K], b: &[Self::Output; K]) -> Self::Output {
let mut acc = Self::Output::zero();
for dim in 0..K {
Self::combine_component(&mut acc, Self::dist1(a[dim], b[dim]));
}
acc
}
#[inline(always)]
fn dist_raw<const K: usize>(a: &[A; K], b: &[A; K]) -> Self::Output {
let mut acc = Self::Output::zero();
for dim in 0..K {
Self::combine_component(&mut acc, Self::dist1_raw(a[dim], b[dim]));
}
acc
}
#[inline(always)]
fn rect_dist_from_off<const K: usize>(off: &[Self::Output; K]) -> Self::Output {
let mut acc = Self::Output::zero();
for off_val in off.iter().copied() {
Self::combine_component(&mut acc, Self::dist1(off_val, Self::Output::zero()));
}
acc
}
#[inline(always)]
fn rect_dist_after_update<const K: usize>(
rd: Self::Output,
off: &[Self::Output; K],
dim: usize,
new_off: Self::Output,
) -> Self::Output {
let new_dist1 = Self::dist1(new_off, Self::Output::zero());
let old_dist1 = Self::dist1(off[dim], Self::Output::zero());
Self::Output::saturating_add(rd - old_dist1, new_dist1)
}
#[inline(always)]
fn better(a: Self::Output, b: Self::Output) -> bool {
match Self::ORDERING {
Ordering::Less => a < b,
Ordering::Greater => a > b,
Ordering::Equal => false,
}
}
#[inline(always)]
fn cmp(a: Self::Output, b: Self::Output) -> Ordering {
a.partial_cmp(&b).unwrap_or(Ordering::Equal)
}
}
#[cfg(test)]
mod tests {
use super::DistanceMetricScalar;
use std::cmp::Ordering;
struct DummyLessMetric;
struct DummyGreaterMetric;
impl DistanceMetricScalar<i16> for DummyLessMetric {
type Output = f64;
const ORDERING: Ordering = Ordering::Less;
fn widen_coord(a: i16) -> Self::Output {
a as f64
}
fn dist1(a: Self::Output, b: Self::Output) -> Self::Output {
(a - b).abs()
}
}
impl DistanceMetricScalar<i16> for DummyGreaterMetric {
type Output = f64;
const ORDERING: Ordering = Ordering::Greater;
fn widen_coord(a: i16) -> Self::Output {
a as f64
}
fn dist1(a: Self::Output, b: Self::Output) -> Self::Output {
a - b
}
}
#[test]
fn default_widen_axis_bulk_widens() {
let axis = [1i16, -2, 7];
let mut out = [0.0f64; 3];
DummyLessMetric::widen_axis(&axis, &mut out);
assert_eq!(out, [1.0, -2.0, 7.0]);
}
#[test]
fn default_better_respects_ordering() {
assert!(DummyLessMetric::better(2.0, 5.0));
assert!(!DummyLessMetric::better(5.0, 2.0));
assert!(DummyGreaterMetric::better(5.0, 2.0));
assert!(!DummyGreaterMetric::better(2.0, 5.0));
}
#[test]
fn default_cmp_uses_partial_cmp() {
assert_eq!(DummyLessMetric::cmp(2.0, 5.0), Ordering::Less);
assert_eq!(DummyLessMetric::cmp(5.0, 2.0), Ordering::Greater);
assert_eq!(DummyLessMetric::cmp(3.0, 3.0), Ordering::Equal);
}
}