use crate::graph::{Graph, Metric};
use crate::Result;
#[derive(Clone, Debug)]
pub enum ScoreExpr {
Const(f32),
Add(Box<ScoreExpr>, Box<ScoreExpr>),
Sub(Box<ScoreExpr>, Box<ScoreExpr>),
Mul(Box<ScoreExpr>, Box<ScoreExpr>),
Div(Box<ScoreExpr>, Box<ScoreExpr>),
Bm25 { field: u64, query: String },
Bm25Norm { field: u64, query: String, k: f32 },
VecSim { field: u64, metric: Metric, query: Vec<f32> },
StDistanceM { field: u64, lat: f64, lon: f64 },
Extern(usize),
}
enum Atom {
Bm25 { field: u64, query: String },
VecSim { field: u64, metric: Metric, query: Vec<f32> },
StDistanceM { field: u64, lat: f64, lon: f64 },
Extern(usize),
}
impl Atom {
fn key(&self) -> String {
match self {
Atom::Bm25 { field, query } => format!("b:{field}:{query}"),
Atom::VecSim { field, metric, query } => {
let h: u64 = query.iter().fold(0u64, |a, x| {
a.wrapping_mul(1099511628211).wrapping_add(x.to_bits() as u64)
});
format!("v:{field}:{:?}:{h}", metric)
}
Atom::StDistanceM { field, lat, lon } => format!("g:{field}:{lat}:{lon}"),
Atom::Extern(i) => format!("x:{i}"),
}
}
}
fn collect_atoms(e: &ScoreExpr, out: &mut Vec<Atom>) {
match e {
ScoreExpr::Const(_) => {}
ScoreExpr::Add(a, b) | ScoreExpr::Sub(a, b)
| ScoreExpr::Mul(a, b) | ScoreExpr::Div(a, b) => {
collect_atoms(a, out);
collect_atoms(b, out);
}
ScoreExpr::Bm25 { field, query } | ScoreExpr::Bm25Norm { field, query, .. } => {
out.push(Atom::Bm25 { field: *field, query: query.clone() })
}
ScoreExpr::VecSim { field, metric, query } => {
out.push(Atom::VecSim { field: *field, metric: *metric, query: query.clone() })
}
ScoreExpr::StDistanceM { field, lat, lon } => {
out.push(Atom::StDistanceM { field: *field, lat: *lat, lon: *lon })
}
ScoreExpr::Extern(i) => out.push(Atom::Extern(*i)),
}
}
impl Graph {
pub fn hybrid_score(&self, cands: &[u64], expr: &ScoreExpr, k: usize,
externs: &[&[f32]]) -> Result<Vec<(u64, f32)>> {
let mut atoms: Vec<Atom> = Vec::new();
collect_atoms(expr, &mut atoms);
let mut cols: std::collections::HashMap<String, Vec<f32>> =
std::collections::HashMap::new();
for a in &atoms {
let key = a.key();
if cols.contains_key(&key) { continue; }
let col = match a {
Atom::Bm25 { field, query } => self.bm25_batch(*field, query, cands)?,
Atom::VecSim { field, metric, query } => {
let mut c = Vec::with_capacity(cands.len());
for &id in cands {
c.push(match self.get_vec(*field, id)? {
Some(v) => similarity(*metric, &v, query),
None => 0.0,
});
}
c
}
Atom::StDistanceM { field, lat, lon } => {
let mut c = Vec::with_capacity(cands.len());
for &id in cands {
c.push(match self.st_distance(*field, id, *lat, *lon)? {
Some(m) => m as f32,
None => f32::INFINITY,
});
}
c
}
Atom::Extern(i) => {
let col = externs.get(*i).copied()
.expect("Extern(i) references a column the caller did not pass");
assert_eq!(col.len(), cands.len(),
"extern column length must equal candidate count");
col.to_vec()
}
};
cols.insert(key, col);
}
let mut out: Vec<(u64, f32)> = Vec::with_capacity(cands.len());
for (i, &id) in cands.iter().enumerate() {
out.push((id, eval(expr, i, &cols)));
}
out.sort_by(|a, b| b.1.total_cmp(&a.1).then(a.0.cmp(&b.0)));
out.truncate(k);
Ok(out)
}
fn bm25_batch(&self, field: u64, query: &str, cands: &[u64]) -> Result<Vec<f32>> {
self.text_score_candidates(field, query, cands)
}
}
fn similarity(metric: Metric, v: &[f32], q: &[f32]) -> f32 {
let dot: f32 = v.iter().zip(q).map(|(a, b)| a * b).sum();
match metric {
Metric::Dot => dot,
Metric::Cosine => {
let nv: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
let nq: f32 = q.iter().map(|x| x * x).sum::<f32>().sqrt();
if nv > 0.0 && nq > 0.0 { dot / (nv * nq) } else { 0.0 }
}
Metric::L2 => -v.iter().zip(q).map(|(a, b)| (a - b) * (a - b)).sum::<f32>(),
Metric::L1 => -v.iter().zip(q).map(|(a, b)| (a - b).abs()).sum::<f32>(),
}
}
fn eval(e: &ScoreExpr, i: usize, cols: &std::collections::HashMap<String, Vec<f32>>) -> f32 {
match e {
ScoreExpr::Const(c) => *c,
ScoreExpr::Add(a, b) => eval(a, i, cols) + eval(b, i, cols),
ScoreExpr::Sub(a, b) => eval(a, i, cols) - eval(b, i, cols),
ScoreExpr::Mul(a, b) => eval(a, i, cols) * eval(b, i, cols),
ScoreExpr::Div(a, b) => {
let d = eval(b, i, cols);
if d == 0.0 { 0.0 } else { eval(a, i, cols) / d }
}
ScoreExpr::Bm25 { field, query } => {
cols[&Atom::Bm25 { field: *field, query: query.clone() }.key()][i]
}
ScoreExpr::Bm25Norm { field, query, k } => {
let s = cols[&Atom::Bm25 { field: *field, query: query.clone() }.key()][i];
if s > 0.0 { s / (s + k.max(f32::MIN_POSITIVE)) } else { 0.0 }
}
ScoreExpr::VecSim { field, metric, query } => {
cols[&Atom::VecSim { field: *field, metric: *metric, query: query.clone() }.key()][i]
}
ScoreExpr::StDistanceM { field, lat, lon } => {
cols[&Atom::StDistanceM { field: *field, lat: *lat, lon: *lon }.key()][i]
}
ScoreExpr::Extern(i_ext) => cols[&Atom::Extern(*i_ext).key()][i],
}
}