1use rand::{rngs::StdRng, RngExt, SeedableRng};
39use serde::{Deserialize, Serialize};
40
41
42const DEFAULT_SEED: u64 = 0x5241_4249_5451_5121; #[derive(Debug, Clone, Serialize, Deserialize)]
47pub struct RaBitQuantizer {
48 dim: usize,
49 centroid: Vec<f32>,
51 rotation: Vec<f32>,
53}
54
55#[derive(Debug, Clone, Serialize, Deserialize)]
57pub struct RaBitCode {
58 pub bits: Vec<u8>,
60 pub dtc_sq: f32,
62 pub est_factor: f32,
64}
65
66pub struct PreparedQuery {
68 rq: Vec<f32>,
70 qn_sq: f32,
72}
73
74impl PreparedQuery {
75 pub fn rq(&self) -> &[f32] {
82 &self.rq
83 }
84
85 pub fn qn_sq(&self) -> f32 {
87 self.qn_sq
88 }
89}
90
91impl RaBitQuantizer {
92 pub fn fit(training_vectors: &[Vec<f32>]) -> Self {
100 Self::fit_with_seed(training_vectors, DEFAULT_SEED)
101 }
102
103 pub fn fit_with_seed(training_vectors: &[Vec<f32>], seed: u64) -> Self {
105 assert!(
106 !training_vectors.is_empty(),
107 "Need at least one training vector"
108 );
109 let dim = training_vectors[0].len();
110 assert!(dim > 0, "Dimension must be positive");
111
112 let mut centroid = vec![0.0f32; dim];
113 for v in training_vectors {
114 assert_eq!(v.len(), dim, "Inconsistent vector dimensions");
115 for (c, &x) in centroid.iter_mut().zip(v.iter()) {
116 *c += x;
117 }
118 }
119 let inv_n = 1.0 / training_vectors.len() as f32;
120 for c in &mut centroid {
121 *c *= inv_n;
122 }
123
124 let rotation = random_orthonormal(dim, seed);
125 Self {
126 dim,
127 centroid,
128 rotation,
129 }
130 }
131
132
133 pub fn dim(&self) -> usize {
135 self.dim
136 }
137
138 pub fn encode(&self, vector: &[f32]) -> RaBitCode {
140 debug_assert_eq!(vector.len(), self.dim);
141
142 let mut res = vec![0.0f32; self.dim];
144 let mut dtc_sq = 0.0f32;
145 for ((r, &v), &c) in res.iter_mut().zip(vector).zip(&self.centroid) {
146 *r = v - c;
147 dtc_sq += (v - c) * (v - c);
148 }
149
150 let ro = self.matvec(&res);
152
153 let mut bits = vec![0u8; self.bytes()];
155 let mut l1 = 0.0f32;
156 for (i, &x) in ro.iter().enumerate() {
157 l1 += x.abs();
158 if x >= 0.0 {
159 bits[i / 8] |= 1 << (i % 8);
160 }
161 }
162
163 let est_factor = if l1 > f32::EPSILON { dtc_sq / l1 } else { 0.0 };
165
166 RaBitCode {
167 bits,
168 dtc_sq,
169 est_factor,
170 }
171 }
172
173 pub fn prepare_query(&self, query: &[f32]) -> PreparedQuery {
175 debug_assert_eq!(query.len(), self.dim);
176 let mut res = vec![0.0f32; self.dim];
177 let mut qn_sq = 0.0f32;
178 for ((r, &q), &c) in res.iter_mut().zip(query).zip(&self.centroid) {
179 *r = q - c;
180 qn_sq += (q - c) * (q - c);
181 }
182 let rq = self.matvec(&res);
183 PreparedQuery { rq, qn_sq }
184 }
185
186 pub fn estimate_dist_sq(&self, query: &PreparedQuery, code: &RaBitCode) -> f32 {
190 let mut s = 0.0f32;
191 for (i, &rq) in query.rq.iter().enumerate() {
192 let bit = (code.bits[i / 8] >> (i % 8)) & 1;
193 if bit == 1 {
195 s += rq;
196 } else {
197 s -= rq;
198 }
199 }
200 let dsq = code.dtc_sq + query.qn_sq - 2.0 * code.est_factor * s;
201 dsq.max(0.0)
202 }
203
204 fn bytes(&self) -> usize {
206 self.dim.div_ceil(8)
207 }
208
209 fn matvec(&self, v: &[f32]) -> Vec<f32> {
211 let d = self.dim;
212 let mut out = vec![0.0f32; d];
213 for (r, o) in out.iter_mut().enumerate() {
214 let row = &self.rotation[r * d..(r + 1) * d];
215 *o = super::simd::dot_product_simd(row, v);
216 }
217 out
218 }
219
220 fn matvec_transpose(&self, v: &[f32]) -> Vec<f32> {
222 let d = self.dim;
223 let mut out = vec![0.0f32; d];
224 for (r, &vr) in v.iter().enumerate() {
225 let row = &self.rotation[r * d..(r + 1) * d];
226 for (o, &rc) in out.iter_mut().zip(row) {
227 *o += rc * vr;
228 }
229 }
230 out
231 }
232}
233
234impl RaBitQuantizer {
244 pub fn quantize(&self, vector: &[f32]) -> RaBitCode {
246 self.encode(vector)
247 }
248
249 pub fn dequantize(&self, quantized: &RaBitCode) -> Vec<f32> {
252 let d = self.dim;
255 let inv_sqrt_d = 1.0 / (d as f32).sqrt();
256 let mut xbar = vec![0.0f32; d];
257 for (i, x) in xbar.iter_mut().enumerate() {
258 let bit = (quantized.bits[i / 8] >> (i % 8)) & 1;
259 *x = if bit == 1 { inv_sqrt_d } else { -inv_sqrt_d };
260 }
261 let dir = self.matvec_transpose(&xbar);
262 let dtc = quantized.dtc_sq.sqrt();
263 dir.iter()
264 .zip(&self.centroid)
265 .map(|(&u, &c)| c + dtc * u)
266 .collect()
267 }
268
269 pub fn distance_quantized(&self, a: &RaBitCode, b: &RaBitCode) -> f32 {
272 let a_full = self.dequantize(a);
273 self.distance_asymmetric(&a_full, b)
274 }
275
276 pub fn distance_asymmetric(&self, query: &[f32], quantized: &RaBitCode) -> f32 {
279 let prepared = self.prepare_query(query);
280 self.estimate_dist_sq(&prepared, quantized).sqrt()
281 }
282}
283
284fn random_orthonormal(dim: usize, seed: u64) -> Vec<f32> {
287 let mut rng = StdRng::seed_from_u64(seed);
288 let mut rows: Vec<Vec<f32>> = Vec::with_capacity(dim);
289
290 for _ in 0..dim {
291 let mut v: Vec<f32> = (0..dim).map(|_| gaussian(&mut rng)).collect();
293
294 for prev in &rows {
296 let proj = dot(&v, prev);
297 for (vi, &pi) in v.iter_mut().zip(prev) {
298 *vi -= proj * pi;
299 }
300 }
301
302 let mut norm = dot(&v, &v).sqrt();
304 while norm < 1e-6 {
305 v = (0..dim).map(|_| gaussian(&mut rng)).collect();
306 for prev in &rows {
307 let proj = dot(&v, prev);
308 for (vi, &pi) in v.iter_mut().zip(prev) {
309 *vi -= proj * pi;
310 }
311 }
312 norm = dot(&v, &v).sqrt();
313 }
314 let inv = 1.0 / norm;
315 for vi in &mut v {
316 *vi *= inv;
317 }
318 rows.push(v);
319 }
320
321 let mut flat = Vec::with_capacity(dim * dim);
322 for row in rows {
323 flat.extend_from_slice(&row);
324 }
325 flat
326}
327
328#[inline]
329fn dot(a: &[f32], b: &[f32]) -> f32 {
330 a.iter().zip(b).map(|(&x, &y)| x * y).sum()
331}
332
333#[inline]
335fn gaussian(rng: &mut StdRng) -> f32 {
336 let u1: f32 = rng.random::<f32>().max(1e-7);
337 let u2: f32 = rng.random::<f32>();
338 (-2.0 * u1.ln()).sqrt() * (std::f32::consts::TAU * u2).cos()
339}
340
341#[cfg(test)]
342mod tests {
343 use super::*;
344
345 fn rng_vec(rng: &mut StdRng, dim: usize) -> Vec<f32> {
346 (0..dim).map(|_| rng.random::<f32>() * 2.0 - 1.0).collect()
347 }
348
349 #[test]
350 fn rotation_is_orthonormal() {
351 let d = 64;
352 let r = random_orthonormal(d, 123);
353 for i in 0..d {
355 for j in 0..d {
356 let ri = &r[i * d..(i + 1) * d];
357 let rj = &r[j * d..(j + 1) * d];
358 let prod = dot(ri, rj);
359 let expected = if i == j { 1.0 } else { 0.0 };
360 assert!(
361 (prod - expected).abs() < 1e-3,
362 "R·Rᵀ[{i},{j}] = {prod}, expected {expected}"
363 );
364 }
365 }
366 }
367
368 #[test]
369 fn rotation_is_deterministic() {
370 assert_eq!(random_orthonormal(32, 42), random_orthonormal(32, 42));
371 }
372
373 #[test]
374 fn estimator_is_approximately_unbiased() {
375 let mut rng = StdRng::seed_from_u64(7);
377 let dim = 128;
378 let train: Vec<Vec<f32>> = (0..500).map(|_| rng_vec(&mut rng, dim)).collect();
379 let q = RaBitQuantizer::fit(&train);
380
381 let mut rel_errs = Vec::new();
382 for _ in 0..200 {
383 let o = rng_vec(&mut rng, dim);
384 let query = rng_vec(&mut rng, dim);
385 let code = q.encode(&o);
386 let prep = q.prepare_query(&query);
387 let est = q.estimate_dist_sq(&prep, &code);
388 let truth: f32 = o.iter().zip(&query).map(|(a, b)| (a - b) * (a - b)).sum();
389 rel_errs.push((est - truth) / truth);
390 }
391 let mean_bias: f32 = rel_errs.iter().sum::<f32>() / rel_errs.len() as f32;
392 assert!(
394 mean_bias.abs() < 0.10,
395 "estimator mean relative bias too large: {mean_bias}"
396 );
397 }
398
399 #[test]
400 fn rerank_recall_beats_hamming_floor() {
401 let mut rng = StdRng::seed_from_u64(99);
405 let dim = 128;
406 let n = 2000;
407 let base: Vec<Vec<f32>> = (0..n).map(|_| rng_vec(&mut rng, dim)).collect();
408 let q = RaBitQuantizer::fit(&base);
409 let codes: Vec<RaBitCode> = base.iter().map(|v| q.encode(v)).collect();
410
411 let k = 10;
412 let rerank = 100; let mut total_recall = 0.0;
414 let trials = 50;
415 for _ in 0..trials {
416 let query = rng_vec(&mut rng, dim);
417
418 let mut exact: Vec<(f32, usize)> = base
420 .iter()
421 .enumerate()
422 .map(|(i, v)| {
423 (
424 v.iter().zip(&query).map(|(a, b)| (a - b) * (a - b)).sum(),
425 i,
426 )
427 })
428 .collect();
429 exact.sort_by(|a, b| a.0.total_cmp(&b.0));
430 let truth: std::collections::HashSet<usize> =
431 exact.iter().take(k).map(|(_, i)| *i).collect();
432
433 let prep = q.prepare_query(&query);
435 let mut est: Vec<(f32, usize)> = codes
436 .iter()
437 .enumerate()
438 .map(|(i, c)| (q.estimate_dist_sq(&prep, c), i))
439 .collect();
440 est.sort_by(|a, b| a.0.total_cmp(&b.0));
441
442 let mut pool: Vec<(f32, usize)> = est
444 .iter()
445 .take(rerank)
446 .map(|&(_, i)| {
447 let d: f32 = base[i]
448 .iter()
449 .zip(&query)
450 .map(|(a, b)| (a - b) * (a - b))
451 .sum();
452 (d, i)
453 })
454 .collect();
455 pool.sort_by(|a, b| a.0.total_cmp(&b.0));
456 let got: std::collections::HashSet<usize> =
457 pool.iter().take(k).map(|(_, i)| *i).collect();
458
459 total_recall += truth.intersection(&got).count() as f32 / k as f32;
460 }
461 let recall = total_recall / trials as f32;
462 assert!(recall > 0.80, "RaBitQ rerank recall@10 too low: {recall}");
463 }
464
465 #[test]
466 fn quantizer_trait_roundtrip() {
467 let mut rng = StdRng::seed_from_u64(5);
468 let dim = 96;
469 let train: Vec<Vec<f32>> = (0..200).map(|_| rng_vec(&mut rng, dim)).collect();
470 let q = RaBitQuantizer::fit(&train);
471
472 let v = rng_vec(&mut rng, dim);
473 let code = q.quantize(&v);
474 assert_eq!(code.bits.len(), dim.div_ceil(8));
475
476 let self_d = q.distance_asymmetric(&v, &code);
478 let other = rng_vec(&mut rng, dim);
479 let other_d = q.distance_asymmetric(&other, &code);
480 assert!(
481 self_d < other_d,
482 "self distance {self_d} should be < cross distance {other_d}"
483 );
484
485 assert_eq!(q.dequantize(&code).len(), dim);
487 }
488}