Skip to main content

maxsim_lut/
scorer.rs

1//! Binding a table and a query to score documents.
2
3use crate::kernel::{self, Args};
4use crate::lut::Lut;
5use crate::query::PreparedQuery;
6use crate::Error;
7
8/// A document's per-token centroid ids, in whatever integer width the host
9/// stores them. [`Codes::None`] is valid when the [`Scorer`] carries no
10/// centroid term.
11#[derive(Debug, Clone, Copy)]
12pub enum Codes<'a> {
13    /// No centroid ids (only valid without a centroid term).
14    None,
15    /// `u32` ids.
16    U32(&'a [u32]),
17    /// `i64` ids (next-plaid's on-disk width). Negative values fail range checks.
18    I64(&'a [i64]),
19    /// `usize` ids.
20    Usize(&'a [usize]),
21}
22
23impl Codes<'_> {
24    #[inline(always)]
25    pub(crate) fn id(&self, t: usize) -> usize {
26        match self {
27            Codes::None => 0,
28            Codes::U32(s) => s[t] as usize,
29            // A negative i64 wraps to a huge usize and fails the range check.
30            Codes::I64(s) => s[t] as usize,
31            Codes::Usize(s) => s[t],
32        }
33    }
34
35    fn len(&self) -> Option<usize> {
36        match self {
37            Codes::None => None,
38            Codes::U32(s) => Some(s.len()),
39            Codes::I64(s) => Some(s.len()),
40            Codes::Usize(s) => Some(s.len()),
41        }
42    }
43}
44
45/// One candidate document, borrowed from wherever the host keeps it.
46#[derive(Debug, Clone, Copy)]
47pub struct DocView<'a> {
48    /// Packed residual rows, one per token, contiguous at `row_stride` bytes.
49    /// Must hold at least `n_tokens · row_stride` bytes.
50    pub packed: &'a [u8],
51    /// Number of tokens.
52    pub n_tokens: usize,
53    /// Bytes from one token's row to the next; at least `dim / keys_per_byte`.
54    pub row_stride: usize,
55    /// Per-token centroid ids, indexing the scorer's centroid term.
56    pub codes: Codes<'a>,
57    /// Optional per-token `1 / ‖reconstructed token‖`. `None` scores the
58    /// unnormalised reconstruction (multiplies by 1).
59    pub inv_norms: Option<&'a [f32]>,
60}
61
62impl<'a> DocView<'a> {
63    /// A document with no centroid ids and no normalisation; add them with
64    /// [`DocView::codes`] and [`DocView::inv_norms`].
65    pub fn new(packed: &'a [u8], n_tokens: usize, row_stride: usize) -> Self {
66        Self {
67            packed,
68            n_tokens,
69            row_stride,
70            codes: Codes::None,
71            inv_norms: None,
72        }
73    }
74
75    /// Attach per-token centroid ids.
76    pub fn codes(mut self, codes: Codes<'a>) -> Self {
77        self.codes = codes;
78        self
79    }
80
81    /// Attach per-token inverse norms.
82    pub fn inv_norms(mut self, inv: &'a [f32]) -> Self {
83        self.inv_norms = Some(inv);
84        self
85    }
86}
87
88/// A [`Lut`] and a [`PreparedQuery`] bound together, with the host's
89/// optional centroid term, ready to score documents.
90///
91/// `Copy`, `Send + Sync`: make one per query and share it across the
92/// threads scoring that query's candidates.
93#[derive(Debug, Clone, Copy)]
94pub struct Scorer<'a> {
95    lut: &'a Lut,
96    query: &'a PreparedQuery,
97    cdot: Option<&'a [f32]>,
98    num_centroids: usize,
99}
100
101impl<'a> Scorer<'a> {
102    /// Bind a table and a query. Panics if the query was prepared against a
103    /// different table (use [`Scorer::try_new`] to get an error instead).
104    pub fn new(lut: &'a Lut, query: &'a PreparedQuery) -> Self {
105        Self::try_new(lut, query).expect("PreparedQuery was built against a different Lut")
106    }
107
108    /// Bind a table and a query.
109    pub fn try_new(lut: &'a Lut, query: &'a PreparedQuery) -> Result<Self, Error> {
110        if !query.matches(lut) {
111            return Err(Error::LutMismatch);
112        }
113        Ok(Self {
114            lut,
115            query,
116            cdot: None,
117            num_centroids: 0,
118        })
119    }
120
121    /// Supply the host's query × centroid scores, **centroid-major**:
122    /// `cdot[cid · n_query_tokens + q]`. One centroid's scores across all
123    /// query rows are then contiguous, so the vectorised fold loads them as
124    /// one vector. Hosts that hold the `[n_query_tokens, num_centroids]`
125    /// orientation transpose once per query.
126    pub fn with_centroid_term(
127        mut self,
128        cdot_centroid_major: &'a [f32],
129        num_centroids: usize,
130    ) -> Result<Self, Error> {
131        let nq = self.query.n_tokens();
132        if cdot_centroid_major.len() != num_centroids * nq {
133            return Err(Error::Shape(format!(
134                "centroid term has {} values, expected num_centroids {num_centroids} × n_query_tokens {nq}",
135                cdot_centroid_major.len()
136            )));
137        }
138        self.cdot = Some(cdot_centroid_major);
139        self.num_centroids = num_centroids;
140        Ok(self)
141    }
142
143    /// The table this scorer uses.
144    pub fn lut(&self) -> &'a Lut {
145        self.lut
146    }
147
148    /// The query this scorer uses.
149    pub fn query(&self) -> &'a PreparedQuery {
150        self.query
151    }
152
153    /// MaxSim of the query against one document. Panics on a shape
154    /// violation (see [`Scorer::try_score`] for the checked form); every
155    /// SIMD path returns the same bits as the scalar reference.
156    #[inline]
157    pub fn score(&self, doc: DocView<'_>) -> f32 {
158        match self.try_score(doc) {
159            Ok(s) => s,
160            Err(e) => panic!("maxsim_lut::Scorer::score: {e}"),
161        }
162    }
163
164    /// MaxSim of the query against one document, with shape validation.
165    pub fn try_score(&self, doc: DocView<'_>) -> Result<f32, Error> {
166        let args = self.validate(doc)?;
167        Ok(kernel::maxsim(self.lut, &args))
168    }
169
170    /// Score many documents into `out` (`out.len()` must equal the number of
171    /// documents). Sequential; wrap the call in the host's parallel iterator
172    /// over chunks to fan out.
173    ///
174    /// Panics if the counts disagree in either direction. Silently dropping
175    /// the tail of a candidate list is the kind of bug that shows up as a
176    /// slightly worse recall number months later, not as a failure.
177    pub fn score_many<'d, I>(&self, docs: I, out: &mut [f32])
178    where
179        I: IntoIterator<Item = DocView<'d>>,
180    {
181        let slots = out.len();
182        let mut n = 0usize;
183        for doc in docs {
184            let slot = out
185                .get_mut(n)
186                .unwrap_or_else(|| panic!("score_many: more than {slots} documents"));
187            *slot = self.score(doc);
188            n += 1;
189        }
190        assert_eq!(n, slots, "score_many: {n} documents for {slots} output slots");
191    }
192
193    fn validate<'d>(&self, doc: DocView<'d>) -> Result<Args<'d>, Error>
194    where
195        'a: 'd,
196    {
197        let q = self.query;
198        let dim = q.dim();
199        let kpb = self.lut.keys_per_byte();
200        let pdim = dim / kpb;
201        if doc.row_stride < pdim {
202            return Err(Error::Shape(format!(
203                "row_stride {} < {pdim} packed bytes for dim {dim} at {kpb} keys/byte",
204                doc.row_stride
205            )));
206        }
207        if doc.n_tokens > 0 && doc.packed.len() < (doc.n_tokens - 1) * doc.row_stride + pdim {
208            return Err(Error::Shape(format!(
209                "packed has {} bytes, need {} for {} tokens at stride {}",
210                doc.packed.len(),
211                (doc.n_tokens - 1) * doc.row_stride + pdim,
212                doc.n_tokens,
213                doc.row_stride
214            )));
215        }
216        if let Some(inv) = doc.inv_norms {
217            if inv.len() != doc.n_tokens {
218                return Err(Error::Shape(format!(
219                    "inv_norms has {} values for {} tokens",
220                    inv.len(),
221                    doc.n_tokens
222                )));
223            }
224        }
225        let (cdot, cdot_stride) = match self.cdot {
226            Some(c) => {
227                let n = doc.codes.len().ok_or_else(|| {
228                    Error::Shape("centroid term supplied but DocView has Codes::None".into())
229                })?;
230                if n != doc.n_tokens {
231                    return Err(Error::Shape(format!(
232                        "codes has {n} ids for {} tokens",
233                        doc.n_tokens
234                    )));
235                }
236                for t in 0..doc.n_tokens {
237                    let cid = doc.codes.id(t);
238                    if cid >= self.num_centroids {
239                        return Err(Error::Shape(format!(
240                            "centroid id {cid} at token {t} out of range {}",
241                            self.num_centroids
242                        )));
243                    }
244                }
245                (c, q.n_tokens())
246            }
247            // No centroid term: every token reads the same zero row.
248            None => (q.zeros(), 0usize),
249        };
250        Ok(Args {
251            query: q,
252            packed: doc.packed,
253            row_stride: doc.row_stride,
254            n_tokens: doc.n_tokens,
255            codes: if self.cdot.is_some() {
256                doc.codes
257            } else {
258                Codes::None
259            },
260            cdot,
261            cdot_stride,
262            inv_norms: doc.inv_norms,
263        })
264    }
265}