1use crate::format::{f16_bits_to_f32, Error, MAGIC, V1_2, V1_3, V1_4, V1_5};
17use crate::rhdh::Rhdh;
18use crate::store::{score_batch4, score_raw_dispatch};
19
20fn rd_u16(b: &[u8]) -> u16 {
21 u16::from_le_bytes([b[0], b[1]])
22}
23
24fn rd_u32(b: &[u8]) -> u32 {
25 u32::from_le_bytes([b[0], b[1], b[2], b[3]])
26}
27
28pub struct ViewQuery {
32 rotated: Vec<f32>,
33 lut: [f32; 16],
34}
35
36pub struct VecqView<'a> {
39 dim: usize,
40 working_dim: usize,
41 padded: usize,
42 n: usize,
43 bits: u8,
44 residual: bool,
45 transform: Rhdh,
46 codes: &'a [u8],
47 scales_raw: &'a [u8], codes2: Option<&'a [u8]>,
49 scales2_raw: Option<&'a [u8]>,
50}
51
52impl<'a> VecqView<'a> {
53 pub fn from_bytes(bytes: &'a [u8]) -> Result<Self, Error> {
57 if bytes.len() < 24 || rd_u32(&bytes[0..4]) != MAGIC {
58 return Err(Error::NotAStableFile);
59 }
60 let version = rd_u16(&bytes[4..6]);
61 if version != V1_2 && version != V1_3 && version != V1_4 && version != V1_5 {
62 return Err(Error::UnsupportedVersion(version));
63 }
64 let dim = rd_u32(&bytes[8..12]) as usize;
65 let seed = u64::from_le_bytes(bytes[12..20].try_into().unwrap());
66 let count = rd_u32(&bytes[20..24]) as usize;
67 let working_dim = match rd_u16(&bytes[6..8]) as usize {
68 0 => dim,
69 w if w <= dim => w,
70 w => {
71 return Err(Error::InvalidWorkingDim {
72 dim,
73 working_dim: w,
74 })
75 }
76 };
77 let mut off = 24usize;
78 let bits = if version == V1_5 {
79 if bytes.len() < 25 {
80 return Err(Error::Truncated);
81 }
82 let w = bytes[24];
83 if !matches!(w, 4..=6) {
84 return Err(Error::InvalidWidth { width: w });
85 }
86 off += 1;
87 w
88 } else {
89 4
90 };
91 let padded = crate::rhdh::padded_dim(working_dim);
92 let codes_bytes = (padded * bits as usize).div_ceil(8);
93 let expected = off + count * (2 + codes_bytes);
94 if bytes.len() < expected {
95 return Err(Error::Truncated);
96 }
97 let scales_raw = &bytes[off..off + count * 2];
98 let codes = &bytes[off + count * 2..off + count * (2 + codes_bytes)];
99 off += count * (2 + codes_bytes);
100 let mut residual = false;
101 let mut codes2 = None;
102 let mut scales2_raw = None;
103 if version == V1_4 {
104 if bytes.len() < off + count * (2 + codes_bytes) {
105 return Err(Error::Truncated);
106 }
107 scales2_raw = Some(&bytes[off..off + count * 2]);
108 codes2 = Some(&bytes[off + count * 2..off + count * (2 + codes_bytes)]);
109 off += count * (2 + codes_bytes);
110 residual = true;
111 }
112 if version == V1_3 || version == V1_4 || version == V1_5 {
116 if bytes.len() < off + 4 {
117 return Err(Error::Truncated);
118 }
119 let entries = rd_u32(&bytes[off..off + 4]) as usize;
120 if bytes.len() < off + 4 + entries * 12 {
121 return Err(Error::Truncated);
122 }
123 }
124 Ok(Self {
125 dim,
126 working_dim,
127 padded,
128 n: count,
129 bits,
130 residual,
131 transform: Rhdh::new(working_dim, seed),
132 codes,
133 scales_raw,
134 codes2,
135 scales2_raw,
136 })
137 }
138
139 pub fn len(&self) -> usize {
140 self.n
141 }
142
143 pub fn is_empty(&self) -> bool {
144 self.n == 0
145 }
146
147 pub fn dim(&self) -> usize {
148 self.dim
149 }
150
151 pub fn working_dim(&self) -> usize {
152 self.working_dim
153 }
154
155 pub fn bits(&self) -> u8 {
157 self.bits
158 }
159
160 pub fn is_residual(&self) -> bool {
162 self.residual
163 }
164
165 pub fn prepare_query(&self, q: &[f32]) -> ViewQuery {
169 assert_eq!(q.len(), self.dim);
170 let norm: f32 = q[..self.working_dim]
171 .iter()
172 .map(|x| x * x)
173 .sum::<f32>()
174 .sqrt();
175 assert!(norm > 0.0, "zero vector");
176 let unit: Vec<f32> = q[..self.working_dim].iter().map(|x| x / norm).collect();
177 let mut rotated = Vec::with_capacity(self.padded);
178 self.transform.apply(&unit, &mut rotated);
179 let rnorm: f32 = rotated.iter().map(|x| x * x).sum::<f32>().sqrt();
180 for x in rotated.iter_mut() {
181 *x /= rnorm;
182 }
183 let mut lut = [0f32; 16];
184 for (c, slot) in lut.iter_mut().enumerate() {
185 *slot = crate::lloyd::dequantize_4bit(c as u8);
186 }
187 ViewQuery { rotated, lut }
188 }
189
190 pub fn score(&self, pq: &ViewQuery, idx: usize) -> f32 {
193 let base = idx * self.bytes_per_vector();
194 let codes = &self.codes[base..base + self.bytes_per_vector()];
195 let q = &pq.rotated[..self.padded];
196 let raw0 = score_raw_dispatch(codes, q, &pq.lut, self.bits);
197 if !self.residual {
198 let s = u16::from_le_bytes(self.scales_raw[idx * 2..idx * 2 + 2].try_into().unwrap());
199 return raw0 * f16_bits_to_f32(s);
200 }
201 let codes1 = &self.codes2.unwrap()[base..base + self.bytes_per_vector()];
202 let raw1 = score_raw_dispatch(codes1, q, &pq.lut, self.bits);
203 let s = u16::from_le_bytes(self.scales_raw[idx * 2..idx * 2 + 2].try_into().unwrap());
204 let s2 = u16::from_le_bytes(
205 self.scales2_raw.unwrap()[idx * 2..idx * 2 + 2]
206 .try_into()
207 .unwrap(),
208 );
209 raw0 * f16_bits_to_f32(s) + raw1 * f16_bits_to_f32(s2)
210 }
211
212 pub fn search(&self, q: &[f32], k: usize) -> Vec<(usize, f32)> {
216 use std::cmp::Reverse;
217 use std::collections::BinaryHeap;
218 let pq = self.prepare_query(q);
219 let k = k.min(self.n).max(1);
220 let bpv = self.bytes_per_vector();
221 let key = |s: f32| -> u32 {
222 let b = s.to_bits();
223 if b & 0x8000_0000 != 0 {
224 !b
225 } else {
226 b ^ 0x8000_0000
227 }
228 };
229 let mut heap: BinaryHeap<Reverse<(u32, usize)>> = BinaryHeap::with_capacity(k + 1);
230 let consider = |s: f32, idx: usize, heap: &mut BinaryHeap<Reverse<(u32, usize)>>| {
231 let ks = key(s);
232 if heap.len() < k {
233 heap.push(Reverse((ks, idx)));
234 } else if ks > heap.peek().map(|r| r.0 .0).unwrap_or(0) {
235 heap.push(Reverse((ks, idx)));
236 heap.pop();
237 }
238 };
239 let combine = |r0: f32, r1: Option<f32>, si: usize| -> f32 {
240 match r1 {
241 Some(r1) => {
242 let s =
243 u16::from_le_bytes(self.scales_raw[si * 2..si * 2 + 2].try_into().unwrap());
244 let s2 = u16::from_le_bytes(
245 self.scales2_raw.unwrap()[si * 2..si * 2 + 2]
246 .try_into()
247 .unwrap(),
248 );
249 r0 * f16_bits_to_f32(s) + r1 * f16_bits_to_f32(s2)
250 }
251 None => {
252 let s =
253 u16::from_le_bytes(self.scales_raw[si * 2..si * 2 + 2].try_into().unwrap());
254 r0 * f16_bits_to_f32(s)
255 }
256 }
257 };
258 let q_rot = &pq.rotated[..self.padded];
259 let mut idx = 0;
260 while idx + 4 <= self.n {
263 let codes4 = &self.codes[idx * bpv..(idx + 4) * bpv];
264 let raw = score_batch4(codes4, q_rot, &pq.lut, self.bits);
265 let raw1 = if self.residual {
266 Some(score_batch4(
267 &self.codes2.unwrap()[idx * bpv..(idx + 4) * bpv],
268 q_rot,
269 &pq.lut,
270 self.bits,
271 ))
272 } else {
273 None
274 };
275 for (v, &r) in raw.iter().enumerate() {
276 consider(combine(r, raw1.map(|a| a[v]), idx + v), idx + v, &mut heap);
277 }
278 idx += 4;
279 }
280 while idx < self.n {
281 consider(self.score(&pq, idx), idx, &mut heap);
282 idx += 1;
283 }
284 let key_undo = |k: u32| -> u32 {
285 if k & 0x8000_0000 != 0 {
286 k ^ 0x8000_0000
287 } else {
288 !k
289 }
290 };
291 let mut out: Vec<(usize, f32)> = heap
292 .into_iter()
293 .map(|r| (r.0 .1, f32::from_bits(key_undo(r.0 .0))))
294 .collect();
295 out.sort_by(|a, b| b.1.partial_cmp(&a.1).expect("no NaN scores"));
296 out
297 }
298
299 fn bytes_per_vector(&self) -> usize {
301 (self.padded * self.bits as usize).div_ceil(8)
302 }
303}
304
305#[cfg(test)]
306mod tests {
307 use super::VecqView;
308 use crate::store::VecqIndex;
309
310 fn rand_unit(dim: usize, salt: u64) -> Vec<f32> {
311 let mut x = salt | 1;
312 let mut v = Vec::with_capacity(dim);
313 for _ in 0..dim {
314 x ^= x << 13;
315 x ^= x >> 7;
316 x ^= x << 17;
317 v.push((x as f32 / u32::MAX as f32 - 0.5) * 2.0);
318 }
319 let n: f32 = v.iter().map(|a| a * a).sum::<f32>().sqrt();
320 v.iter_mut().for_each(|a| *a /= n);
321 v
322 }
323
324 #[test]
325 fn view_matches_loaded_index_bitwise() {
326 let dim = 128;
329 for bits in [4u8, 5, 6] {
330 let mut idx = VecqIndex::new(dim, 42);
331 idx.set_bits(bits);
332 for i in 0..20 {
333 idx.add(&rand_unit(dim, i + 11));
334 }
335 let bytes = idx.to_bytes();
336 let loaded = VecqIndex::from_bytes(&bytes).unwrap();
337 let view = VecqView::from_bytes(&bytes).unwrap();
338 assert_eq!(view.len(), 20);
339 assert_eq!(view.bits(), bits);
340 for qi in 0..5 {
341 let q = rand_unit(dim, 900 + qi);
342 let a = loaded.search(&q, 7);
343 let b = view.search(&q, 7);
344 assert_eq!(a.len(), b.len(), "bits {bits} q{qi}");
345 for ((sa, fa), (sb, fb)) in a.iter().zip(b.iter()) {
346 assert_eq!(sa, sb, "bits {bits} q{qi}");
347 assert_eq!(fa.to_bits(), fb.to_bits(), "bits {bits} q{qi}");
348 }
349 }
350 }
351 }
352
353 #[test]
354 fn view_supports_residual_and_working_dim() {
355 let dim = 128;
358 let mut resid = VecqIndex::with_residual(dim, 7);
359 for i in 0..12 {
360 resid.add(&rand_unit(dim, i + 300));
361 }
362 let bytes = resid.to_bytes();
363 let loaded = VecqIndex::from_bytes(&bytes).unwrap();
364 let view = VecqView::from_bytes(&bytes).unwrap();
365 assert!(view.is_residual());
366 let q = rand_unit(dim, 500);
367 for (sa, fa) in loaded.search(&q, 5) {
368 let (_, fb) = view
369 .search(&q, 5)
370 .into_iter()
371 .find(|(sb, _)| *sb == sa)
372 .unwrap();
373 assert_eq!(fa.to_bits(), fb.to_bits());
374 }
375
376 let mut wd = VecqIndex::with_working_dim(256, 64, 21);
377 for i in 0..10 {
378 wd.add(&rand_unit(256, i + 700));
379 }
380 let bytes = wd.to_bytes();
381 let loaded = VecqIndex::from_bytes(&bytes).unwrap();
382 let view = VecqView::from_bytes(&bytes).unwrap();
383 assert_eq!(view.working_dim(), 64);
384 let q = rand_unit(256, 999);
385 for (sa, fa) in loaded.search(&q, 5) {
386 let (_, fb) = view
387 .search(&q, 5)
388 .into_iter()
389 .find(|(sb, _)| *sb == sa)
390 .unwrap();
391 assert_eq!(fa.to_bits(), fb.to_bits());
392 }
393 }
394
395 #[test]
396 fn view_rejects_v1_and_truncated_bytes() {
397 let mut idx = VecqIndex::new(64, 3);
398 idx.add(&rand_unit(64, 1));
399 let bytes = idx.to_bytes();
400 let mut v1 = bytes.clone();
403 v1[4] = 1;
404 v1[5] = 0;
405 assert!(matches!(
406 VecqView::from_bytes(&v1),
407 Err(crate::format::Error::UnsupportedVersion(1))
408 ));
409 assert!(VecqView::from_bytes(&bytes[..10]).is_err());
411 let mut bad = bytes.clone();
412 bad[0] = b'X';
413 assert!(VecqView::from_bytes(&bad).is_err());
414 assert!(VecqView::from_bytes(&bytes[..bytes.len() - 1]).is_err());
416 }
417}