urna_runtime/bm25/
index.rs1use std::collections::HashMap;
5
6use super::tokenize::tokenize;
7
8pub const BM25_PAYLOAD_VERSION: u32 = 2;
15pub const DEFAULT_K1: f32 = 1.5;
16pub const DEFAULT_B: f32 = 0.75;
17
18#[derive(Clone, Debug)]
19pub(super) struct Posting {
20 pub doc: u32,
21 pub tf: u32,
22}
23
24#[derive(Clone, Debug)]
25pub(super) struct TermEntry {
26 pub df: u32,
27 pub postings: Vec<Posting>,
28}
29
30pub struct Bm25Index {
31 pub k1: f32,
32 pub b: f32,
33 pub avgdl: f32,
34 pub n_docs: usize,
35 pub n_terms: usize,
36 pub(super) doc_lengths: Vec<u32>,
37 pub(super) terms: HashMap<String, TermEntry>,
39}
40
41impl Bm25Index {
42 pub fn build(docs: &[String], k1: f32, b: f32) -> Self {
44 let n_docs = docs.len();
45 let mut doc_lengths = Vec::with_capacity(n_docs);
46 let mut term_postings: HashMap<String, Vec<Posting>> = HashMap::new();
47 for (doc_id, doc) in docs.iter().enumerate() {
48 let tokens = tokenize(doc);
49 doc_lengths.push(tokens.len() as u32);
50 let mut tf_map: HashMap<&str, u32> = HashMap::new();
51 for t in &tokens {
52 *tf_map.entry(t.as_str()).or_insert(0) += 1;
53 }
54 for (term, tf) in tf_map {
55 term_postings
56 .entry(term.to_string())
57 .or_default()
58 .push(Posting {
59 doc: doc_id as u32,
60 tf,
61 });
62 }
63 }
64 let total_dl: u64 = doc_lengths.iter().map(|x| *x as u64).sum();
65 let avgdl = if n_docs == 0 {
66 0.0
67 } else {
68 total_dl as f32 / n_docs as f32
69 };
70 let mut entries: Vec<(String, Vec<Posting>)> = term_postings.into_iter().collect();
72 entries.sort_by(|a, b| a.0.cmp(&b.0));
73 let mut terms: HashMap<String, TermEntry> = HashMap::with_capacity(entries.len());
74 for (key, mut postings) in entries {
75 postings.sort_by_key(|p| p.doc);
76 terms.insert(
77 key,
78 TermEntry {
79 df: postings.len() as u32,
80 postings,
81 },
82 );
83 }
84 let n_terms = terms.len();
85 Self {
86 k1,
87 b,
88 avgdl,
89 n_docs,
90 n_terms,
91 doc_lengths,
92 terms,
93 }
94 }
95
96 pub fn search(&self, query_text: &str, k: usize) -> Vec<(usize, f32)> {
99 if k == 0 || self.n_docs == 0 {
100 return Vec::new();
101 }
102 let q_tokens = tokenize(query_text);
103 if q_tokens.is_empty() {
104 return Vec::new();
105 }
106 let mut scores: HashMap<u32, f32> = HashMap::new();
107 for term in q_tokens {
108 let Some(entry) = self.terms.get(&term) else {
109 continue;
110 };
111 let idf =
112 ((self.n_docs as f32 - entry.df as f32 + 0.5) / (entry.df as f32 + 0.5) + 1.0).ln();
113 for p in &entry.postings {
114 let dl = self.doc_lengths.get(p.doc as usize).copied().unwrap_or(0) as f32;
117 let denom =
118 p.tf as f32 + self.k1 * (1.0 - self.b + self.b * dl / self.avgdl.max(1.0));
119 let term_score = idf * (p.tf as f32 * (self.k1 + 1.0)) / denom;
120 *scores.entry(p.doc).or_insert(0.0) += term_score;
121 }
122 }
123 let mut all: Vec<(usize, f32)> = scores.into_iter().map(|(d, s)| (d as usize, s)).collect();
124 all.retain(|p| p.1.is_finite());
127 all.sort_by(|a, b| crate::order::cmp_score_desc(*a, *b));
128 all.truncate(k);
129 all
130 }
131}