Skip to main content

kernel/
score.rs

1//! 2l: hybrid scoring -- the kernel primitive SGQL's ScoreExpr lowers onto.
2//!
3//! A newcomer's map. Retrieval FINDS candidates (any family: a text
4//! search, a vector walk, a spatial radius, a graph traversal). Scoring
5//! RANKS them: an arithmetic expression over per-candidate atoms --
6//! BM25 relevance, vector similarity, geodesic distance -- evaluated in
7//! one pass. The structural rule (the contract's): scoring never re-runs
8//! retrieval. Candidates arrive as ids; every atom is a point evaluation
9//! against rows those ids name; cost is candidates x atoms, resident
10//! state is one f32 column per distinct atom, freed at return.
11//!
12//! Atom semantics (documented, not configurable):
13//! - Bm25: Okapi BM25 (K1=1.2, B=0.75), same math as text_search; a
14//!   candidate without the term(s) scores 0.0.
15//! - Bm25Norm: s/(s+k) saturation of the Bm25 score, bounded [0,1] so it
16//!   blends with cosine on equal footing (e1's BM25_NORM).
17//! - VecSim: HIGHER IS BETTER for every metric -- cosine and dot as-is,
18//!   L2/L1 negated. A candidate without a vector IN THAT FIELD scores 0.0.
19//! - StDistanceM: PostGIS-geography metres (Vincenty). A candidate
20//!   without geometry is INFINITELY far (it sorts last under any
21//!   positive distance weighting) -- unknown location never ranks near.
22//! - Extern(i): the caller's per-candidate column (payload-derived
23//!   values computed above the kernel; the payload stays opaque here).
24//! - Division by zero yields 0.0 -- a poisoned NaN would make the final
25//!   ordering depend on sort internals instead of the expression.
26
27use crate::graph::{Graph, Metric};
28use crate::Result;
29
30/// The expression tree. Leaves are atoms or constants; interior nodes
31/// are arithmetic. Built once per query, evaluated per candidate.
32#[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
46/// One resolved atom: what to batch-evaluate. Two Bm25 atoms over the
47/// same (field, query) resolve to ONE column (the tree may reference a
48/// score twice, e.g. raw and saturated; the postings are read once).
49enum 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    /// Rank `cands` by `expr`, descending, top `k`. Ties break on id
95    /// ascending (deterministic under snapshots). `externs` supplies the
96    /// Extern(i) columns; each must be cands.len() long.
97    pub fn hybrid_score(&self, cands: &[u64], expr: &ScoreExpr, k: usize,
98                        externs: &[&[f32]]) -> Result<Vec<(u64, f32)>> {
99        // 1) resolve distinct atoms
100        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        // 2) one batch pass per distinct atom -> a column aligned to cands
105        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        // 3) fold the tree per candidate
141        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    /// BM25 for `cands` only. Live field/term statistics and each candidate's
151    /// compact norm row are point-read; posting extents are retrieval data and
152    /// are never reopened at the scoring boundary required by D33.
153    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}