fn boundary_mask(n_sentences: usize, breaks: &[usize]) -> Vec<bool> {
if n_sentences == 0 {
return Vec::new();
}
let mut mask = vec![false; n_sentences - 1];
for &b in breaks {
if b >= 1 && b <= mask.len() {
mask[b - 1] = true;
}
}
mask
}
pub fn default_window(n_sentences: usize, reference: &[usize]) -> usize {
if n_sentences == 0 {
return 2;
}
let segments = reference.len() + 1;
let mean_len = n_sentences as f64 / segments as f64;
((mean_len / 2.0).round() as usize).max(2)
}
pub fn pk(n_sentences: usize, reference: &[usize], hypothesis: &[usize], k: usize) -> f64 {
let r = boundary_mask(n_sentences, reference);
let h = boundary_mask(n_sentences, hypothesis);
if r.len() < k || k == 0 {
return f64::NAN;
}
let mut disagreements = 0usize;
let mut windows = 0usize;
for start in 0..=(r.len() - k) {
let r_same = !r[start..start + k].iter().any(|b| *b);
let h_same = !h[start..start + k].iter().any(|b| *b);
if r_same != h_same {
disagreements += 1;
}
windows += 1;
}
if windows == 0 {
f64::NAN
} else {
disagreements as f64 / windows as f64
}
}
pub fn window_diff(
n_sentences: usize,
reference: &[usize],
hypothesis: &[usize],
k: usize,
) -> f64 {
let r = boundary_mask(n_sentences, reference);
let h = boundary_mask(n_sentences, hypothesis);
if r.len() < k || k == 0 {
return f64::NAN;
}
let mut disagreements = 0usize;
let mut windows = 0usize;
for start in 0..=(r.len() - k) {
let rc = r[start..start + k].iter().filter(|b| **b).count();
let hc = h[start..start + k].iter().filter(|b| **b).count();
if rc != hc {
disagreements += 1;
}
windows += 1;
}
if windows == 0 {
f64::NAN
} else {
disagreements as f64 / windows as f64
}
}
pub fn boundary_counts(reference: &[usize], hypothesis: &[usize]) -> (usize, usize, usize) {
let tp = hypothesis.iter().filter(|b| reference.contains(b)).count();
let fp = hypothesis.len() - tp;
let fn_ = reference.len() - tp;
(tp, fp, fn_)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn perfect_segmentation_scores_zero() {
let reference = vec![4, 8];
assert_eq!(pk(12, &reference, &reference, 3), 0.0);
assert_eq!(window_diff(12, &reference, &reference, 3), 0.0);
}
#[test]
fn a_near_miss_is_penalised_less_than_a_wild_guess() {
let reference = vec![6];
let near = window_diff(12, &reference, &[7], 3);
let far = window_diff(12, &reference, &[1], 3);
assert!(near < far, "near {near} should beat far {far}");
}
#[test]
fn missing_every_boundary_is_penalised() {
let reference = vec![4, 8];
assert!(window_diff(12, &reference, &[], 3) > 0.0);
assert!(pk(12, &reference, &[], 3) > 0.0);
}
#[test]
fn window_diff_penalises_spurious_boundaries_pk_can_miss() {
let reference = vec![6];
let over = vec![5, 6, 7];
assert!(
window_diff(12, &reference, &over, 3) > 0.0,
"WindowDiff should notice the extra boundaries"
);
}
#[test]
fn boundary_mask_maps_breaks_to_gaps() {
let mask = boundary_mask(6, &[4]);
assert_eq!(mask, vec![false, false, false, true, false]);
}
#[test]
fn boundary_mask_ignores_out_of_range_breaks() {
let mask = boundary_mask(4, &[0, 4, 99]);
assert_eq!(mask, vec![false, false, false]);
}
#[test]
fn default_window_is_half_the_mean_segment() {
assert_eq!(default_window(12, &[4, 8]), 2);
assert_eq!(default_window(30, &[10, 20]), 5);
}
#[test]
fn default_window_never_degenerates_to_zero() {
assert!(default_window(2, &[1]) >= 2);
assert!(default_window(0, &[]) >= 2);
}
#[test]
fn metrics_are_nan_when_the_window_does_not_fit() {
assert!(pk(3, &[1], &[1], 10).is_nan());
assert!(window_diff(3, &[1], &[1], 10).is_nan());
}
#[test]
fn boundary_counts_are_exact_match() {
let (tp, fp, fn_) = boundary_counts(&[4, 8], &[4, 9]);
assert_eq!((tp, fp, fn_), (1, 1, 1));
}
#[test]
fn boundary_counts_with_no_hypothesis() {
assert_eq!(boundary_counts(&[4, 8], &[]), (0, 0, 2));
}
#[test]
fn splitting_everywhere_is_worse_than_splitting_nowhere_here() {
let reference = vec![6];
let none = window_diff(12, &reference, &[], 3);
let all: Vec<usize> = (1..12).collect();
assert!(
window_diff(12, &reference, &all, 3) > none,
"over-splitting should score worse than not splitting"
);
}
}