use crate::text::Rect;
pub fn levenshtein<T: PartialEq>(a: &[T], b: &[T]) -> usize {
if a.is_empty() {
return b.len();
}
if b.is_empty() {
return a.len();
}
let mut prev: Vec<usize> = (0..=b.len()).collect();
let mut curr = vec![0usize; b.len() + 1];
for (i, ta) in a.iter().enumerate() {
curr[0] = i + 1;
for (j, tb) in b.iter().enumerate() {
let cost = usize::from(ta != tb);
curr[j + 1] = (prev[j] + cost).min(prev[j + 1] + 1).min(curr[j] + 1);
}
core::mem::swap(&mut prev, &mut curr);
}
prev[b.len()]
}
pub fn cer(reference: &str, hypothesis: &str) -> f64 {
let r: Vec<char> = reference.chars().collect();
let h: Vec<char> = hypothesis.chars().collect();
if r.is_empty() {
return if h.is_empty() { 0.0 } else { 1.0 };
}
levenshtein(&r, &h) as f64 / r.len() as f64
}
pub fn wer(reference: &str, hypothesis: &str) -> f64 {
let r: Vec<&str> = reference.split_whitespace().collect();
let h: Vec<&str> = hypothesis.split_whitespace().collect();
if r.is_empty() {
return if h.is_empty() { 0.0 } else { 1.0 };
}
levenshtein(&r, &h) as f64 / r.len() as f64
}
pub fn iou(a: &Rect, b: &Rect) -> f64 {
let ax1 = a.x + a.width;
let ay1 = a.y + a.height;
let bx1 = b.x + b.width;
let by1 = b.y + b.height;
let ix = ax1.min(bx1).saturating_sub(a.x.max(b.x)) as f64;
let iy = ay1.min(by1).saturating_sub(a.y.max(b.y)) as f64;
let inter = ix * iy;
if inter <= 0.0 {
return 0.0;
}
let union = (a.width as f64 * a.height as f64) + (b.width as f64 * b.height as f64) - inter;
inter / union
}
pub fn mean_line_iou(truth: &[Rect], detected: &[Rect]) -> f64 {
if truth.is_empty() {
return 1.0;
}
let mut pairs: Vec<(f64, usize, usize)> = Vec::new();
for (ti, t) in truth.iter().enumerate() {
for (di, d) in detected.iter().enumerate() {
let v = iou(t, d);
if v > 0.0 {
pairs.push((v, ti, di));
}
}
}
pairs.sort_by(|a, b| b.0.total_cmp(&a.0).then(a.1.cmp(&b.1)).then(a.2.cmp(&b.2)));
let mut truth_used = vec![false; truth.len()];
let mut det_used = vec![false; detected.len()];
let mut sum = 0.0;
for (v, ti, di) in pairs {
if !truth_used[ti] && !det_used[di] {
truth_used[ti] = true;
det_used[di] = true;
sum += v;
}
}
sum / truth.len() as f64
}
#[cfg(test)]
mod tests {
use super::*;
fn rect(x: u32, y: u32, width: u32, height: u32) -> Rect {
Rect {
x,
y,
width,
height,
}
}
#[test]
fn levenshtein_basics() {
assert_eq!(levenshtein::<char>(&[], &[]), 0);
assert_eq!(levenshtein(&['a'], &[]), 1);
let kitten: Vec<char> = "kitten".chars().collect();
let sitting: Vec<char> = "sitting".chars().collect();
assert_eq!(levenshtein(&kitten, &sitting), 3);
}
#[test]
fn cer_counts_unicode_chars_not_bytes() {
assert_eq!(cer("тест", "техт"), 0.25);
assert_eq!(cer("тест", "тест"), 0.0);
}
#[test]
fn cer_empty_reference() {
assert_eq!(cer("", ""), 0.0);
assert_eq!(cer("", "мусор"), 1.0);
}
#[test]
fn wer_treats_newlines_as_separators() {
assert_eq!(wer("один два\nтри", "один два три"), 0.0);
assert!((wer("один два три", "один дваа три") - 1.0 / 3.0).abs() < 1e-12);
}
#[test]
fn iou_disjoint_identical_partial() {
let a = rect(0, 0, 10, 10);
assert_eq!(iou(&a, &rect(20, 20, 10, 10)), 0.0);
assert_eq!(iou(&a, &a), 1.0);
let b = rect(5, 0, 10, 10);
assert!((iou(&a, &b) - 50.0 / 150.0).abs() < 1e-12);
}
#[test]
fn mean_line_iou_matches_greedily_once() {
let truth = [rect(0, 0, 10, 10), rect(0, 20, 10, 10)];
let detected = [rect(0, 0, 10, 12)];
let m = mean_line_iou(&truth, &detected);
let best = iou(&truth[0], &detected[0]);
assert!((m - best / 2.0).abs() < 1e-12);
}
#[test]
fn mean_line_iou_empty_cases() {
assert_eq!(mean_line_iou(&[], &[rect(0, 0, 1, 1)]), 1.0);
assert_eq!(mean_line_iou(&[rect(0, 0, 1, 1)], &[]), 0.0);
}
}