1use crate::lut::Lut;
5use crate::{padded_stride, Error, MAX_DIM};
6
7#[derive(Debug, Clone)]
12pub struct PreparedQuery {
13 nq: usize,
14 dim: usize,
15 values: Vec<i8>,
17 scales: Vec<f32>,
19 #[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 #[cfg(target_arch = "x86_64")]
39 tiles: Vec<u8>,
40 #[cfg(target_arch = "aarch64")]
47 pairs: Vec<i8>,
48 sqw: Vec<f32>,
51 zeros: Vec<f32>,
53 lut_fingerprint: (usize, u32),
54}
55
56impl PreparedQuery {
57 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 #[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 pub fn n_tokens(&self) -> usize {
154 self.nq
155 }
156
157 pub fn dim(&self) -> usize {
159 self.dim
160 }
161
162 pub fn codes(&self) -> &[i8] {
164 &self.values
165 }
166
167 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 #[cfg(target_arch = "x86_64")]
183 pub(crate) fn tiles_u8(&self) -> &[u8] {
184 &self.tiles
185 }
186 #[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 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 #[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 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 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 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}