From 8d9ab0e2995246cbe57a36cab9b8fa3a8fca5830 Mon Sep 17 00:00:00 2001 From: Brummel Date: Mon, 18 May 2026 16:54:25 +0200 Subject: [PATCH] feat: RRF and cross-segment max-score dedupe --- src/fusion.rs | 55 ++++++++++++++++++++++++++++++++++++++++++- tests/fusion_tests.rs | 31 ++++++++++++++++++++++++ 2 files changed, 85 insertions(+), 1 deletion(-) create mode 100644 tests/fusion_tests.rs diff --git a/src/fusion.rs b/src/fusion.rs index 4b0fd9d..9cb18ab 100644 --- a/src/fusion.rs +++ b/src/fusion.rs @@ -1 +1,54 @@ -// implemented in a later task +use crate::model::Candidate; +use std::collections::HashMap; + +/// Reciprocal Rank Fusion of two ranked lists of entry indices. +/// Returns deduped entry indices ordered by fused score desc. +pub fn rrf_merge(a: &[usize], b: &[usize], k: usize) -> Vec { + let mut score: HashMap = HashMap::new(); + for (rank, &id) in a.iter().enumerate() { + *score.entry(id).or_default() += 1.0 / (k as f64 + rank as f64 + 1.0); + } + for (rank, &id) in b.iter().enumerate() { + *score.entry(id).or_default() += 1.0 / (k as f64 + rank as f64 + 1.0); + } + let mut v: Vec<(usize, f64)> = score.into_iter().collect(); + v.sort_by(|x, y| y.1.partial_cmp(&x.1).unwrap()); + v.into_iter().map(|(id, _)| id).collect() +} + +/// Group candidates by normalized ICD code across all segments. +/// Score = max over occurrences; retain all source segments and +/// the matched phrase of the best-scoring occurrence. +pub fn cross_segment_dedupe(cands: Vec) -> Vec<(DedupCandidate, f32)> { + let mut by_code: HashMap = HashMap::new(); + for c in cands { + let s = c.rerank.or(c.semantic).or(c.lexical).unwrap_or(0.0); + let e = by_code.entry(c.icd_code.clone()).or_insert_with(|| DedupCandidate { + icd_code: c.icd_code.clone(), + best_phrase: c.alpha_text.clone(), + best_score: f32::MIN, + source_segments: Vec::new(), + }); + if !e.source_segments.contains(&c.segment_idx) { + e.source_segments.push(c.segment_idx); + } + if s > e.best_score { + e.best_score = s; + e.best_phrase = c.alpha_text.clone(); + } + } + let mut out: Vec<(DedupCandidate, f32)> = by_code + .into_values() + .map(|mut d| { d.source_segments.sort(); let s = d.best_score; (d, s) }) + .collect(); + out.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap()); + out +} + +#[derive(Debug, Clone)] +pub struct DedupCandidate { + pub icd_code: String, + pub best_phrase: String, + pub best_score: f32, + pub source_segments: Vec, +} diff --git a/tests/fusion_tests.rs b/tests/fusion_tests.rs new file mode 100644 index 0000000..5a70d5c --- /dev/null +++ b/tests/fusion_tests.rs @@ -0,0 +1,31 @@ +use alpha_id::fusion::{rrf_merge, cross_segment_dedupe}; +use alpha_id::model::Candidate; + +fn cand(code: &str, seg: usize, rerank: Option) -> Candidate { + Candidate { icd_code: code.into(), alpha_text: code.into(), + segment_idx: seg, lexical: None, semantic: None, rerank } +} + +#[test] +fn rrf_merges_two_rankings_by_reciprocal_rank() { + let lex = vec![0usize, 1, 2]; // entry indices ranked + let sem = vec![2usize, 0, 3]; + let merged = rrf_merge(&lex, &sem, 60); + // entry 0 appears rank0 in lex and rank1 in sem -> highest + assert_eq!(merged[0], 0); + assert!(merged.contains(&3)); +} + +#[test] +fn cross_segment_keeps_max_score_and_all_sources() { + let cands = vec![ + cand("E11.9", 0, Some(0.4)), + cand("E11.9", 2, Some(0.9)), + cand("I10.90", 1, Some(0.5)), + ]; + let mut merged = cross_segment_dedupe(cands); + merged.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap()); + assert_eq!(merged[0].0.icd_code, "E11.9"); + assert_eq!(merged[0].1, 0.9); // max + assert_eq!(merged[0].0.source_segments, vec![0, 2]); // all sources +}