feat: RRF and cross-segment max-score dedupe

This commit is contained in:
2026-05-18 16:54:25 +02:00
parent a896c5bbde
commit 8d9ab0e299
2 changed files with 85 additions and 1 deletions
+54 -1
View File
@@ -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<usize> {
let mut score: HashMap<usize, f64> = 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<Candidate>) -> Vec<(DedupCandidate, f32)> {
let mut by_code: HashMap<String, DedupCandidate> = 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<usize>,
}
+31
View File
@@ -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<f32>) -> 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
}