Skip to main content

maxsim_lut/
lut.rs

1//! The document-side table: packed byte → int8 bucket weights.
2
3use crate::kernel::{self, Kernel};
4use crate::packing::Packing;
5use crate::Error;
6
7/// Highest embedding dimension the kernels support. The SIMD expansion
8/// buffer is `[i8; MAX_DIM]` and the AVX-512 dot reads it in 64-lane chunks,
9/// which `dim ≤ 256` keeps in bounds for every byte-aligned `dim`.
10pub const MAX_DIM: usize = 256;
11
12/// The fused table factored per key position into 16-entry nibble tables,
13/// the shape NEON `tbl` / SSE `pshufb` consume: one in-register lookup per
14/// key position per 16 packed bytes.
15///
16/// Codes of width 1, 2 or 4 never cross a nibble boundary, so key `k` of a
17/// byte is a function of exactly one of its nibbles. The factorisation is
18/// *verified* over all 256 byte values when the [`Lut`] is built, so a
19/// [`Packing`] the tables cannot represent falls back to the scalar path
20/// instead of silently diverging from it.
21#[derive(Debug, Clone, PartialEq, Eq)]
22pub struct NibbleTables {
23    /// Per key position: weights indexed by the source nibble's value.
24    pub tables: [[i8; 16]; 8],
25    /// Whether key `k` reads the byte's high nibble (else the low one).
26    pub from_hi: [bool; 8],
27}
28
29/// The document-side lookup state for one residual codec: a table turning
30/// each packed residual byte directly into its `8/nbits` int8 bucket
31/// weights, plus the dequantisation scale.
32///
33/// Build once per index (it depends only on the bucket weights and the
34/// packing), share across threads.
35#[derive(Debug, Clone)]
36pub struct Lut {
37    /// `[256 · keys_per_byte]` int8 weights; row `b` is the expansion of byte `b`.
38    fused: Vec<i8>,
39    keys_per_byte: usize,
40    nbits: usize,
41    /// `fused as f32 · scale ≈ bucket_weight`.
42    scale: f32,
43    nibble: Option<NibbleTables>,
44    force_scalar: bool,
45    pin: Option<Kernel>,
46}
47
48impl Lut {
49    /// Build the table from a packing layout and the codec's `2^nbits` bucket
50    /// weights (f32, in bucket-index order).
51    ///
52    /// Weights are quantised symmetrically to int8 with `scale = max|w| / 127`.
53    pub fn new<P: Packing>(packing: &P, bucket_weights: &[f32]) -> Result<Self, Error> {
54        let nbits = packing.nbits();
55        if !matches!(nbits, 1 | 2 | 4 | 8) {
56            return Err(Error::NbitsUnsupported(nbits));
57        }
58        let expected = 1usize << nbits;
59        if bucket_weights.len() != expected {
60            return Err(Error::BucketCount {
61                expected,
62                got: bucket_weights.len(),
63            });
64        }
65        let max_abs = bucket_weights.iter().fold(0.0f32, |m, &x| m.max(x.abs()));
66        let scale = (max_abs / 127.0).max(1e-12);
67        let vals: Vec<i8> = bucket_weights
68            .iter()
69            .map(|&w| (w / scale).round().clamp(-127.0, 127.0) as i8)
70            .collect();
71        let keys_per_byte = 8 / nbits;
72        let mut fused = vec![0i8; 256 * keys_per_byte];
73        for byte in 0..256usize {
74            for k in 0..keys_per_byte {
75                let bi = packing.bucket_index(byte as u8, k);
76                assert!(
77                    bi < expected,
78                    "Packing::bucket_index returned {bi} >= 2^nbits for byte {byte} key {k}"
79                );
80                fused[byte * keys_per_byte + k] = vals[bi];
81            }
82        }
83        let nibble = derive_nibble_tables(&fused, keys_per_byte);
84        Ok(Self {
85            fused,
86            keys_per_byte,
87            nbits,
88            scale,
89            nibble,
90            force_scalar: false,
91            pin: None,
92        })
93    }
94
95    /// Convenience for the ColBERT / PLAID layout.
96    pub fn colbert(nbits: usize, bucket_weights: &[f32]) -> Result<Self, Error> {
97        Self::new(&crate::ColbertPacking::new(nbits)?, bucket_weights)
98    }
99
100    /// Pin every score to the scalar reference kernel. For tests and for
101    /// measuring what the SIMD is worth; the results are bit-identical either
102    /// way. The environment variable `MAXSIM_LUT_FORCE_SCALAR=1` has the same
103    /// effect process-wide.
104    pub fn force_scalar(mut self, yes: bool) -> Self {
105        self.force_scalar = yes;
106        self
107    }
108
109    /// Pin dispatch to one kernel instead of the calibrated choice.
110    ///
111    /// Only a kernel this CPU can execute is honoured (check with
112    /// [`crate::supported_kernels`]); anything else is ignored and dispatch
113    /// proceeds normally, because a kernel the CPU cannot run has no
114    /// meaningful behaviour to fall back to. Scoring is bit-identical
115    /// whichever kernel runs, so this only affects speed.
116    ///
117    /// Use it to benchmark one path, or to make dispatch deterministic on a
118    /// fleet of mixed cores. `None` restores the default.
119    pub fn pin_kernel(mut self, kernel: Option<Kernel>) -> Self {
120        self.pin = kernel;
121        self
122    }
123
124    /// Code width this table was built for.
125    pub fn nbits(&self) -> usize {
126        self.nbits
127    }
128
129    /// `8 / nbits`: how many dims one packed byte carries.
130    pub fn keys_per_byte(&self) -> usize {
131        self.keys_per_byte
132    }
133
134    /// Dequantisation scale: `fused as f32 · scale ≈ bucket_weight`.
135    pub fn scale(&self) -> f32 {
136        self.scale
137    }
138
139    /// The `8/nbits` int8 weights a packed byte expands to, in dim order.
140    #[inline]
141    pub fn expand(&self, byte: u8) -> &[i8] {
142        let base = byte as usize * self.keys_per_byte;
143        &self.fused[base..base + self.keys_per_byte]
144    }
145
146    /// The whole fused table, `[256 · keys_per_byte]`, row `b` = byte `b`.
147    pub fn fused_table(&self) -> &[i8] {
148        &self.fused
149    }
150
151    /// The nibble-factored tables, if the layout admits them (always, for
152    /// [`crate::ColbertPacking`] at nbits 1, 2 or 4; never at nbits 8).
153    pub fn nibble_tables(&self) -> Option<&NibbleTables> {
154        self.nibble.as_ref()
155    }
156
157    /// Which kernel [`crate::Scorer::score`] will run for this table and
158    /// `dim` on this CPU. Print it next to any benchmark number: a speedup
159    /// attributed to a path that never executed is the easiest measurement
160    /// error to make and the hardest to notice.
161    pub fn kernel(&self, dim: usize) -> Kernel {
162        kernel::select(self, dim)
163    }
164
165    pub(crate) fn force_scalar_set(&self) -> bool {
166        self.force_scalar
167    }
168
169    pub(crate) fn pinned_kernel(&self) -> Option<Kernel> {
170        self.pin
171    }
172
173    /// Two [`Lut`]s prepared from the same weights and packing are
174    /// interchangeable for a [`crate::PreparedQuery`]; this is the identity
175    /// the query checks.
176    pub(crate) fn fingerprint(&self) -> (usize, u32) {
177        (self.nbits, self.scale.to_bits())
178    }
179}
180
181/// Factor `fused` into per-key nibble tables; `None` if any key position is
182/// not a function of a single nibble.
183fn derive_nibble_tables(fused: &[i8], keys_per_byte: usize) -> Option<NibbleTables> {
184    if keys_per_byte > 8 || keys_per_byte == 1 {
185        // nbits 8: a key spans the whole byte, no 16-entry factorisation.
186        return None;
187    }
188    let mut tables = [[0i8; 16]; 8];
189    let mut from_hi = [false; 8];
190    for k in 0..keys_per_byte {
191        let hi: [i8; 16] = std::array::from_fn(|x| fused[(x << 4) * keys_per_byte + k]);
192        if (0..256).all(|b| fused[b * keys_per_byte + k] == hi[b >> 4]) {
193            tables[k] = hi;
194            from_hi[k] = true;
195            continue;
196        }
197        let lo: [i8; 16] = std::array::from_fn(|x| fused[x * keys_per_byte + k]);
198        if (0..256).all(|b| fused[b * keys_per_byte + k] == lo[b & 15]) {
199            tables[k] = lo;
200            from_hi[k] = false;
201            continue;
202        }
203        return None;
204    }
205    Some(NibbleTables { tables, from_hi })
206}
207
208#[cfg(test)]
209mod tests {
210    use super::*;
211    use crate::ColbertPacking;
212
213    fn weights(nbits: usize) -> Vec<f32> {
214        let n = 1usize << nbits;
215        (0..n)
216            .map(|i| -0.35 + 0.7 * (i as f32 + 0.5) / n as f32)
217            .collect()
218    }
219
220    #[test]
221    fn fused_table_expands_to_quantised_bucket_weights() {
222        for nbits in [1usize, 2, 4, 8] {
223            let p = ColbertPacking::new(nbits).unwrap();
224            let w = weights(nbits);
225            let lut = Lut::new(&p, &w).unwrap();
226            for byte in 0..=255u8 {
227                for k in 0..lut.keys_per_byte() {
228                    let bi = p.bucket_index(byte, k);
229                    let want = (w[bi] / lut.scale()).round().clamp(-127.0, 127.0) as i8;
230                    assert_eq!(lut.expand(byte)[k], want, "nbits={nbits} byte={byte} k={k}");
231                }
232            }
233        }
234    }
235
236    #[test]
237    fn nibble_factorisation_holds_for_sub_byte_codes() {
238        for nbits in [1usize, 2, 4] {
239            let lut = Lut::colbert(nbits, &weights(nbits)).unwrap();
240            let nib = lut
241                .nibble_tables()
242                .unwrap_or_else(|| panic!("nbits={nbits}: not nibble-separable"));
243            for b in 0..256usize {
244                for k in 0..lut.keys_per_byte() {
245                    let nibble = if nib.from_hi[k] { b >> 4 } else { b & 15 };
246                    assert_eq!(
247                        lut.fused_table()[b * lut.keys_per_byte() + k],
248                        nib.tables[k][nibble]
249                    );
250                }
251            }
252        }
253        assert!(Lut::colbert(8, &weights(8)).unwrap().nibble_tables().is_none());
254    }
255
256    #[test]
257    fn rejects_bad_inputs() {
258        assert_eq!(ColbertPacking::new(3).unwrap_err(), Error::NbitsUnsupported(3));
259        assert_eq!(
260            Lut::colbert(4, &[0.0; 15]).unwrap_err(),
261            Error::BucketCount {
262                expected: 16,
263                got: 15
264            }
265        );
266    }
267}