1use crate::graph::{Graph, Metric};
28use crate::Result;
29
30#[derive(Clone, Debug)]
33pub enum ScoreExpr {
34 Const(f32),
35 Add(Box<ScoreExpr>, Box<ScoreExpr>),
36 Sub(Box<ScoreExpr>, Box<ScoreExpr>),
37 Mul(Box<ScoreExpr>, Box<ScoreExpr>),
38 Div(Box<ScoreExpr>, Box<ScoreExpr>),
39 Bm25 { field: u64, query: String },
40 Bm25Norm { field: u64, query: String, k: f32 },
41 VecSim { field: u64, metric: Metric, query: Vec<f32> },
42 StDistanceM { field: u64, lat: f64, lon: f64 },
43 Extern(usize),
44}
45
46enum Atom {
50 Bm25 { field: u64, query: String },
51 VecSim { field: u64, metric: Metric, query: Vec<f32> },
52 StDistanceM { field: u64, lat: f64, lon: f64 },
53 Extern(usize),
54}
55
56impl Atom {
57 fn key(&self) -> String {
58 match self {
59 Atom::Bm25 { field, query } => format!("b:{field}:{query}"),
60 Atom::VecSim { field, metric, query } => {
61 let h: u64 = query.iter().fold(0u64, |a, x| {
62 a.wrapping_mul(1099511628211).wrapping_add(x.to_bits() as u64)
63 });
64 format!("v:{field}:{:?}:{h}", metric)
65 }
66 Atom::StDistanceM { field, lat, lon } => format!("g:{field}:{lat}:{lon}"),
67 Atom::Extern(i) => format!("x:{i}"),
68 }
69 }
70}
71
72fn collect_atoms(e: &ScoreExpr, out: &mut Vec<Atom>) {
73 match e {
74 ScoreExpr::Const(_) => {}
75 ScoreExpr::Add(a, b) | ScoreExpr::Sub(a, b)
76 | ScoreExpr::Mul(a, b) | ScoreExpr::Div(a, b) => {
77 collect_atoms(a, out);
78 collect_atoms(b, out);
79 }
80 ScoreExpr::Bm25 { field, query } | ScoreExpr::Bm25Norm { field, query, .. } => {
81 out.push(Atom::Bm25 { field: *field, query: query.clone() })
82 }
83 ScoreExpr::VecSim { field, metric, query } => {
84 out.push(Atom::VecSim { field: *field, metric: *metric, query: query.clone() })
85 }
86 ScoreExpr::StDistanceM { field, lat, lon } => {
87 out.push(Atom::StDistanceM { field: *field, lat: *lat, lon: *lon })
88 }
89 ScoreExpr::Extern(i) => out.push(Atom::Extern(*i)),
90 }
91}
92
93impl Graph {
94 pub fn hybrid_score(&self, cands: &[u64], expr: &ScoreExpr, k: usize,
98 externs: &[&[f32]]) -> Result<Vec<(u64, f32)>> {
99 let mut atoms: Vec<Atom> = Vec::new();
101 collect_atoms(expr, &mut atoms);
102 let mut cols: std::collections::HashMap<String, Vec<f32>> =
103 std::collections::HashMap::new();
104 for a in &atoms {
106 let key = a.key();
107 if cols.contains_key(&key) { continue; }
108 let col = match a {
109 Atom::Bm25 { field, query } => self.bm25_batch(*field, query, cands)?,
110 Atom::VecSim { field, metric, query } => {
111 let mut c = Vec::with_capacity(cands.len());
112 for &id in cands {
113 c.push(match self.get_vec(*field, id)? {
114 Some(v) => similarity(*metric, &v, query),
115 None => 0.0,
116 });
117 }
118 c
119 }
120 Atom::StDistanceM { field, lat, lon } => {
121 let mut c = Vec::with_capacity(cands.len());
122 for &id in cands {
123 c.push(match self.st_distance(*field, id, *lat, *lon)? {
124 Some(m) => m as f32,
125 None => f32::INFINITY,
126 });
127 }
128 c
129 }
130 Atom::Extern(i) => {
131 let col = externs.get(*i).copied()
132 .expect("Extern(i) references a column the caller did not pass");
133 assert_eq!(col.len(), cands.len(),
134 "extern column length must equal candidate count");
135 col.to_vec()
136 }
137 };
138 cols.insert(key, col);
139 }
140 let mut out: Vec<(u64, f32)> = Vec::with_capacity(cands.len());
142 for (i, &id) in cands.iter().enumerate() {
143 out.push((id, eval(expr, i, &cols)));
144 }
145 out.sort_by(|a, b| b.1.total_cmp(&a.1).then(a.0.cmp(&b.0)));
146 out.truncate(k);
147 Ok(out)
148 }
149
150 fn bm25_batch(&self, field: u64, query: &str, cands: &[u64]) -> Result<Vec<f32>> {
154 self.text_score_candidates(field, query, cands)
155 }
156}
157
158fn similarity(metric: Metric, v: &[f32], q: &[f32]) -> f32 {
159 let dot: f32 = v.iter().zip(q).map(|(a, b)| a * b).sum();
160 match metric {
161 Metric::Dot => dot,
162 Metric::Cosine => {
163 let nv: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
164 let nq: f32 = q.iter().map(|x| x * x).sum::<f32>().sqrt();
165 if nv > 0.0 && nq > 0.0 { dot / (nv * nq) } else { 0.0 }
166 }
167 Metric::L2 => -v.iter().zip(q).map(|(a, b)| (a - b) * (a - b)).sum::<f32>(),
168 Metric::L1 => -v.iter().zip(q).map(|(a, b)| (a - b).abs()).sum::<f32>(),
169 }
170}
171
172fn eval(e: &ScoreExpr, i: usize, cols: &std::collections::HashMap<String, Vec<f32>>) -> f32 {
173 match e {
174 ScoreExpr::Const(c) => *c,
175 ScoreExpr::Add(a, b) => eval(a, i, cols) + eval(b, i, cols),
176 ScoreExpr::Sub(a, b) => eval(a, i, cols) - eval(b, i, cols),
177 ScoreExpr::Mul(a, b) => eval(a, i, cols) * eval(b, i, cols),
178 ScoreExpr::Div(a, b) => {
179 let d = eval(b, i, cols);
180 if d == 0.0 { 0.0 } else { eval(a, i, cols) / d }
181 }
182 ScoreExpr::Bm25 { field, query } => {
183 cols[&Atom::Bm25 { field: *field, query: query.clone() }.key()][i]
184 }
185 ScoreExpr::Bm25Norm { field, query, k } => {
186 let s = cols[&Atom::Bm25 { field: *field, query: query.clone() }.key()][i];
187 if s > 0.0 { s / (s + k.max(f32::MIN_POSITIVE)) } else { 0.0 }
188 }
189 ScoreExpr::VecSim { field, metric, query } => {
190 cols[&Atom::VecSim { field: *field, metric: *metric, query: query.clone() }.key()][i]
191 }
192 ScoreExpr::StDistanceM { field, lat, lon } => {
193 cols[&Atom::StDistanceM { field: *field, lat: *lat, lon: *lon }.key()][i]
194 }
195 ScoreExpr::Extern(i_ext) => cols[&Atom::Extern(*i_ext).key()][i],
196 }
197}