Skip to main content

maxsim_lut/
query.rs

1//! The query side: symmetric int8 codes, one scale per row, laid out the
2//! way the kernels read them.
3
4use crate::lut::Lut;
5use crate::{padded_stride, Error, MAX_DIM};
6
7/// A query quantised to int8 and pre-arranged for the kernels.
8///
9/// Built once per query against a [`Lut`]; scoring thousands of candidates
10/// reuses it. `Send + Sync`, no interior mutability.
11#[derive(Debug, Clone)]
12pub struct PreparedQuery {
13    nq: usize,
14    dim: usize,
15    /// Row-major int8 codes in dim order, `[nq · dim]` (the scalar kernel's layout).
16    values: Vec<i8>,
17    /// Per-row `max|q| / 127`.
18    scales: Vec<f32>,
19    /// Codes permuted to *plane order* at a padded row stride: plane `k`
20    /// holds the dims byte position `i` carries at key `k`
21    /// (`d = i·keys_per_byte + k`), so the SIMD expand stores each `tbl`
22    /// result contiguously. A dot product is permutation-invariant and the
23    /// integer accumulator is order-invariant, so this changes no result.
24    /// The NEON kernel reads `planes` directly; the other layouts below are
25    /// rearrangements of it, each built only on the architecture whose dot
26    /// instruction needs it.
27    #[cfg_attr(not(target_arch = "aarch64"), allow(dead_code))]
28    planes: Vec<i8>,
29    #[cfg_attr(not(any(target_arch = "aarch64", target_arch = "x86_64")), allow(dead_code))]
30    stride: usize,
31    /// The same plane-order codes as *unsigned* GEMM tiles for the
32    /// `u8 × s8` dot instructions (`vpdpbusd`): `[⌈nq/16⌉ tiles][stride/4
33    /// groups][16 rows][4 bytes]`, each byte `code + 128`. A 64-byte load is
34    /// then 16 rows × 4 consecutive plane dims, multiplied against a 4-byte
35    /// weight broadcast, so the accumulator lanes *are* the row sums and no
36    /// horizontal reduction is needed. The +128 offset is exact: the kernel
37    /// subtracts `128 · Σw` per token. Rows past `nq` hold 128 (code 0).
38    #[cfg(target_arch = "x86_64")]
39    tiles: Vec<u8>,
40    /// The same plane-order codes as *row pairs* for the `smmla` matrix
41    /// instruction: `[⌈nq/2⌉ pairs][stride/8 groups][16 bytes]`, each 16-byte
42    /// group holding 8 consecutive plane dims of row `2p` followed by the
43    /// same 8 dims of row `2p+1`. `smmla` multiplies such a 2×8 query block
44    /// against a 2×8 block of two tokens' weights into a 2×2 accumulator.
45    /// The odd row of an odd `nq` is all zeros.
46    #[cfg(target_arch = "aarch64")]
47    pairs: Vec<i8>,
48    /// Per row: `scales[q] · lut.scale`, the query-constant factor the fold
49    /// applies to each integer accumulator.
50    sqw: Vec<f32>,
51    /// `[nq]` zeros: the centroid row used when the host supplies no centroid term.
52    zeros: Vec<f32>,
53    lut_fingerprint: (usize, u32),
54}
55
56impl PreparedQuery {
57    /// Quantise a query of `n_tokens` rows × `dim` (row-major f32).
58    ///
59    /// Each row is scaled by `max|q| / 127` so its largest component maps to
60    /// ±127; an all-zero row gets scale 0 and codes 0.
61    pub fn new(lut: &Lut, query: &[f32], n_tokens: usize, dim: usize) -> Result<Self, Error> {
62        if dim > MAX_DIM {
63            return Err(Error::DimTooLarge(dim));
64        }
65        if !(dim * lut.nbits()).is_multiple_of(8) {
66            return Err(Error::DimNotByteAligned {
67                dim,
68                nbits: lut.nbits(),
69            });
70        }
71        if query.len() != n_tokens * dim {
72            return Err(Error::Shape(format!(
73                "query has {} values, expected n_tokens {n_tokens} × dim {dim}",
74                query.len()
75            )));
76        }
77        let nq = n_tokens;
78        let mut values = vec![0i8; nq * dim];
79        let mut scales = vec![0.0f32; nq];
80        for (qi, row) in query.chunks_exact(dim).enumerate() {
81            let max_abs = row.iter().fold(0.0f32, |m, &x| m.max(x.abs()));
82            if max_abs <= 0.0 {
83                continue;
84            }
85            let scale = max_abs / 127.0;
86            scales[qi] = scale;
87            for (d, &x) in row.iter().enumerate() {
88                values[qi * dim + d] = (x / scale).round().clamp(-127.0, 127.0) as i8;
89            }
90        }
91        let kpb = lut.keys_per_byte();
92        let pdim = dim / kpb;
93        let stride = padded_stride(dim);
94        let mut planes = vec![0i8; nq * stride];
95        for qi in 0..nq {
96            let row = &values[qi * dim..(qi + 1) * dim];
97            let out = &mut planes[qi * stride..qi * stride + dim];
98            for i in 0..pdim {
99                for k in 0..kpb {
100                    out[k * pdim + i] = row[i * kpb + k];
101                }
102            }
103        }
104        // Each alternative layout is a rearrangement of `planes` that one
105        // architecture's dot instruction needs, so it is built only where a
106        // kernel can consume it. Preparing a query is on the interactive
107        // latency path, and an aarch64 host has no use for GEMM tiles.
108        #[cfg(target_arch = "x86_64")]
109        let tiles = {
110            let d4n = stride / 4;
111            let n16 = nq.div_ceil(16);
112            let mut tiles = vec![128u8; n16 * d4n * 64];
113            for qi in 0..nq {
114                let (t, r) = (qi / 16, qi % 16);
115                for d in 0..dim {
116                    let v = planes[qi * stride + d];
117                    tiles[(t * d4n + d / 4) * 64 + r * 4 + (d % 4)] = (v as i16 + 128) as u8;
118                }
119            }
120            tiles
121        };
122        #[cfg(target_arch = "aarch64")]
123        let pairs = {
124            let npairs = nq.div_ceil(2);
125            let mut pairs = vec![0i8; npairs * 2 * stride];
126            for qi in 0..nq {
127                let (p, r) = (qi / 2, qi % 2);
128                for d in 0..dim {
129                    pairs[p * 2 * stride + (d / 8) * 16 + r * 8 + (d % 8)] = planes[qi * stride + d];
130                }
131            }
132            pairs
133        };
134        let sqw = scales.iter().map(|&s| s * lut.scale()).collect();
135        Ok(Self {
136            nq,
137            dim,
138            values,
139            scales,
140            planes,
141            stride,
142            #[cfg(target_arch = "x86_64")]
143            tiles,
144            #[cfg(target_arch = "aarch64")]
145            pairs,
146            sqw,
147            zeros: vec![0.0f32; nq],
148            lut_fingerprint: lut.fingerprint(),
149        })
150    }
151
152    /// Number of query tokens (rows).
153    pub fn n_tokens(&self) -> usize {
154        self.nq
155    }
156
157    /// Embedding dimension.
158    pub fn dim(&self) -> usize {
159        self.dim
160    }
161
162    /// Row-major int8 codes, `[n_tokens · dim]`.
163    pub fn codes(&self) -> &[i8] {
164        &self.values
165    }
166
167    /// Per-row dequantisation scales.
168    pub fn scales(&self) -> &[f32] {
169        &self.scales
170    }
171
172    #[cfg_attr(not(target_arch = "aarch64"), allow(dead_code))]
173    pub(crate) fn planes(&self) -> &[i8] {
174        &self.planes
175    }
176    #[cfg_attr(not(any(target_arch = "aarch64", target_arch = "x86_64")), allow(dead_code))]
177    pub(crate) fn stride(&self) -> usize {
178        self.stride
179    }
180    /// Unsigned GEMM tiles; see the field docs. Tile `t` starts at
181    /// `t · (stride/4) · 64`; dim group `g` of it at `+ g · 64`.
182    #[cfg(target_arch = "x86_64")]
183    pub(crate) fn tiles_u8(&self) -> &[u8] {
184        &self.tiles
185    }
186    /// Row-pair layout for `smmla`; see the field docs. Pair `p` starts at
187    /// `p · 2 · stride`; dim group `g` (8 dims) of it at `+ g · 16`.
188    #[cfg(target_arch = "aarch64")]
189    pub(crate) fn pairs(&self) -> &[i8] {
190        &self.pairs
191    }
192    pub(crate) fn sqw(&self) -> &[f32] {
193        &self.sqw
194    }
195    pub(crate) fn zeros(&self) -> &[f32] {
196        &self.zeros
197    }
198    pub(crate) fn matches(&self, lut: &Lut) -> bool {
199        self.lut_fingerprint == lut.fingerprint()
200    }
201}
202
203#[cfg(test)]
204mod tests {
205    use super::*;
206
207    #[test]
208    fn quantisation_is_symmetric_per_row() {
209        let lut = Lut::colbert(4, &[0.0; 16]).unwrap();
210        let dim = 8;
211        let q = vec![
212            0.5, -1.0, 0.25, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0,
213        ];
214        let p = PreparedQuery::new(&lut, &q, 2, dim).unwrap();
215        assert_eq!(p.codes()[..3], [64, -127, 32]);
216        assert!((p.scales()[0] - 1.0 / 127.0).abs() < 1e-9);
217        assert_eq!(p.scales()[1], 0.0);
218        assert!(p.codes()[dim..].iter().all(|&c| c == 0));
219        // Planes: nbits 4 → 2 keys/byte, plane 0 = even dims, plane 1 = odd dims.
220        assert_eq!(p.planes()[..4], [64, 32, 0, 0]);
221        assert_eq!(p.planes()[4..8], [-127, 0, 0, 0]);
222        assert_eq!(p.stride(), 64);
223    }
224
225    /// The x86 tile layout. Built only where a kernel reads it, so this test
226    /// runs on the x86 CI runners; `row_pairs_*` is its aarch64 counterpart.
227    #[cfg(target_arch = "x86_64")]
228    #[test]
229    fn gemm_tiles_place_every_row_and_dim() {
230        let lut = Lut::colbert(4, &[0.0; 16]).unwrap();
231        let dim = 8;
232        let q = vec![
233            0.5, -1.0, 0.25, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0, 0.0,
234        ];
235        let p = PreparedQuery::new(&lut, &q, 2, dim).unwrap();
236        // One 16-row tile, 16 dim groups of 64 bytes. Group 0 row 0 =
237        // planes[0..4] + 128; group 1 row 0 = planes[4..8] + 128; row 1 (zero
238        // query) and the 14 padding rows are 128 everywhere.
239        let t = p.tiles_u8();
240        assert_eq!(t.len(), 16 * 64);
241        assert_eq!(&t[0..4], &[192, 160, 128, 128]);
242        assert!(t[4..64].iter().all(|&b| b == 128));
243        assert_eq!(&t[64..68], &[1, 128, 128, 128]);
244        assert!(t[68..].iter().all(|&b| b == 128));
245        // Exhaustive: every (row, dim) lands where the kernel will read it.
246        for qi in 0..2 {
247            for d in 0..dim {
248                let (tile, r) = (qi / 16, qi % 16);
249                let idx = (tile * 16 + d / 4) * 64 + r * 4 + d % 4;
250                assert_eq!(t[idx] as i16 - 128, p.planes()[qi * 64 + d] as i16);
251            }
252        }
253    }
254
255    #[cfg(target_arch = "aarch64")]
256    #[test]
257    fn row_pairs_interleave_eight_dims_at_a_time() {
258        let lut = Lut::colbert(4, &[0.0; 16]).unwrap();
259        let (nq, dim) = (3, 24);
260        let q: Vec<f32> = (0..nq * dim).map(|i| (i % 13) as f32 - 6.0).collect();
261        let p = PreparedQuery::new(&lut, &q, nq, dim).unwrap();
262        let pr = p.pairs();
263        assert_eq!(pr.len(), 2 * 2 * 64, "⌈3/2⌉ pairs × 2 · stride");
264        for qi in 0..nq {
265            for d in 0..dim {
266                let idx = (qi / 2) * 128 + (d / 8) * 16 + (qi % 2) * 8 + d % 8;
267                assert_eq!(pr[idx], p.planes()[qi * 64 + d], "row {qi} dim {d}");
268            }
269        }
270        // The phantom fourth row is zero.
271        for g in 0..3 {
272            assert!(pr[128 + g * 16 + 8..128 + g * 16 + 16].iter().all(|&b| b == 0));
273        }
274    }
275
276    #[test]
277    fn rejects_misaligned_dim() {
278        let lut = Lut::colbert(2, &[0.0; 4]).unwrap();
279        assert_eq!(
280            PreparedQuery::new(&lut, &[0.0; 10], 1, 10).unwrap_err(),
281            Error::DimNotByteAligned { dim: 10, nbits: 2 }
282        );
283        assert_eq!(
284            PreparedQuery::new(&lut, &[0.0; 264], 1, 264).unwrap_err(),
285            Error::DimTooLarge(264)
286        );
287    }
288}