use crate::kernel::{self, Args};
use crate::lut::Lut;
use crate::query::PreparedQuery;
use crate::Error;
#[derive(Debug, Clone, Copy)]
pub enum Codes<'a> {
None,
U32(&'a [u32]),
I64(&'a [i64]),
Usize(&'a [usize]),
}
impl Codes<'_> {
#[inline(always)]
pub(crate) fn id(&self, t: usize) -> usize {
match self {
Codes::None => 0,
Codes::U32(s) => s[t] as usize,
Codes::I64(s) => s[t] as usize,
Codes::Usize(s) => s[t],
}
}
fn len(&self) -> Option<usize> {
match self {
Codes::None => None,
Codes::U32(s) => Some(s.len()),
Codes::I64(s) => Some(s.len()),
Codes::Usize(s) => Some(s.len()),
}
}
}
#[derive(Debug, Clone, Copy)]
pub struct DocView<'a> {
pub packed: &'a [u8],
pub n_tokens: usize,
pub row_stride: usize,
pub codes: Codes<'a>,
pub inv_norms: Option<&'a [f32]>,
}
impl<'a> DocView<'a> {
pub fn new(packed: &'a [u8], n_tokens: usize, row_stride: usize) -> Self {
Self {
packed,
n_tokens,
row_stride,
codes: Codes::None,
inv_norms: None,
}
}
pub fn codes(mut self, codes: Codes<'a>) -> Self {
self.codes = codes;
self
}
pub fn inv_norms(mut self, inv: &'a [f32]) -> Self {
self.inv_norms = Some(inv);
self
}
}
#[derive(Debug, Clone, Copy)]
pub struct Scorer<'a> {
lut: &'a Lut,
query: &'a PreparedQuery,
cdot: Option<&'a [f32]>,
num_centroids: usize,
}
impl<'a> Scorer<'a> {
pub fn new(lut: &'a Lut, query: &'a PreparedQuery) -> Self {
Self::try_new(lut, query).expect("PreparedQuery was built against a different Lut")
}
pub fn try_new(lut: &'a Lut, query: &'a PreparedQuery) -> Result<Self, Error> {
if !query.matches(lut) {
return Err(Error::LutMismatch);
}
Ok(Self {
lut,
query,
cdot: None,
num_centroids: 0,
})
}
pub fn with_centroid_term(
mut self,
cdot_centroid_major: &'a [f32],
num_centroids: usize,
) -> Result<Self, Error> {
let nq = self.query.n_tokens();
if cdot_centroid_major.len() != num_centroids * nq {
return Err(Error::Shape(format!(
"centroid term has {} values, expected num_centroids {num_centroids} × n_query_tokens {nq}",
cdot_centroid_major.len()
)));
}
self.cdot = Some(cdot_centroid_major);
self.num_centroids = num_centroids;
Ok(self)
}
pub fn lut(&self) -> &'a Lut {
self.lut
}
pub fn query(&self) -> &'a PreparedQuery {
self.query
}
#[inline]
pub fn score(&self, doc: DocView<'_>) -> f32 {
match self.try_score(doc) {
Ok(s) => s,
Err(e) => panic!("maxsim_lut::Scorer::score: {e}"),
}
}
pub fn try_score(&self, doc: DocView<'_>) -> Result<f32, Error> {
let args = self.validate(doc)?;
Ok(kernel::maxsim(self.lut, &args))
}
pub fn score_many<'d, I>(&self, docs: I, out: &mut [f32])
where
I: IntoIterator<Item = DocView<'d>>,
{
let slots = out.len();
let mut n = 0usize;
for doc in docs {
let slot = out
.get_mut(n)
.unwrap_or_else(|| panic!("score_many: more than {slots} documents"));
*slot = self.score(doc);
n += 1;
}
assert_eq!(n, slots, "score_many: {n} documents for {slots} output slots");
}
fn validate<'d>(&self, doc: DocView<'d>) -> Result<Args<'d>, Error>
where
'a: 'd,
{
let q = self.query;
let dim = q.dim();
let kpb = self.lut.keys_per_byte();
let pdim = dim / kpb;
if doc.row_stride < pdim {
return Err(Error::Shape(format!(
"row_stride {} < {pdim} packed bytes for dim {dim} at {kpb} keys/byte",
doc.row_stride
)));
}
if doc.n_tokens > 0 && doc.packed.len() < (doc.n_tokens - 1) * doc.row_stride + pdim {
return Err(Error::Shape(format!(
"packed has {} bytes, need {} for {} tokens at stride {}",
doc.packed.len(),
(doc.n_tokens - 1) * doc.row_stride + pdim,
doc.n_tokens,
doc.row_stride
)));
}
if let Some(inv) = doc.inv_norms {
if inv.len() != doc.n_tokens {
return Err(Error::Shape(format!(
"inv_norms has {} values for {} tokens",
inv.len(),
doc.n_tokens
)));
}
}
let (cdot, cdot_stride) = match self.cdot {
Some(c) => {
let n = doc.codes.len().ok_or_else(|| {
Error::Shape("centroid term supplied but DocView has Codes::None".into())
})?;
if n != doc.n_tokens {
return Err(Error::Shape(format!(
"codes has {n} ids for {} tokens",
doc.n_tokens
)));
}
for t in 0..doc.n_tokens {
let cid = doc.codes.id(t);
if cid >= self.num_centroids {
return Err(Error::Shape(format!(
"centroid id {cid} at token {t} out of range {}",
self.num_centroids
)));
}
}
(c, q.n_tokens())
}
None => (q.zeros(), 0usize),
};
Ok(Args {
query: q,
packed: doc.packed,
row_stride: doc.row_stride,
n_tokens: doc.n_tokens,
codes: if self.cdot.is_some() {
doc.codes
} else {
Codes::None
},
cdot,
cdot_stride,
inv_norms: doc.inv_norms,
})
}
}