nodedb_codec/vector_quant/
rabitq.rs1use crate::error::CodecError;
33use crate::vector_quant::codec::VectorCodec;
34use crate::vector_quant::codec_envelope;
35use crate::vector_quant::hamming::hamming_distance;
36use crate::vector_quant::layout::{QuantHeader, QuantMode, UnifiedQuantizedVector};
37use serde::{Deserialize, Serialize};
38
39#[inline]
43fn xorshift64(state: &mut u64) -> u64 {
44 let mut x = *state;
45 x ^= x << 13;
46 x ^= x >> 7;
47 x ^= x << 17;
48 *state = x;
49 x
50}
51
52#[inline]
56fn next_pow2(n: usize) -> usize {
57 if n.is_power_of_two() {
58 n
59 } else {
60 n.next_power_of_two()
61 }
62}
63
64fn wht_inplace(buf: &mut [f32]) {
68 let n = buf.len();
69 debug_assert!(n.is_power_of_two());
70 let mut step = 1usize;
71 while step < n {
72 let mut i = 0usize;
73 while i < n {
74 for j in i..i + step {
75 let a = buf[j];
76 let b = buf[j + step];
77 buf[j] = a + b;
78 buf[j + step] = a - b;
79 }
80 i += step * 2;
81 }
82 step *= 2;
83 }
84}
85
86fn sign_pack(rotated: &[f32], dim: usize) -> Vec<u8> {
91 let nbytes = dim.div_ceil(8);
92 let mut out = vec![0u8; nbytes];
93 for (i, &v) in rotated.iter().take(dim).enumerate() {
94 if v < 0.0 {
95 out[i / 8] |= 1 << (i % 8);
96 }
97 }
98 out
99}
100
101fn sign_unpack(packed: &[u8], dim: usize) -> Vec<f32> {
103 (0..dim)
104 .map(|i| {
105 if packed[i / 8] & (1 << (i % 8)) != 0 {
106 -1.0f32
107 } else {
108 1.0f32
109 }
110 })
111 .collect()
112}
113
114#[derive(Serialize, Deserialize, zerompk::ToMessagePack, zerompk::FromMessagePack)]
120pub struct RaBitQCodec {
121 pub dim: usize,
122 centroid: Vec<f32>,
124 rotation_seed: u64,
126 pub bias_correct: bool,
129}
130
131impl RaBitQCodec {
132 pub fn calibrate(vectors: &[&[f32]], dim: usize, rotation_seed: u64) -> Self {
141 let centroid = if vectors.is_empty() {
142 vec![0.0f32; dim]
143 } else {
144 let n = vectors.len() as f32;
145 let mut c = vec![0.0f32; dim];
146 for v in vectors {
147 for (ci, &vi) in c.iter_mut().zip(v.iter()) {
148 *ci += vi;
149 }
150 }
151 c.iter_mut().for_each(|x| *x /= n);
152 c
153 };
154 Self {
155 dim,
156 centroid,
157 rotation_seed,
158 bias_correct: false,
159 }
160 }
161
162 pub const ENVELOPE_MAGIC: &'static [u8; codec_envelope::MAGIC_LEN] = b"NDRBQ";
164
165 pub const ENVELOPE_VERSION: u8 = 1;
167
168 pub fn to_bytes(&self) -> Result<Vec<u8>, CodecError> {
170 codec_envelope::encode(Self::ENVELOPE_MAGIC, Self::ENVELOPE_VERSION, self)
171 }
172
173 pub fn from_bytes(buf: &[u8]) -> Result<Self, CodecError> {
175 codec_envelope::decode(Self::ENVELOPE_MAGIC, Self::ENVELOPE_VERSION, buf)
176 }
177
178 pub fn apply_rotation(&self, v: &[f32]) -> Vec<f32> {
186 let dim = self.dim;
187 let pow2 = next_pow2(dim);
188
189 let mut seed = self.rotation_seed;
191 let mut buf = vec![0.0f32; pow2];
192 for (i, &vi) in v.iter().take(dim).enumerate() {
193 let flip = if xorshift64(&mut seed) & 1 == 0 {
194 1.0f32
195 } else {
196 -1.0f32
197 };
198 buf[i] = vi * flip;
199 }
200 wht_inplace(&mut buf);
203 buf.truncate(dim);
204 buf
205 }
206
207 fn encode_inner(&self, v: &[f32]) -> UnifiedQuantizedVector {
218 let dim = self.dim;
219
220 let residual: Vec<f32> = v
222 .iter()
223 .zip(self.centroid.iter())
224 .map(|(&vi, &ci)| vi - ci)
225 .collect();
226
227 let residual_norm = residual.iter().map(|x| x * x).sum::<f32>().sqrt();
229
230 let rotated = self.apply_rotation(&residual);
232
233 let packed = sign_pack(&rotated, dim);
235
236 let signs_fp = sign_unpack(&packed, dim);
239 let pow2 = next_pow2(dim);
241 let mut sign_buf = vec![0.0f32; pow2];
242 for (i, &s) in signs_fp.iter().enumerate() {
243 sign_buf[i] = s;
244 }
245 wht_inplace(&mut sign_buf);
246 let mut seed = self.rotation_seed;
248 for x in sign_buf.iter_mut().take(dim) {
249 let flip = if xorshift64(&mut seed) & 1 == 0 {
250 1.0f32
251 } else {
252 -1.0f32
253 };
254 *x *= flip;
255 }
256 let dot_raw: f32 = residual
257 .iter()
258 .zip(sign_buf.iter().take(dim))
259 .map(|(&r, &s)| r * s)
260 .sum();
261 let dot_quantized = if residual_norm > 0.0 {
262 dot_raw / residual_norm
263 } else {
264 0.0
265 };
266
267 let header = QuantHeader {
268 quant_mode: QuantMode::RaBitQ as u16,
269 dim: dim as u16,
270 global_scale: residual_norm,
271 residual_norm,
272 dot_quantized,
273 outlier_bitmask: 0,
274 reserved: [0u8; 8],
275 };
276
277 UnifiedQuantizedVector::new(header, &packed, &[])
278 .expect("RaBitQ encode: layout construction must succeed")
279 }
280}
281
282pub struct RaBitQQuantized(UnifiedQuantizedVector);
286
287impl AsRef<UnifiedQuantizedVector> for RaBitQQuantized {
288 #[inline]
289 fn as_ref(&self) -> &UnifiedQuantizedVector {
290 &self.0
291 }
292}
293
294pub struct RaBitQQuery {
296 pub rotated_signs: Vec<u8>,
298 pub query_norm: f32,
300}
301
302impl VectorCodec for RaBitQCodec {
305 type Quantized = RaBitQQuantized;
306 type Query = RaBitQQuery;
307
308 fn encode(&self, v: &[f32]) -> Self::Quantized {
309 RaBitQQuantized(self.encode_inner(v))
310 }
311
312 fn prepare_query(&self, q: &[f32]) -> Self::Query {
313 let dim = self.dim;
314 let residual: Vec<f32> = q
315 .iter()
316 .zip(self.centroid.iter())
317 .map(|(&qi, &ci)| qi - ci)
318 .collect();
319 let query_norm = residual.iter().map(|x| x * x).sum::<f32>().sqrt();
320 let rotated = self.apply_rotation(&residual);
321 let rotated_signs = sign_pack(&rotated, dim);
322 RaBitQQuery {
323 rotated_signs,
324 query_norm,
325 }
326 }
327
328 fn fast_symmetric_distance(&self, q: &Self::Quantized, v: &Self::Quantized) -> f32 {
335 let qh = q.0.header();
336 let vh = v.0.header();
337 let qb = q.0.packed_bits();
338 let vb = v.0.packed_bits();
339 let h = hamming_distance(qb, vb);
340 let dim = self.dim as f32;
341 let dot_estimate = 1.0 - 2.0 * h as f32 / dim;
342 let approx = qh.residual_norm * qh.residual_norm + vh.residual_norm * vh.residual_norm
343 - 2.0 * qh.residual_norm * vh.residual_norm * dot_estimate;
344 approx.max(0.0)
345 }
346
347 fn exact_asymmetric_distance(&self, q: &Self::Query, v: &Self::Quantized) -> f32 {
356 let vh = v.0.header();
357 let vb = v.0.packed_bits();
358 let h = hamming_distance(&q.rotated_signs, vb);
359 let dim = self.dim as f32;
360 let dot_estimate = 1.0 - 2.0 * h as f32 / dim;
361 let mut approx = q.query_norm * q.query_norm + vh.residual_norm * vh.residual_norm
362 - 2.0 * q.query_norm * vh.residual_norm * dot_estimate;
363 if self.bias_correct {
364 approx -= vh.dot_quantized;
365 }
366 approx.max(0.0)
367 }
368}
369
370#[cfg(test)]
373mod tests {
374 use super::*;
375
376 fn random_vec(seed: u64, dim: usize) -> Vec<f32> {
377 let mut s = seed | 1;
378 (0..dim)
379 .map(|_| {
380 let v = xorshift64(&mut s);
381 (v as f32 / u64::MAX as f32) * 2.0 - 1.0
383 })
384 .collect()
385 }
386
387 #[test]
388 fn to_bytes_from_bytes_roundtrip() {
389 let dim = 64;
390 let vecs: Vec<Vec<f32>> = (0..4).map(|i| random_vec(i as u64, dim)).collect();
391 let refs: Vec<&[f32]> = vecs.iter().map(|v| v.as_slice()).collect();
392 let codec = RaBitQCodec::calibrate(&refs, dim, 0xABCD_1234_5678_EF01);
393 let bytes = codec.to_bytes().expect("to_bytes should succeed");
394 let restored = RaBitQCodec::from_bytes(&bytes).expect("from_bytes should succeed");
395 assert_eq!(restored.dim, codec.dim);
396 assert_eq!(restored.rotation_seed, codec.rotation_seed);
397 assert_eq!(restored.bias_correct, codec.bias_correct);
398 assert_eq!(restored.centroid.len(), codec.centroid.len());
399 for (a, b) in restored.centroid.iter().zip(codec.centroid.iter()) {
400 assert!((a - b).abs() < 1e-6, "centroid mismatch: {a} vs {b}");
401 }
402 }
403
404 #[test]
405 fn from_bytes_rejects_bad_magic() {
406 let mut bytes = b"WRONG".to_vec();
407 bytes.push(1);
408 bytes.extend_from_slice(&[0u8; 4]);
409 assert!(RaBitQCodec::from_bytes(&bytes).is_err());
410 }
411
412 #[test]
413 fn from_bytes_rejects_bad_version() {
414 let codec = RaBitQCodec::calibrate(&[], 4, 1);
415 let mut bytes = codec.to_bytes().unwrap();
416 bytes[5] = 99;
417 assert!(RaBitQCodec::from_bytes(&bytes).is_err());
418 }
419
420 #[test]
421 fn apply_rotation_different_seeds_differ() {
422 let dim = 64;
423 let v: Vec<f32> = (0..dim).map(|i| i as f32 * 0.1).collect();
424 let codec_a = RaBitQCodec::calibrate(&[], dim, 0xDEAD_BEEF_1234_5678);
425 let codec_b = RaBitQCodec::calibrate(&[], dim, 0xCAFE_BABE_0000_0001);
426 let rot_a = codec_a.apply_rotation(&v);
427 let rot_b = codec_b.apply_rotation(&v);
428 let differ = rot_a
430 .iter()
431 .zip(rot_b.iter())
432 .any(|(a, b)| (a - b).abs() > 1e-6);
433 assert!(differ, "different seeds must produce different rotations");
434 }
435
436 #[test]
437 fn encode_roundtrip_preserves_residual_norm() {
438 let dim = 128;
439 let vecs: Vec<Vec<f32>> = (0..16).map(|i| random_vec(i as u64, dim)).collect();
440 let refs: Vec<&[f32]> = vecs.iter().map(|v| v.as_slice()).collect();
441 let codec = RaBitQCodec::calibrate(&refs, dim, 42);
442 let v = random_vec(99, dim);
443 let q = codec.encode(&v);
444 let h = q.0.header();
445 assert!(h.residual_norm.is_finite() && h.residual_norm >= 0.0);
447 assert!((h.global_scale - h.residual_norm).abs() < 1e-6);
448 }
449
450 #[test]
451 fn distance_non_negative_finite() {
452 let dim = 64;
453 let vecs: Vec<Vec<f32>> = (0..8).map(|i| random_vec(i as u64, dim)).collect();
454 let refs: Vec<&[f32]> = vecs.iter().map(|v| v.as_slice()).collect();
455 let codec = RaBitQCodec::calibrate(&refs, dim, 7);
456 let v1 = codec.encode(&random_vec(100, dim));
457 let v2 = codec.encode(&random_vec(200, dim));
458 let sym = codec.fast_symmetric_distance(&v1, &v2);
459 assert!(sym.is_finite() && sym >= 0.0, "sym distance: {sym}");
460 let q = codec.prepare_query(&random_vec(300, dim));
461 let asym = codec.exact_asymmetric_distance(&q, &v2);
462 assert!(asym.is_finite() && asym >= 0.0, "asym distance: {asym}");
463 }
464
465 #[test]
466 fn calibrate_identical_vectors_zero_residual() {
467 let dim = 32;
468 let v: Vec<f32> = (0..dim).map(|i| i as f32).collect();
469 let refs = vec![v.as_slice(); 16];
470 let codec = RaBitQCodec::calibrate(&refs, dim, 1);
471 let q = codec.encode(&v);
473 assert!(
474 q.0.header().residual_norm < 1e-5,
475 "residual_norm should be ~0 for vector equal to centroid"
476 );
477 }
478}