use serde::{Deserialize, Serialize};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default, Serialize, Deserialize)]
pub enum Metric {
#[default]
L2,
InnerProduct,
Cosine,
}
impl Metric {
#[inline]
pub const fn requires_normalized_input(self) -> bool {
matches!(self, Metric::Cosine)
}
#[inline]
pub const fn as_str(self) -> &'static str {
match self {
Metric::L2 => "l2",
Metric::InnerProduct => "inner_product",
Metric::Cosine => "cosine",
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn normalization_requirement() {
assert!(Metric::Cosine.requires_normalized_input());
assert!(!Metric::L2.requires_normalized_input());
assert!(!Metric::InnerProduct.requires_normalized_input());
}
#[test]
fn stable_string_tags() {
assert_eq!(Metric::L2.as_str(), "l2");
assert_eq!(Metric::InnerProduct.as_str(), "inner_product");
assert_eq!(Metric::Cosine.as_str(), "cosine");
}
}