1use crate::kernel::{self, Args};
4use crate::lut::Lut;
5use crate::query::PreparedQuery;
6use crate::Error;
7
8#[derive(Debug, Clone, Copy)]
12pub enum Codes<'a> {
13 None,
15 U32(&'a [u32]),
17 I64(&'a [i64]),
19 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 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#[derive(Debug, Clone, Copy)]
47pub struct DocView<'a> {
48 pub packed: &'a [u8],
51 pub n_tokens: usize,
53 pub row_stride: usize,
55 pub codes: Codes<'a>,
57 pub inv_norms: Option<&'a [f32]>,
60}
61
62impl<'a> DocView<'a> {
63 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 pub fn codes(mut self, codes: Codes<'a>) -> Self {
77 self.codes = codes;
78 self
79 }
80
81 pub fn inv_norms(mut self, inv: &'a [f32]) -> Self {
83 self.inv_norms = Some(inv);
84 self
85 }
86}
87
88#[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 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 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 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 pub fn lut(&self) -> &'a Lut {
145 self.lut
146 }
147
148 pub fn query(&self) -> &'a PreparedQuery {
150 self.query
151 }
152
153 #[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 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 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 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}