use rustc_hash::FxHashSet;
use verbora_core::whitespace::is_whitespace;
pub fn dice_coefficient(s1: &str, s2: &str) -> f64 {
let a = sanitize(s1);
let b = sanitize(s2);
let bigrams_a = bigrams(&a);
let bigrams_b = bigrams(&b);
let (small, large) = if bigrams_a.len() <= bigrams_b.len() {
(&bigrams_a, &bigrams_b)
} else {
(&bigrams_b, &bigrams_a)
};
let shared = small.iter().filter(|g| large.contains(*g)).count();
(2 * shared) as f64 / (bigrams_a.len() + bigrams_b.len()) as f64
}
fn sanitize(s: &str) -> Vec<u16> {
let lowered = s.to_lowercase();
let mut out: Vec<u16> = Vec::with_capacity(lowered.len());
let mut pending_space = false;
for c in lowered.chars() {
if is_whitespace(c) {
pending_space = true;
continue;
}
if pending_space && !out.is_empty() {
out.push(u16::from(b' '));
}
pending_space = false;
let mut buf = [0u16; 2];
out.extend_from_slice(c.encode_utf16(&mut buf));
}
out
}
fn bigrams(units: &[u16]) -> FxHashSet<(u16, u16)> {
let mut set = FxHashSet::default();
if units.len() == 1 {
set.insert((units[0], u16::from(b' ')));
return set;
}
if units.is_empty() {
return set;
}
set.reserve(units.len() - 1);
for w in units.windows(2) {
set.insert((w[0], w[1]));
}
set
}
#[cfg(feature = "parallel")]
pub fn par_dice_coefficient_batch(pairs: &[(&str, &str)]) -> Vec<f64> {
use rayon::prelude::*;
pairs
.par_iter()
.map(|(a, b)| dice_coefficient(a, b))
.collect()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn identical_strings_score_one() {
assert_eq!(dice_coefficient("abc", "abc"), 1.0);
assert_eq!(dice_coefficient("a", "a"), 1.0);
}
#[test]
fn empty_pair_is_nan() {
assert!(dice_coefficient("", "").is_nan());
}
#[test]
fn empty_against_nonempty_is_zero() {
assert_eq!(dice_coefficient("", "abc"), 0.0);
assert_eq!(dice_coefficient("abc", ""), 0.0);
}
#[test]
fn bigrams_are_a_set_so_repeats_collapse() {
assert_eq!(dice_coefficient("aaaa", "aa"), 1.0);
}
#[test]
fn sanitize_folds_case_and_collapses_space() {
assert_eq!(dice_coefficient("Hello World", "hello world"), 1.0);
assert_eq!(dice_coefficient(" padded ", "padded"), 1.0);
}
#[test]
fn partial_overlap_is_between_zero_and_one() {
let d = dice_coefficient("night", "nacht");
assert!(d > 0.0 && d < 1.0, "got {d}");
}
#[test]
fn single_char_is_padded_not_dropped() {
assert_eq!(dice_coefficient("a", "b"), 0.0);
assert_eq!(bigrams(&[u16::from(b'a')]).len(), 1);
}
#[test]
fn astral_characters_use_code_unit_pairs() {
let units: Vec<u16> = "😀".encode_utf16().collect();
assert_eq!(units.len(), 2);
assert_eq!(bigrams(&units).len(), 1);
assert_eq!(dice_coefficient("😀", "😀"), 1.0);
}
}