Skip to main content

sie_sdk/
scoring.rs

1//! Client-side `MaxSim` scoring for late-interaction (ColBERT-style) models.
2//!
3//! For a query `Q` and a document `D`, both token-level matrices:
4//!
5//! ```text
6//! score(Q, D) = Σ over query tokens i of ( max over document tokens j of Q[i] · D[j] )
7//! ```
8//!
9//! There is no normalization and no division by the query length: the encoder is expected
10//! to have L2-normalized its output already.
11//!
12//! # Precision
13//!
14//! Late-interaction corpora are usually stored as `f16` to halve their memory cost, but the
15//! dot products must accumulate in `f32` or the scores drift. Both functions here widen to
16//! `f32` immediately before the dot product and never accumulate at `f16`.
17//! [`maxsim_batch`] additionally iterates documents in the outer loop, so an `f16` corpus is
18//! never materialized as `f32` all at once.
19
20use crate::types::Multivector;
21
22/// Rows widened to `f32`, ready for the dot products.
23fn rows(multivector: &Multivector) -> Vec<Vec<f32>> {
24    multivector.to_f32()
25}
26
27fn score_against(query: &[Vec<f32>], document: &[Vec<f32>]) -> f32 {
28    query
29        .iter()
30        .map(|query_token| {
31            document
32                .iter()
33                .map(|document_token| dot(query_token, document_token))
34                .fold(f32::NEG_INFINITY, f32::max)
35        })
36        // A document with no tokens scores zero rather than negative infinity.
37        .map(|best| if best.is_finite() { best } else { 0.0 })
38        .sum()
39}
40
41fn dot(left: &[f32], right: &[f32]) -> f32 {
42    left.iter().zip(right).map(|(a, b)| a * b).sum()
43}
44
45/// Score one query against each document, in the order given.
46///
47/// The result is not sorted; ranking is the caller's decision.
48pub fn maxsim(query: &Multivector, documents: &[Multivector]) -> Vec<f32> {
49    let query = rows(query);
50    documents
51        .iter()
52        // Widened per document, so only one document is held at f32 at a time.
53        .map(|document| score_against(&query, &rows(document)))
54        .collect()
55}
56
57/// Score every query against every document.
58///
59/// Returns one row per query, each holding one score per document, so
60/// `maxsim_batch(queries, docs)[i][j] == maxsim(&queries[i], docs)[j]`.
61pub fn maxsim_batch(queries: &[Multivector], documents: &[Multivector]) -> Vec<Vec<f32>> {
62    let queries: Vec<Vec<Vec<f32>>> = queries.iter().map(rows).collect();
63    let mut scores = vec![vec![0.0f32; documents.len()]; queries.len()];
64
65    // Documents outer, queries inner: the corpus is the large side, and this widens one
66    // document at a time rather than the whole corpus.
67    for (document_index, document) in documents.iter().enumerate() {
68        let document = rows(document);
69        for (query_index, query) in queries.iter().enumerate() {
70            scores[query_index][document_index] = score_against(query, &document);
71        }
72    }
73    scores
74}
75
76#[cfg(test)]
77mod tests {
78    // These assertions are about exact values, so exact comparison is the point.
79    #![allow(clippy::float_cmp)]
80
81    use super::*;
82    use half::f16;
83
84    fn f32_mv(rows: &[&[f32]]) -> Multivector {
85        Multivector::F32(rows.iter().map(|row| row.to_vec()).collect())
86    }
87
88    fn f16_mv(rows: &[&[f32]]) -> Multivector {
89        Multivector::F16(
90            rows.iter()
91                .map(|row| row.iter().map(|v| f16::from_f32(*v)).collect())
92                .collect(),
93        )
94    }
95
96    #[test]
97    fn identical_orthonormal_matrices_score_the_query_length() {
98        let query = f32_mv(&[&[1.0, 0.0], &[0.0, 1.0]]);
99        let scores = maxsim(&query, std::slice::from_ref(&query));
100        assert_eq!(scores.len(), 1);
101        approx::assert_relative_eq!(scores[0], 2.0, epsilon = 1e-5);
102    }
103
104    #[test]
105    fn scores_rank_by_similarity() {
106        let query = f32_mv(&[&[1.0, 0.0]]);
107        let documents = [
108            f32_mv(&[&[1.0, 0.0]]),                           // identical
109            f32_mv(&[&[std::f32::consts::FRAC_1_SQRT_2; 2]]), // 45 degrees
110            f32_mv(&[&[0.0, 1.0]]),                           // orthogonal
111        ];
112        let scores = maxsim(&query, &documents);
113        assert!(scores[0] > scores[1] && scores[1] > scores[2]);
114        approx::assert_relative_eq!(scores[2], 0.0, epsilon = 1e-6);
115    }
116
117    #[test]
118    fn max_is_over_document_tokens_and_sum_is_over_query_tokens() {
119        // Each query token matches a different document token exactly.
120        let query = f32_mv(&[&[1.0, 0.0], &[0.0, 1.0]]);
121        let document = f32_mv(&[&[0.0, 1.0], &[1.0, 0.0], &[0.0, 0.0]]);
122        approx::assert_relative_eq!(maxsim(&query, &[document])[0], 2.0, epsilon = 1e-6);
123    }
124
125    #[test]
126    fn f16_inputs_score_exactly_as_their_widened_selves() {
127        let query = f16_mv(&[&[0.3, -0.7], &[0.1, 0.9]]);
128        let documents = [f16_mv(&[&[0.2, 0.5], &[-0.4, 0.8]]), f16_mv(&[&[1.0, 0.0]])];
129
130        let widened_query = Multivector::F32(query.to_f32());
131        let widened_documents: Vec<Multivector> = documents
132            .iter()
133            .map(|d| Multivector::F32(d.to_f32()))
134            .collect();
135
136        // Bit-exact, not approximately equal: widening must happen before the dot product,
137        // never after partial accumulation at f16.
138        assert_eq!(
139            maxsim(&query, &documents),
140            maxsim(&widened_query, &widened_documents)
141        );
142    }
143
144    #[test]
145    fn batch_agrees_with_the_single_query_form() {
146        let queries = [f32_mv(&[&[1.0, 0.0]]), f32_mv(&[&[0.0, 1.0], &[1.0, 0.0]])];
147        let documents = [f32_mv(&[&[1.0, 0.0]]), f32_mv(&[&[0.6, 0.8]])];
148
149        let batch = maxsim_batch(&queries, &documents);
150        assert_eq!(batch.len(), 2);
151        for (index, query) in queries.iter().enumerate() {
152            let single = maxsim(query, &documents);
153            for (document_index, score) in single.iter().enumerate() {
154                approx::assert_relative_eq!(batch[index][document_index], score, epsilon = 1e-6);
155            }
156        }
157    }
158
159    #[test]
160    fn variable_token_counts_all_produce_finite_scores() {
161        let query = f32_mv(&[&[1.0, 0.0]]);
162        for count in [1usize, 10, 100] {
163            let document = Multivector::F32(
164                (0..count)
165                    .map(|i| vec![i as f32 / count as f32, 0.5])
166                    .collect(),
167            );
168            assert!(maxsim(&query, &[document])[0].is_finite());
169        }
170    }
171
172    #[test]
173    fn empty_inputs_are_handled_rather_than_producing_infinities() {
174        let query = f32_mv(&[&[1.0, 0.0]]);
175        assert!(maxsim(&query, &[]).is_empty());
176        assert_eq!(maxsim(&query, &[Multivector::F32(Vec::new())]), vec![0.0]);
177        assert_eq!(maxsim(&Multivector::F32(Vec::new()), &[query]), vec![0.0]);
178        assert!(maxsim_batch(&[], &[]).is_empty());
179    }
180}