GCC Code Coverage Report


Directory: src/
File: src/quant/rabitq.c
Date: 2026-09-30 11:11:31
Exec Total Coverage
Lines: 492 607 81.1%
Functions: 46 52 88.5%
Branches: 230 393 58.5%

Line Branch Exec Source
1 /*
2 * Copyright (c) 2026 Tiger Data, Inc.
3 * Licensed under the PostgreSQL License. See LICENSE for details.
4 *
5 * rabitq.c - RaBitQ (Randomized Binary Quantization) implementation
6 *
7 * Core implementation with SIMD dispatch for the binary inner product
8 * operation. The hot path during search is computing:
9 *
10 * sum = 0
11 * for each bit i:
12 * if bits[i] == 1:
13 * sum += transformed_query[i]
14 *
15 * This is optimized with AVX2/AVX-512/NEON intrinsics.
16 */
17
18 #include "vs_config.h"
19
20 #include <math.h>
21 #include <stdatomic.h>
22 #include <string.h>
23
24 #include "algo/simd_utils.h"
25 #include "algo/vecops.h"
26 #include "core/memory.h"
27 #include "core/platform.h"
28 #include "quant/matrix.h"
29 #include "quant/rabitq.h"
30 #include "types/vec16.h"
31 #include "types/vec32.h"
32
33 /*
34 * Compiler-Vectorized Implementation
35 *
36 * Uses target_clones to generate multiple versions for different ISAs.
37 * The dynamic linker selects the best version at load time.
38 *
39 * Structured for auto-vectorization: processes 8 floats per byte with
40 * explicit mask expansion that compilers can optimize.
41 */
42
43 VS_TARGET_CLONES static float
44 324 rabitq_inner_product_compiler(
45 const float *transformed, const uint8_t *bits, Dimension dim)
46 {
47 324 float sum = 0.0f;
48
49 /* Main loop: process 8 floats (1 byte) at a time.
50 * This structure helps auto-vectorization by:
51 * 1. Fixed iteration count inner loop (8 iterations)
52 * 2. Simple bit test with mask array lookup
53 * 3. Multiply instead of conditional for branchless code
54 */
55 324 Dimension i = 0;
56
2/2
✓ Branch 0 taken 9870 times.
✓ Branch 1 taken 324 times.
10194 for (; i + 8 <= dim; i += 8)
57 {
58 9870 uint8_t byte = bits[i / 8];
59
60 /* Unrolled: multiply by bit value (0 or 1) */
61 9870 sum += transformed[i + 0] * ((byte >> 0) & 1);
62 9870 sum += transformed[i + 1] * ((byte >> 1) & 1);
63 9870 sum += transformed[i + 2] * ((byte >> 2) & 1);
64 9870 sum += transformed[i + 3] * ((byte >> 3) & 1);
65 9870 sum += transformed[i + 4] * ((byte >> 4) & 1);
66 9870 sum += transformed[i + 5] * ((byte >> 5) & 1);
67 9870 sum += transformed[i + 6] * ((byte >> 6) & 1);
68 9870 sum += transformed[i + 7] * ((byte >> 7) & 1);
69 }
70
71 /* Handle tail elements */
72
2/2
✓ Branch 0 taken 18 times.
✓ Branch 1 taken 324 times.
342 for (; i < dim; i++)
73 {
74 18 int byte_idx = i / 8;
75 18 int bit_idx = i % 8;
76 18 int bit = (bits[byte_idx] >> bit_idx) & 1;
77 18 sum += transformed[i] * bit;
78 }
79
80 324 return sum;
81 }
82
83 /*
84 * Compiler-vectorized multi-candidate inner product baseline.
85 *
86 * Simple loop calling the single-candidate function. This is the
87 * baseline that hand-optimized vertical SIMD kernels must beat.
88 */
89 VS_TARGET_CLONES static void
90 16 rabitq_inner_product_multi_compiler(
91 const float *transformed,
92 const uint8_t *bits,
93 uint32_t stride,
94 Dimension dim,
95 uint32_t count,
96 float *results)
97 {
98
2/2
✓ Branch 0 taken 132 times.
✓ Branch 1 taken 16 times.
148 for (uint32_t i = 0; i < count; i++)
99 132 results[i] = rabitq_inner_product_compiler(
100 132 transformed, bits + (size_t)i * stride, dim);
101 16 }
102
103 /*
104 * Function pointer dispatch for inner product and sign extraction
105 */
106 typedef float (*InnerProductFn)(const float *, const uint8_t *, Dimension);
107 typedef void (*ExtractSignsFn)(const float *, uint8_t *, Dimension);
108 static InnerProductFn g_inner_product_fn = NULL;
109 static InnerProductMultiFn g_inner_product_multi_fn = NULL;
110 static ExtractSignsFn g_extract_signs_fn = NULL;
111 static const char *g_impl_name = NULL;
112 static _Atomic(bool) g_rabitq_initialized = false;
113
114 /* Hamming distance function pointer dispatch */
115 typedef uint32_t (*HammingFn)(const uint8_t *, const uint8_t *, uint32_t);
116 typedef void (*HammingMultiFn)(
117 const uint8_t *,
118 const uint8_t *,
119 uint32_t,
120 uint32_t,
121 uint32_t,
122 uint32_t *);
123 static HammingFn g_hamming_fn = NULL;
124 static HammingMultiFn g_hamming_multi_fn = NULL;
125 static const char *g_hamming_impl_name = NULL;
126
127 /* Forward declaration for compiler-vectorized fallback */
128 VS_TARGET_CLONES static void rabitq_extract_signs_compiler(
129 const float *transformed, uint8_t *bits, Dimension dim);
130
131 /* Forward declarations for compiler-vectorized hamming */
132 VS_TARGET_CLONES static uint32_t rabitq_hamming_compiler(
133 const uint8_t *a, const uint8_t *b, uint32_t packed_bytes);
134 VS_TARGET_CLONES static void rabitq_hamming_multi_compiler(
135 const uint8_t *query_bits,
136 const uint8_t *data_bits,
137 uint32_t stride,
138 uint32_t packed_bytes,
139 uint32_t count,
140 uint32_t *results);
141
142 void
143 220 vs_rabitq_force_reinit(void)
144 {
145 220 g_rabitq_initialized = false;
146 220 g_inner_product_fn = NULL;
147 220 g_inner_product_multi_fn = NULL;
148 220 g_extract_signs_fn = NULL;
149 220 g_impl_name = NULL;
150 220 g_hamming_fn = NULL;
151 220 g_hamming_multi_fn = NULL;
152 220 g_hamming_impl_name = NULL;
153 220 }
154
155 int
156 485 vs_rabitq_init_simd(void)
157 {
158
2/2
✓ Branch 0 taken 257 times.
✓ Branch 1 taken 228 times.
485 if (g_rabitq_initialized)
159 2 return 0;
160
161 #ifdef VS_SIMD_FULL
162 483 SimdCapability caps = vs_detect_simd();
163
164 #if defined(__x86_64__) || defined(_M_X64)
165
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 483 times.
483 if ((caps & VS_SIMD_AVX512_DQ) == VS_SIMD_AVX512_DQ)
166 {
167 ✗ g_inner_product_fn = vs_rabitq_inner_product_avx512;
168 ✗ g_inner_product_multi_fn = vs_rabitq_inner_product_multi_avx512;
169 ✗ g_extract_signs_fn = vs_rabitq_extract_signs_avx512;
170 ✗ g_impl_name = "avx512";
171 }
172
2/2
✓ Branch 0 taken 339 times.
✓ Branch 1 taken 144 times.
483 else if (caps & SIMD_AVX2)
173 {
174 339 g_inner_product_fn = vs_rabitq_inner_product_avx2;
175 339 g_inner_product_multi_fn = vs_rabitq_inner_product_multi_avx2;
176 339 g_extract_signs_fn = vs_rabitq_extract_signs_avx2;
177 339 g_impl_name = "avx2";
178 }
179 else
180 {
181 144 g_inner_product_fn = rabitq_inner_product_compiler;
182 144 g_inner_product_multi_fn = rabitq_inner_product_multi_compiler;
183 144 g_extract_signs_fn = rabitq_extract_signs_compiler;
184 144 g_impl_name = "compiler";
185 }
186 #elif defined(__aarch64__) || defined(_M_ARM64)
187 if (caps & SIMD_NEON)
188 {
189 g_inner_product_fn = vs_rabitq_inner_product_neon;
190 g_inner_product_multi_fn = vs_rabitq_inner_product_multi_neon;
191 g_extract_signs_fn = vs_rabitq_extract_signs_neon;
192 g_impl_name = "neon";
193 }
194 else
195 {
196 g_inner_product_fn = rabitq_inner_product_compiler;
197 g_inner_product_multi_fn = rabitq_inner_product_multi_compiler;
198 g_extract_signs_fn = rabitq_extract_signs_compiler;
199 g_impl_name = "compiler";
200 }
201 #else
202 /* No hand-optimized kernels for this architecture. */
203 (void)caps;
204 g_inner_product_fn = rabitq_inner_product_compiler;
205 g_inner_product_multi_fn = rabitq_inner_product_multi_compiler;
206 g_extract_signs_fn = rabitq_extract_signs_compiler;
207 g_impl_name = "compiler";
208 #endif
209
210 /* Hamming dispatch (VPOPCNTDQ > AVX2 > compiler) */
211 #if defined(__x86_64__) || defined(_M_X64)
212
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 483 times.
483 if (caps & SIMD_AVX512_VPOPCNTDQ)
213 {
214 ✗ g_hamming_fn = vs_rabitq_hamming_avx512;
215 ✗ g_hamming_multi_fn = vs_rabitq_hamming_multi_avx512;
216 ✗ g_hamming_impl_name = "avx512-vpopcntdq";
217 }
218
2/2
✓ Branch 0 taken 339 times.
✓ Branch 1 taken 144 times.
483 else if (caps & SIMD_AVX2)
219 {
220 339 g_hamming_fn = vs_rabitq_hamming_avx2;
221 339 g_hamming_multi_fn = vs_rabitq_hamming_multi_avx2;
222 339 g_hamming_impl_name = "avx2";
223 }
224 else
225 {
226 144 g_hamming_fn = rabitq_hamming_compiler;
227 144 g_hamming_multi_fn = rabitq_hamming_multi_compiler;
228 144 g_hamming_impl_name = "compiler";
229 }
230 #elif defined(__aarch64__) || defined(_M_ARM64)
231 g_hamming_fn = rabitq_hamming_compiler;
232 g_hamming_multi_fn = rabitq_hamming_multi_compiler;
233 g_hamming_impl_name = "compiler";
234 #else
235 g_hamming_fn = rabitq_hamming_compiler;
236 g_hamming_multi_fn = rabitq_hamming_multi_compiler;
237 g_hamming_impl_name = "compiler";
238 #endif
239
240 #else
241 /* simd=compiler or simd=none mode */
242 g_inner_product_fn = rabitq_inner_product_compiler;
243 g_inner_product_multi_fn = rabitq_inner_product_multi_compiler;
244 g_extract_signs_fn = rabitq_extract_signs_compiler;
245 g_impl_name = "compiler";
246 g_hamming_fn = rabitq_hamming_compiler;
247 g_hamming_multi_fn = rabitq_hamming_multi_compiler;
248 g_hamming_impl_name = "compiler";
249 #endif
250
251 483 g_rabitq_initialized = true;
252 483 return 0;
253 }
254
255 const char *
256 2 vs_rabitq_impl_name(void)
257 {
258
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 2 times.
2 if (vs_unlikely(!g_rabitq_initialized))
259 ✗ vs_rabitq_init_simd();
260 2 return g_impl_name;
261 }
262
263 const char *
264 2 vs_rabitq_hamming_impl_name(void)
265 {
266
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 2 times.
2 if (vs_unlikely(!g_rabitq_initialized))
267 ✗ vs_rabitq_init_simd();
268 2 return g_hamming_impl_name;
269 }
270
271 /*
272 * Internal helper: compute inner product using dispatched function
273 */
274 static inline float
275 7138 rabitq_inner_product(
276 const float *transformed, const uint8_t *bits, Dimension dim)
277 {
278
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 7138 times.
7138 if (vs_unlikely(!g_rabitq_initialized))
279 ✗ vs_rabitq_init_simd();
280 7138 return g_inner_product_fn(transformed, bits, dim);
281 }
282
283 /*
284 * Compiler-Vectorized Sign Extraction
285 *
286 * Extracts sign bits from transformed floats. Structured for
287 * auto-vectorization by processing 8 floats at a time with explicit comparison
288 * and bit packing.
289 */
290 VS_TARGET_CLONES static void
291 196 rabitq_extract_signs_compiler(
292 const float *transformed, uint8_t *bits, Dimension dim)
293 {
294 196 Dimension packed_bytes = (dim + 7) / 8;
295 196 memset(bits, 0, packed_bytes);
296
297 /* Main loop: process 8 floats -> 1 byte at a time */
298 196 Dimension i = 0;
299
2/2
✓ Branch 0 taken 5914 times.
✓ Branch 1 taken 196 times.
6110 for (; i + 8 <= dim; i += 8)
300 {
301 5914 uint8_t byte = 0;
302 /* Explicit bit tests help vectorization */
303 5914 byte |= (transformed[i + 0] > 0) ? (1 << 0) : 0;
304
2/2
✓ Branch 0 taken 2922 times.
✓ Branch 1 taken 2992 times.
5914 byte |= (transformed[i + 1] > 0) ? (1 << 1) : 0;
305
2/2
✓ Branch 0 taken 2840 times.
✓ Branch 1 taken 3074 times.
5914 byte |= (transformed[i + 2] > 0) ? (1 << 2) : 0;
306
2/2
✓ Branch 0 taken 2750 times.
✓ Branch 1 taken 3164 times.
5914 byte |= (transformed[i + 3] > 0) ? (1 << 3) : 0;
307
2/2
✓ Branch 0 taken 2744 times.
✓ Branch 1 taken 3170 times.
5914 byte |= (transformed[i + 4] > 0) ? (1 << 4) : 0;
308
2/2
✓ Branch 0 taken 2782 times.
✓ Branch 1 taken 3132 times.
5914 byte |= (transformed[i + 5] > 0) ? (1 << 5) : 0;
309
2/2
✓ Branch 0 taken 2824 times.
✓ Branch 1 taken 3090 times.
5914 byte |= (transformed[i + 6] > 0) ? (1 << 6) : 0;
310
2/2
✓ Branch 0 taken 2528 times.
✓ Branch 1 taken 3386 times.
5914 byte |= (transformed[i + 7] > 0) ? (1 << 7) : 0;
311 5914 bits[i / 8] = byte;
312 }
313
314 /* Handle tail elements */
315
2/2
✓ Branch 0 taken 24 times.
✓ Branch 1 taken 196 times.
220 for (; i < dim; i++)
316 {
317 24 int byte_idx = i / 8;
318 24 int bit_idx = i % 8;
319
2/2
✓ Branch 0 taken 6 times.
✓ Branch 1 taken 18 times.
24 if (transformed[i] > 0)
320 6 bits[byte_idx] |= (1 << bit_idx);
321 }
322 196 }
323
324 /*
325 * Internal helper: extract sign bits using dispatched function
326 */
327 static inline void
328 685332 rabitq_extract_signs(const float *transformed, uint8_t *bits, Dimension dim)
329 {
330
2/2
✓ Branch 0 taken 8 times.
✓ Branch 1 taken 685324 times.
685332 if (vs_unlikely(!g_rabitq_initialized))
331 8 vs_rabitq_init_simd();
332 685332 g_extract_signs_fn(transformed, bits, dim);
333 685332 }
334
335 /*
336 * Error bound helpers
337 *
338 * Used by all with_bound distance functions (scalar and batch, asymmetric
339 * and symmetric). Centralizes the f_error derivation and lower_bound
340 * computation that were previously duplicated across 4 functions.
341 */
342
343 static inline float
344 681308 rabitq_derive_f_error(
345 float f_add, float f_rescale, float c_error, Dimension dim)
346 {
347 681308 float f_rsq = f_rescale * f_rescale;
348 681308 float ratio = f_rsq / f_add;
349
3/4
✓ Branch 0 taken 279942 times.
✓ Branch 1 taken 401366 times.
✗ Branch 2 not taken.
✓ Branch 3 taken 245420 times.
681308 if (ratio <= 1.0f || dim <= 1)
350 34522 return 2e-4f * sqrtf(f_add);
351 646786 return c_error * sqrtf(f_rsq - f_add);
352 }
353
354 float
355 674548 vs_rabitq_derive_f_error(float f_add, float f_rescale, Dimension dim)
356 {
357 1110436 float c_error = (dim > 1)
358 674548 ? 2.0f * VS_RABITQ_EPSILON / sqrtf((float)(dim - 1))
359
1/2
✓ Branch 0 taken 674548 times.
✗ Branch 1 not taken.
674548 : 0.0f;
360 674548 return rabitq_derive_f_error(f_add, f_rescale, c_error, dim);
361 }
362
363 static inline Distance
364 6760 rabitq_lower_bound(
365 Distance est_dist, float f_error, float g_error, float multiplier)
366 {
367 6760 float err_margin = multiplier * f_error * g_error;
368 6760 float fp_margin = 1e-5f * fabsf(est_dist);
369 6760 return est_dist - err_margin - fp_margin;
370 }
371
372 /*
373 * Compiler-Vectorized Hamming Distance
374 *
375 * Computes XOR + popcount between two packed bit vectors.
376 * Uses target_clones to generate multiple versions for different ISAs.
377 */
378 VS_TARGET_CLONES static uint32_t
379 60 rabitq_hamming_compiler(
380 const uint8_t *a, const uint8_t *b, uint32_t packed_bytes)
381 {
382 60 uint32_t count = 0;
383
384 /* Main loop: process 8 bytes (64 bits) at a time */
385 60 uint32_t i = 0;
386
2/2
✓ Branch 0 taken 408 times.
✓ Branch 1 taken 60 times.
468 for (; i + 8 <= packed_bytes; i += 8)
387 {
388 ✗ uint64_t va, vb;
389 408 memcpy(&va, a + i, 8);
390 408 memcpy(&vb, b + i, 8);
391 408 count += (uint32_t)__builtin_popcountll(va ^ vb);
392 }
393
394 /* Byte tail */
395
2/2
✓ Branch 0 taken 54 times.
✓ Branch 1 taken 60 times.
114 for (; i < packed_bytes; i++)
396 54 count += (uint32_t)__builtin_popcount(a[i] ^ b[i]);
397
398 60 return count;
399 }
400
401 /*
402 * Compiler-Vectorized Multi-Candidate Hamming Distance
403 */
404 VS_TARGET_CLONES static void
405 ✗ rabitq_hamming_multi_compiler(
406 const uint8_t *query_bits,
407 const uint8_t *data_bits,
408 uint32_t stride,
409 uint32_t packed_bytes,
410 uint32_t count,
411 uint32_t *results)
412 {
413 ✗ for (uint32_t c = 0; c < count; c++)
414 {
415 ✗ results[c] = rabitq_hamming_compiler(
416 ✗ query_bits, data_bits + c * stride, packed_bytes);
417 }
418 ✗ }
419
420 /*
421 * Lifecycle functions
422 */
423
424 /*
425 * Apply the index rotation P^T*in -> out. Uses the O(d log d) Randomized
426 * Hadamard Transform when armed (supported dims), else the dense matrix.
427 */
428 static inline void
429 763778 rabitq_rotate(const RaBitQParams *p, const float *in, float *out)
430 {
431
2/2
✓ Branch 0 taken 701166 times.
✓ Branch 1 taken 62612 times.
763778 if (p->use_fast_rotate)
432 701166 vs_fast_rotate_apply(&p->fr, in, out);
433 else
434 62612 vs_matrix_transpose_vector_mul(p->P, in, out, p->dim);
435 763778 }
436
437 RaBitQParams *
438 783 vs_rabitq_create(Dimension dim, uint64_t seed)
439 {
440 783 size_t size = VS_RABITQ_PARAMS_SIZE(dim);
441 783 RaBitQParams *params = vs_alloc(size);
442
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 783 times.
783 if (params == NULL)
443 ✗ return NULL;
444
445
1/2
✗ Branch 1 not taken.
✓ Branch 2 taken 783 times.
783 if (vs_rabitq_init(params, dim, seed) != 0)
446 {
447 ✗ vs_free(params);
448 ✗ return NULL;
449 }
450
451 504 return params;
452 }
453
454 RaBitQParams *
455 12 vs_rabitq_create_from_matrix(Dimension dim, uint64_t seed, const float *P)
456 {
457 /* Stores an explicit dense matrix, so always allocate the full layout. */
458 12 size_t size = VS_RABITQ_PARAMS_DENSE_SIZE(dim);
459 12 RaBitQParams *params = vs_alloc(size);
460
2/2
✓ Branch 0 taken 4 times.
✓ Branch 1 taken 8 times.
12 if (params == NULL)
461 ✗ return NULL;
462
463 12 params->dim = dim;
464 12 params->seed = seed;
465 12 params->packed_bytes = VS_RABITQ_BYTES(dim);
466 12 params->use_fast_rotate = false;
467 12 memcpy(params->P, P, (size_t)dim * dim * sizeof(float));
468
469 12 return params;
470 }
471
472 int
473 882 vs_rabitq_init(RaBitQParams *params, Dimension dim, uint64_t seed)
474 {
475
4/4
✓ Branch 0 taken 506 times.
✓ Branch 1 taken 376 times.
✓ Branch 2 taken 2 times.
✓ Branch 3 taken 504 times.
882 if (params == NULL || dim == 0)
476 4 return -1;
477
478 878 params->dim = dim;
479 878 params->seed = seed;
480 878 params->packed_bytes = VS_RABITQ_BYTES(dim);
481
482 /*
483 * Prefer the O(d log d) Randomized Hadamard rotation where the dim
484 * supports it: the dense P is neither built (no O(d^3) QR) nor stored
485 * (no O(d^2) per-backend matrix). Same seed, so build-encode and
486 * query-rotate stay consistent. Unsupported dims fall back to the
487 * dense random orthogonal matrix in P[].
488 */
489 878 params->use_fast_rotate = vs_fast_rotate_supported(dim);
490
2/2
✓ Branch 0 taken 658 times.
✓ Branch 1 taken 220 times.
878 if (params->use_fast_rotate)
491 658 vs_fast_rotate_init(&params->fr, dim, seed);
492
1/2
✗ Branch 1 not taken.
✓ Branch 2 taken 220 times.
220 else if (vs_random_orthogonal_matrix(params->P, dim, seed) != 0)
493 ✗ return -1;
494
495 504 return 0;
496 }
497
498 void
499 4 vs_rabitq_cleanup(RaBitQParams *params)
500 {
501 ✗ (void)params;
502 /* P is now inline — nothing to free */
503 4 }
504
505 void
506 230 vs_rabitq_destroy(RaBitQParams *params)
507 {
508
2/2
✓ Branch 0 taken 2 times.
✓ Branch 1 taken 228 times.
230 if (params == NULL)
509 2 return;
510
511 228 vs_free(params);
512 }
513
514 /*
515 * Encoding functions
516 */
517
518 RaBitQData *
519 638 vs_rabitq_encode(const RaBitQParams *params, Vec32Ref input, Vec32Ref centroid)
520 {
521
6/6
✓ Branch 0 taken 636 times.
✓ Branch 1 taken 2 times.
✓ Branch 2 taken 634 times.
✓ Branch 3 taken 2 times.
✓ Branch 4 taken 2 times.
✓ Branch 5 taken 632 times.
638 if (params == NULL || input.data == NULL || centroid.data == NULL)
522 6 return NULL;
523
524
3/4
✓ Branch 0 taken 632 times.
✗ Branch 1 not taken.
✓ Branch 2 taken 2 times.
✓ Branch 3 taken 630 times.
632 if (input.dim != params->dim || centroid.dim != params->dim)
525 2 return NULL;
526
527 630 size_t size = VS_RABITQ_DATA_SIZE(params->dim);
528 630 RaBitQData *output = vs_alloc(size);
529
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 630 times.
630 if (output == NULL)
530 ✗ return NULL;
531
532
1/2
✗ Branch 1 not taken.
✓ Branch 2 taken 630 times.
630 if (vs_rabitq_encode_into(params, input, centroid, output) != 0)
533 {
534 ✗ vs_free(output);
535 ✗ return NULL;
536 }
537
538 630 return output;
539 }
540
541 void
542 29595 vs_rabitq_scratch_init(RaBitQScratch *scratch, Dimension dim)
543 {
544 29595 scratch->residual = vs_alloc_aligned(dim * sizeof(float), 64);
545 29595 scratch->transformed = vs_alloc_aligned(dim * sizeof(float), 64);
546 29595 scratch->xu_cb = vs_alloc_aligned(dim * sizeof(float), 64);
547 29595 }
548
549 void
550 12132 vs_rabitq_scratch_cleanup(RaBitQScratch *scratch)
551 {
552 12132 vs_free_aligned(scratch->residual);
553 12132 vs_free_aligned(scratch->transformed);
554 12132 vs_free_aligned(scratch->xu_cb);
555 12132 scratch->residual = NULL;
556 12132 scratch->transformed = NULL;
557 12132 scratch->xu_cb = NULL;
558 12132 }
559
560 int
561 71546 vs_rabitq_encode_into_ex(
562 const RaBitQParams *params,
563 Vec32Ref input,
564 Vec32Ref centroid,
565 RaBitQData *output,
566 RaBitQScratch *scratch)
567 {
568
4/8
✓ Branch 0 taken 71546 times.
✗ Branch 1 not taken.
✓ Branch 2 taken 71546 times.
✗ Branch 3 not taken.
✓ Branch 4 taken 71546 times.
✗ Branch 5 not taken.
✓ Branch 6 taken 56846 times.
✗ Branch 7 not taken.
71546 if (params == NULL || input.data == NULL || centroid.data == NULL ||
569
2/2
✓ Branch 0 taken 14700 times.
✓ Branch 1 taken 56846 times.
71546 output == NULL || scratch == NULL)
570 ✗ return -1;
571
572
3/4
✓ Branch 0 taken 71546 times.
✗ Branch 1 not taken.
✓ Branch 2 taken 14700 times.
✓ Branch 3 taken 56846 times.
71546 if (input.dim != params->dim || centroid.dim != params->dim)
573 ✗ return -1;
574
575 71546 Dimension dim = params->dim;
576
577 71546 float *residual = scratch->residual;
578 71546 float *transformed = scratch->transformed;
579
580 /* residual = input - centroid; transformed = P^T * residual. The signs +
581 * factor math is shared with encode_from_pt (which the insert path calls
582 * directly with a pre-rotated residual). */
583 71546 vec32_sub(input.data, centroid.data, residual, dim);
584 71546 rabitq_rotate(params, residual, transformed);
585
586 71546 return vs_rabitq_encode_from_pt(params, transformed, output, scratch);
587 }
588
589 int
590 654286 vs_rabitq_encode_from_pt(
591 const RaBitQParams *params,
592 const float *pt_residual,
593 RaBitQData *output,
594 RaBitQScratch *scratch)
595 {
596
4/8
✓ Branch 0 taken 654286 times.
✗ Branch 1 not taken.
✓ Branch 2 taken 227512 times.
✗ Branch 3 not taken.
✓ Branch 4 taken 227512 times.
✗ Branch 5 not taken.
✗ Branch 6 not taken.
✓ Branch 7 taken 227512 times.
654286 if (params == NULL || pt_residual == NULL || output == NULL ||
597
1/2
✓ Branch 0 taken 426774 times.
✗ Branch 1 not taken.
426774 scratch == NULL)
598 ✗ return -1;
599
600 654286 Dimension dim = params->dim;
601 654286 float *xu_cb = scratch->xu_cb;
602
603 /* pt_residual is the rotated residual P^T*(input-centroid); everything
604 * below operates on it exactly as encode_into_ex did on `transformed`. */
605 654286 rabitq_extract_signs(pt_residual, output->bits, dim);
606
607 654286 float cb = -0.5f;
608
3/3
✓ Branch 0 taken 7040352 times.
✓ Branch 1 taken 31115390 times.
✓ Branch 2 taken 426774 times.
38582516 for (Dimension i = 0; i < dim; i++)
609 {
610 37928230 int byte_idx = i / 8;
611 37928230 int bit_idx = i % 8;
612 37928230 int bit = (output->bits[byte_idx] >> bit_idx) & 1;
613 37928230 xu_cb[i] = (float)bit + cb;
614 }
615
616 654286 float l2_sqr = vs_l2_norm_squared(pt_residual, dim);
617 654286 float ip_resi_xucb = vs_dot_product(pt_residual, xu_cb, dim);
618
619
2/2
✓ Branch 0 taken 1609 times.
✓ Branch 1 taken 652677 times.
654286 if (fabsf(ip_resi_xucb) < 1e-10f)
620 1609 ip_resi_xucb = 1e-10f;
621
622 654286 float sqrt_d = sqrtf((float)dim);
623 654286 float l1_norm = 2.0f * fabsf(ip_resi_xucb);
624 654286 output->f_add = l2_sqr;
625 654286 output->f_rescale = l2_sqr * sqrt_d / l1_norm;
626
627 654286 return 0;
628 }
629
630 int
631 24018 vs_rabitq_encode_into(
632 const RaBitQParams *params,
633 Vec32Ref input,
634 Vec32Ref centroid,
635 RaBitQData *output)
636 {
637
5/8
✓ Branch 0 taken 24018 times.
✗ Branch 1 not taken.
✓ Branch 2 taken 24018 times.
✗ Branch 3 not taken.
✓ Branch 4 taken 24018 times.
✗ Branch 5 not taken.
✓ Branch 6 taken 2 times.
✓ Branch 7 taken 24016 times.
24018 if (params == NULL || input.data == NULL || centroid.data == NULL ||
638 output == NULL)
639 2 return -1;
640
641
4/4
✓ Branch 0 taken 13788 times.
✓ Branch 1 taken 10228 times.
✓ Branch 2 taken 2 times.
✓ Branch 3 taken 24012 times.
24016 if (input.dim != params->dim || centroid.dim != params->dim)
642 4 return -1;
643
644 24012 Dimension dim = params->dim;
645
646 /* Allocate temporary buffers */
647 24012 float *residual = vs_alloc_aligned(dim * sizeof(float), 64);
648 24012 float *transformed = vs_alloc_aligned(dim * sizeof(float), 64);
649 24012 float *xu_cb = vs_alloc_aligned(dim * sizeof(float), 64);
650
651
4/6
✓ Branch 0 taken 24012 times.
✗ Branch 1 not taken.
✓ Branch 2 taken 13786 times.
✓ Branch 3 taken 10226 times.
✗ Branch 4 not taken.
✓ Branch 5 taken 13786 times.
24012 if (residual == NULL || transformed == NULL || xu_cb == NULL)
652 {
653 ✗ if (residual)
654 ✗ vs_free_aligned(residual);
655 ✗ if (transformed)
656 ✗ vs_free_aligned(transformed);
657 ✗ if (xu_cb)
658 ✗ vs_free_aligned(xu_cb);
659 ✗ return -1;
660 }
661
662 /* Step 1: Compute residual = input - centroid */
663 24012 vec32_sub(input.data, centroid.data, residual, dim);
664
665 /* Step 2: Transform residual through P^T
666 * RaBitQ applies the same rotation to all vectors (data, centroid, query).
667 * The library stores Q^T and multiplies by it, which is effectively P^T.
668 * Since P^T * (data - centroid) = P^T * data - P^T * centroid, applying
669 * P^T to the residual is equivalent to rotating both vectors then
670 * subtracting.
671 */
672 24012 rabitq_rotate(params, residual, transformed);
673
674 /* Step 3: Extract sign bits (LSB-first packing, FAISS-compatible) */
675 24012 rabitq_extract_signs(transformed, output->bits, dim);
676
677 /* Step 4: Compute xu_cb = binary_code - 0.5 (for 1-bit quantization) */
678 /* In 1-bit RaBitQ, cb = -0.5, so xu_cb[i] = bit[i] - 0.5 */
679 24012 float cb = -0.5f;
680
3/3
✓ Branch 0 taken 457582 times.
✓ Branch 1 taken 778672 times.
✓ Branch 2 taken 10226 times.
1246480 for (Dimension i = 0; i < dim; i++)
681 {
682 1222468 int byte_idx = i / 8;
683 1222468 int bit_idx = i % 8; /* LSB-first */
684 1222468 int bit = (output->bits[byte_idx] >> bit_idx) & 1;
685 1222468 xu_cb[i] = (float)bit + cb;
686 }
687
688 /* Step 5: Compute factors for distance estimation */
689 24012 float l2_sqr = vs_l2_norm_squared(transformed, dim);
690
691 24012 float ip_resi_xucb = vs_dot_product(transformed, xu_cb, dim);
692
693 /* Handle corner case: avoid division by near-zero in f_rescale */
694
2/2
✓ Branch 0 taken 120 times.
✓ Branch 1 taken 23892 times.
24012 if (fabsf(ip_resi_xucb) < 1e-10f)
695 120 ip_resi_xucb = 1e-10f;
696
697 /* L2 distance factors (FAISS-style formula for better accuracy)
698 *
699 * FAISS uses a correction factor based on the actual distribution of
700 * values: est_dist = ||v-c||² + ||q-c||² - 2 * dp_multiplier * final_dot
701 * where:
702 * dp_multiplier = ||v-c||² * sqrt(d) / ||v-c||_1
703 *
704 * The L1 norm is encoded in ip_resi_xucb (which equals 0.5 * ||v-c||_1
705 * because xu_cb = bit - 0.5, so the dot product accumulates ±0.5 * |val|).
706 *
707 * f_add = ||v-c||² (the L2 squared distance to centroid)
708 * f_rescale = dp_multiplier = ||v-c||² * sqrt(d) / ||v-c||_1
709 *
710 * f_error is derived at query time from f_add and f_rescale.
711 */
712 24012 float sqrt_d = sqrtf((float)dim);
713 24012 float l1_norm = 2.0f * fabsf(ip_resi_xucb); /* ||v-c||_1 */
714 24012 output->f_add = l2_sqr;
715 24012 output->f_rescale = l2_sqr * sqrt_d / l1_norm; /* dp_multiplier */
716
717 /* Cleanup */
718 24012 vs_free_aligned(residual);
719 24012 vs_free_aligned(transformed);
720 24012 vs_free_aligned(xu_cb);
721
722 24012 return 0;
723 }
724
725 /*
726 * Batch encode implementation — always_inline so that the specialized
727 * wrappers below pass a static const Vec32TypeOps from the header,
728 * enabling the compiler to inline through every vtable function pointer.
729 * VS_TARGET_CLONES on the wrappers generates AVX2/AVX-512 variants.
730 *
731 * For f32 input, to_float_block returns the input pointer (zero-copy).
732 * For f16, it bulk-converts all vectors in one SIMD-dispatched call.
733 * Bulk conversion via to_float_block amortizes call overhead for f16
734 * (one SIMD-dispatched call vs N per-vector calls).
735 */
736 __attribute__((always_inline)) static inline int
737 ✗ rabitq_encode_batch_impl(
738 const RaBitQParams *params,
739 const void *vectors,
740 Vec32Ref centroid,
741 float *f_add,
742 float *f_rescale,
743 uint8_t *bits,
744 uint16_t count,
745 const Vec32TypeOps *ops)
746 {
747 40 Dimension dim = params->dim;
748 40 uint32_t packed_bytes = params->packed_bytes;
749
750 /* Bulk-convert to f32 if needed. For f32, returns input pointer
751 * (zero-copy). For f16, converts into conv_buf via SIMD. */
752 ✗ float *conv_buf =
753 40 vs_alloc_aligned((size_t)count * dim * sizeof(float), 64);
754 40 const float *fvecs = ops->to_float_block(vectors, conv_buf, count, dim);
755
756 /* Allocate batch buffers */
757 ✗ float *residuals =
758 40 vs_alloc_aligned((size_t)count * dim * sizeof(float), 64);
759 ✗ float *transformed =
760 40 vs_alloc_aligned((size_t)count * dim * sizeof(float), 64);
761 40 float *cent_rotated = vs_alloc_aligned(dim * sizeof(float), 64);
762
763
3/18
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✗ Branch 3 not taken.
✗ Branch 4 not taken.
✗ Branch 5 not taken.
✗ Branch 6 not taken.
✗ Branch 7 not taken.
✗ Branch 8 not taken.
✗ Branch 9 not taken.
✗ Branch 10 not taken.
✗ Branch 11 not taken.
✓ Branch 12 taken 40 times.
✗ Branch 13 not taken.
✓ Branch 14 taken 40 times.
✗ Branch 15 not taken.
✗ Branch 16 not taken.
✓ Branch 17 taken 40 times.
40 if (residuals == NULL || transformed == NULL || cent_rotated == NULL)
764 {
765 ✗ if (residuals)
766 ✗ vs_free_aligned(residuals);
767 ✗ if (transformed)
768 ✗ vs_free_aligned(transformed);
769 ✗ if (cent_rotated)
770 ✗ vs_free_aligned(cent_rotated);
771 ✗ vs_free_aligned(conv_buf);
772 ✗ return -1;
773 }
774
775 /* Step 1: Compute all residuals = vectors[i] - centroid */
776
2/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✗ Branch 3 not taken.
✓ Branch 4 taken 328 times.
✓ Branch 5 taken 40 times.
368 for (uint16_t i = 0; i < count; i++)
777 {
778 328 const float *vec = fvecs + i * dim;
779 328 float *res = residuals + i * dim;
780 328 vec32_sub(vec, centroid.data, res, dim);
781 }
782
783 /* Step 2: Batch transform all residuals through P^T
784 * This is the key optimization: matrix P stays in cache while
785 * processing all vectors.
786 */
787
1/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✗ Branch 3 not taken.
✓ Branch 4 taken 40 times.
✗ Branch 5 not taken.
40 if (params->use_fast_rotate)
788
2/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✗ Branch 3 not taken.
✓ Branch 4 taken 328 times.
✓ Branch 5 taken 40 times.
368 for (uint16_t i = 0; i < count; i++)
789 328 vs_fast_rotate_apply(
790 &params->fr,
791 328 residuals + (size_t)i * dim,
792 328 transformed + (size_t)i * dim);
793 else
794 ✗ vs_matrix_transpose_vector_mul_batch(
795 ✗ params->P, residuals, transformed, count, dim);
796
797 /* Step 3: Rotate centroid once (shared across all vectors) */
798 40 rabitq_rotate(params, centroid.data, cent_rotated);
799
800 /* Step 4: Process each transformed vector to extract bits and factors */
801
2/8
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✗ Branch 3 not taken.
✓ Branch 4 taken 328 times.
✓ Branch 5 taken 40 times.
✗ Branch 7 not taken.
✗ Branch 8 not taken.
368 for (uint16_t i = 0; i < count; i++)
802 {
803 328 const float *trans = transformed + i * dim;
804 328 uint8_t *vec_bits = bits + (size_t)i * packed_bytes;
805
806 /* Extract sign bits (LSB-first packing, FAISS-compatible) */
807 328 rabitq_extract_signs(trans, vec_bits, dim);
808
809 /* Compute xu_cb = binary_code - 0.5 */
810 328 float cb = -0.5f;
811 328 float ip_resi_xucb = 0.0f;
812 328 float ip_cent_xucb = 0.0f;
813 328 float l2_sqr = 0.0f;
814
815
2/8
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✗ Branch 3 not taken.
✓ Branch 4 taken 74752 times.
✓ Branch 5 taken 328 times.
✗ Branch 7 not taken.
✗ Branch 8 not taken.
75080 for (Dimension j = 0; j < dim; j++)
816 {
817 74752 int byte_idx = j / 8;
818 74752 int bit_idx = j % 8;
819 74752 int bit = (vec_bits[byte_idx] >> bit_idx) & 1;
820 74752 float xu_cb_j = (float)bit + cb;
821
822 74752 ip_resi_xucb += trans[j] * xu_cb_j;
823 74752 ip_cent_xucb += cent_rotated[j] * xu_cb_j;
824 74752 l2_sqr += trans[j] * trans[j];
825 }
826
827 /* ip_cent_xucb computed for future centered distance support */
828 ✗ (void)ip_cent_xucb;
829
830 /* Handle corner case: avoid division by near-zero in f_rescale */
831
1/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✗ Branch 3 not taken.
✗ Branch 4 not taken.
✓ Branch 5 taken 328 times.
328 if (fabsf(ip_resi_xucb) < 1e-10f)
832 ✗ ip_resi_xucb = 1e-10f;
833
834 /* L2 distance factors (FAISS-style)
835 * dp_multiplier = ||v-c||² * sqrt(d) / ||v-c||_1
836 * ip_resi_xucb = 0.5 * ||v-c||_1, so ||v-c||_1 = 2 * |ip_resi_xucb|
837 */
838 328 float sqrt_d = sqrtf((float)dim);
839 328 float l1_norm = 2.0f * fabsf(ip_resi_xucb);
840 328 f_add[i] = l2_sqr;
841 328 f_rescale[i] = l2_sqr * sqrt_d / l1_norm; /* dp_multiplier */
842 }
843
844 /* Cleanup */
845 40 vs_free_aligned(residuals);
846 40 vs_free_aligned(transformed);
847 40 vs_free_aligned(cent_rotated);
848 40 vs_free_aligned(conv_buf);
849
850 40 return 0;
851 }
852
853 /* Specialized wrappers — VS_TARGET_CLONES generates SIMD variants */
854
855 VS_TARGET_CLONES static int
856 40 rabitq_encode_batch_f32(
857 const RaBitQParams *params,
858 const void *vectors,
859 Vec32Ref centroid,
860 float *f_add,
861 float *f_rescale,
862 uint8_t *bits,
863 uint16_t count)
864 {
865 80 return rabitq_encode_batch_impl(
866 params,
867 vectors,
868 centroid,
869 f_add,
870 f_rescale,
871 bits,
872 count,
873 &vs_f32_type_ops);
874 }
875
876 VS_TARGET_CLONES static int
877 ✗ rabitq_encode_batch_f16(
878 const RaBitQParams *params,
879 const void *vectors,
880 Vec32Ref centroid,
881 float *f_add,
882 float *f_rescale,
883 uint8_t *bits,
884 uint16_t count)
885 {
886 ✗ return rabitq_encode_batch_impl(
887 params,
888 vectors,
889 centroid,
890 f_add,
891 f_rescale,
892 bits,
893 count,
894 &vs_f16_type_ops);
895 }
896
897 #if defined(VS_F16C_SUPPORT) && !defined(VS_SIMD_NONE)
898 VS_TARGET_F16C_AVX2 static int
899 ✗ rabitq_encode_batch_f16c(
900 const RaBitQParams *params,
901 const void *vectors,
902 Vec32Ref centroid,
903 float *f_add,
904 float *f_rescale,
905 uint8_t *bits,
906 uint16_t count)
907 {
908 ✗ return rabitq_encode_batch_impl(
909 params,
910 vectors,
911 centroid,
912 f_add,
913 f_rescale,
914 bits,
915 count,
916 &vs_f16c_type_ops);
917 }
918 #endif
919
920 int
921 56 vs_rabitq_encode_batch(
922 const RaBitQParams *params,
923 const void *vectors,
924 VecType vec_type,
925 Vec32Ref centroid,
926 float *f_add,
927 float *f_rescale,
928 uint8_t *bits,
929 uint16_t count)
930 {
931
8/8
✓ Branch 0 taken 54 times.
✓ Branch 1 taken 2 times.
✓ Branch 2 taken 52 times.
✓ Branch 3 taken 2 times.
✓ Branch 4 taken 50 times.
✓ Branch 5 taken 2 times.
✓ Branch 6 taken 48 times.
✓ Branch 7 taken 2 times.
56 if (params == NULL || vectors == NULL || centroid.data == NULL ||
932
6/6
✓ Branch 0 taken 46 times.
✓ Branch 1 taken 2 times.
✓ Branch 2 taken 44 times.
✓ Branch 3 taken 2 times.
✓ Branch 4 taken 2 times.
✓ Branch 5 taken 42 times.
48 f_add == NULL || f_rescale == NULL || bits == NULL || count == 0)
933 14 return -1;
934
935
2/2
✓ Branch 0 taken 2 times.
✓ Branch 1 taken 40 times.
42 if (centroid.dim != params->dim)
936 2 return -1;
937
938 /* Single dispatch point — selects the inline vtable once */
939
1/3
✓ Branch 0 taken 40 times.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
40 switch (vec_type)
940 {
941 40 case VS_VEC_F32:
942 40 return rabitq_encode_batch_f32(
943 params, vectors, centroid, f_add, f_rescale, bits, count);
944 #if defined(VS_F16C_SUPPORT) && !defined(VS_SIMD_NONE)
945 ✗ case VS_VEC_F16C:
946 ✗ return rabitq_encode_batch_f16c(
947 params, vectors, centroid, f_add, f_rescale, bits, count);
948 #endif
949 ✗ default:
950 ✗ return rabitq_encode_batch_f16(
951 params, vectors, centroid, f_add, f_rescale, bits, count);
952 }
953 }
954
955 /*
956 * Batch helpers
957 */
958
959 RaBitQBatch *
960 ✗ vs_rabitq_batch_create(uint16_t count, Dimension dim)
961 {
962 ✗ if (count == 0 || dim == 0)
963 ✗ return NULL;
964
965 ✗ RaBitQBatch *batch = vs_alloc(sizeof(RaBitQBatch));
966 ✗ if (batch == NULL)
967 ✗ return NULL;
968
969 ✗ uint32_t packed_bytes = VS_RABITQ_BYTES(dim);
970
971 ✗ batch->count = count;
972 ✗ batch->packed_bytes = (uint16_t)packed_bytes;
973 ✗ batch->f_add = vs_alloc(count * sizeof(float));
974 ✗ batch->f_rescale = vs_alloc(count * sizeof(float));
975 ✗ batch->bits = vs_alloc((size_t)count * packed_bytes);
976
977 ✗ if (batch->f_add == NULL || batch->f_rescale == NULL ||
978 ✗ batch->bits == NULL)
979 {
980 ✗ vs_rabitq_batch_destroy(batch);
981 ✗ return NULL;
982 }
983
984 ✗ return batch;
985 }
986
987 void
988 ✗ vs_rabitq_batch_destroy(RaBitQBatch *batch)
989 {
990 ✗ if (batch == NULL)
991 ✗ return;
992
993 ✗ if (batch->f_add)
994 ✗ vs_free(batch->f_add);
995 ✗ if (batch->f_rescale)
996 ✗ vs_free(batch->f_rescale);
997 ✗ if (batch->bits)
998 ✗ vs_free(batch->bits);
999 ✗ vs_free(batch);
1000 }
1001
1002 RaBitQBatch *
1003 ✗ vs_rabitq_encode_batch_alloc(
1004 const RaBitQParams *params,
1005 const void *vectors,
1006 VecType vec_type,
1007 Vec32Ref centroid,
1008 uint16_t count)
1009 {
1010 ✗ if (params == NULL || vectors == NULL || centroid.data == NULL ||
1011 count == 0)
1012 ✗ return NULL;
1013
1014 ✗ RaBitQBatch *batch = vs_rabitq_batch_create(count, params->dim);
1015 ✗ if (batch == NULL)
1016 ✗ return NULL;
1017
1018 ✗ if (vs_rabitq_encode_batch(
1019 params,
1020 vectors,
1021 vec_type,
1022 centroid,
1023 batch->f_add,
1024 batch->f_rescale,
1025 batch->bits,
1026 count) != 0)
1027 {
1028 ✗ vs_rabitq_batch_destroy(batch);
1029 ✗ return NULL;
1030 }
1031
1032 ✗ return batch;
1033 }
1034
1035 /*
1036 * Query preparation
1037 */
1038
1039 RaBitQQueryState *
1040 6704 vs_rabitq_prepare_query_ex(
1041 const RaBitQParams *params,
1042 Vec32Ref query,
1043 Vec32Ref centroid,
1044 VsDistanceMode mode)
1045 {
1046
6/6
✓ Branch 0 taken 6702 times.
✓ Branch 1 taken 2 times.
✓ Branch 2 taken 6700 times.
✓ Branch 3 taken 2 times.
✓ Branch 4 taken 2 times.
✓ Branch 5 taken 6698 times.
6704 if (params == NULL || query.data == NULL || centroid.data == NULL)
1047 6 return NULL;
1048
1049
4/4
✓ Branch 0 taken 6696 times.
✓ Branch 1 taken 2 times.
✓ Branch 2 taken 2 times.
✓ Branch 3 taken 6694 times.
6698 if (query.dim != params->dim || centroid.dim != params->dim)
1050 4 return NULL;
1051
1052 6694 Dimension dim = params->dim;
1053
1054 6694 RaBitQQueryState *state = vs_alloc(sizeof(RaBitQQueryState));
1055
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 6694 times.
6694 if (state == NULL)
1056 ✗ return NULL;
1057
1058 6694 state->transformed = vs_alloc_aligned(dim * sizeof(float), 64);
1059
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 6694 times.
6694 if (state->transformed == NULL)
1060 {
1061 ✗ vs_free(state);
1062 ✗ return NULL;
1063 }
1064
1065 6694 state->dim = dim;
1066
1067 /* Compute residual = query - centroid */
1068 6694 float *residual = vs_alloc_aligned(dim * sizeof(float), 64);
1069
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 6694 times.
6694 if (residual == NULL)
1070 {
1071 ✗ vs_free_aligned(state->transformed);
1072 ✗ vs_free(state);
1073 ✗ return NULL;
1074 }
1075
1076 6694 vec32_sub(query.data, centroid.data, residual, dim);
1077
1078 /* Transform through P^T */
1079 6694 rabitq_rotate(params, residual, state->transformed);
1080
1081 /* Compute g_add = ||query - centroid||^2 */
1082 6694 state->g_add = vs_l2_norm_squared(state->transformed, dim);
1083
1084 /* Compute g_error = sqrt(g_add) for error bound */
1085 6694 state->g_error = sqrtf(state->g_add);
1086
1087 /* Compute sum of transformed values for FAISS-style distance formula */
1088 6694 state->sum_transformed = vec32_sum(state->transformed, dim);
1089
1090 /* Precompute 1/sqrt(dim) for distance formula */
1091 6694 state->inv_sqrt_d = 1.0f / sqrtf((float)dim);
1092
1093 /* Precompute C_error = 2*ε/√(d-1) for deriving f_error from compact
1094 * data */
1095
1/2
✓ Branch 0 taken 6694 times.
✗ Branch 1 not taken.
6694 if (dim > 1)
1096 6694 state->c_error = 2.0f * VS_RABITQ_EPSILON / sqrtf((float)(dim - 1));
1097 else
1098 ✗ state->c_error = 0.0f;
1099
1100 /* Compute symmetric search fields: query sign bits and g_scale */
1101 6694 uint32_t packed_bytes = VS_RABITQ_BYTES(dim);
1102 6694 state->query_bits = vs_alloc_aligned(packed_bytes, 64);
1103
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 6694 times.
6694 if (state->query_bits == NULL)
1104 {
1105 ✗ vs_free_aligned(state->transformed);
1106 ✗ vs_free_aligned(residual);
1107 ✗ vs_free(state);
1108 ✗ return NULL;
1109 }
1110 6694 rabitq_extract_signs(state->transformed, state->query_bits, dim);
1111
1112 /* g_scale = mean(|transformed|) = L1(transformed) / dim */
1113 6694 float l1_sum = 0.0f;
1114
2/3
✓ Branch 0 taken 774410 times.
✓ Branch 1 taken 6694 times.
✗ Branch 2 not taken.
781104 for (Dimension i = 0; i < dim; i++)
1115 774410 l1_sum += fabsf(state->transformed[i]);
1116 6694 state->g_scale = l1_sum / (float)dim;
1117
1118 /* Set dispatch function pointers and error multiplier based on mode */
1119 6694 state->mode = mode;
1120
2/2
✓ Branch 0 taken 6 times.
✓ Branch 1 taken 6688 times.
6694 if (mode == VS_DISTANCE_MODE_SYMMETRIC)
1121 {
1122 6 state->distance_fn = vs_rabitq_distance_symmetric;
1123 6 state->distance_with_bound_fn =
1124 vs_rabitq_distance_symmetric_with_bound;
1125 6 state->error_multiplier = 3.0f;
1126 }
1127 else
1128 {
1129 6688 state->distance_fn = vs_rabitq_distance;
1130 6688 state->distance_with_bound_fn = vs_rabitq_distance_with_bound;
1131 6688 state->error_multiplier = 1.0f;
1132 }
1133
1134 6694 vs_free_aligned(residual);
1135
1136 6694 return state;
1137 }
1138
1139 RaBitQQueryState *
1140 6694 vs_rabitq_prepare_query(
1141 const RaBitQParams *params, Vec32Ref query, Vec32Ref centroid)
1142 {
1143 6694 return vs_rabitq_prepare_query_ex(
1144 params, query, centroid, VS_DISTANCE_MODE_ASYMMETRIC);
1145 }
1146
1147 void
1148 661486 vs_rabitq_rotate(const RaBitQParams *params, const float *input, float *output)
1149 {
1150 661486 rabitq_rotate(params, input, output);
1151 661486 }
1152
1153 void
1154 36250 vs_rabitq_init_query_constants(RaBitQQueryState *state, Dimension dim)
1155 {
1156 36250 state->dim = dim;
1157 36250 state->inv_sqrt_d = 1.0f / sqrtf((float)dim);
1158
1/2
✓ Branch 0 taken 36250 times.
✗ Branch 1 not taken.
36250 if (dim > 1)
1159 36250 state->c_error = 2.0f * VS_RABITQ_EPSILON / sqrtf((float)(dim - 1));
1160 else
1161 ✗ state->c_error = 0.0f;
1162 36250 }
1163
1164 void
1165 614070 vs_rabitq_init_query_state(
1166 RaBitQQueryState *state,
1167 const float *pt_query,
1168 const float *pt_centroid,
1169 Dimension dim,
1170 VsDistanceMode mode)
1171 {
1172 /* transformed = pt_query - pt_centroid (O(dim) vector subtraction) */
1173 614070 vec32_sub(pt_query, pt_centroid, state->transformed, dim);
1174
1175 /* Compute per-centroid scalar fields from transformed.
1176 * inv_sqrt_d and c_error are dim-dependent constants set once
1177 * via vs_rabitq_init_query_constants(). */
1178 614070 state->g_add = vs_l2_norm_squared(state->transformed, dim);
1179 614070 state->g_error = sqrtf(state->g_add);
1180 614070 state->sum_transformed = vec32_sum(state->transformed, dim);
1181
1182 /* Dispatch pointers */
1183 614070 state->mode = mode;
1184
2/2
✓ Branch 0 taken 12 times.
✓ Branch 1 taken 614058 times.
614070 if (mode == VS_DISTANCE_MODE_SYMMETRIC)
1185 {
1186 /* Symmetric mode needs sign bits and L1 norm */
1187 12 rabitq_extract_signs(state->transformed, state->query_bits, dim);
1188
1189 12 float l1_sum = 0.0f;
1190
2/3
✓ Branch 0 taken 384 times.
✓ Branch 1 taken 12 times.
✗ Branch 2 not taken.
396 for (Dimension i = 0; i < dim; i++)
1191 384 l1_sum += fabsf(state->transformed[i]);
1192 12 state->g_scale = l1_sum / (float)dim;
1193
1194 12 state->distance_fn = vs_rabitq_distance_symmetric;
1195 12 state->distance_with_bound_fn =
1196 vs_rabitq_distance_symmetric_with_bound;
1197 12 state->error_multiplier = 3.0f;
1198 }
1199 else
1200 {
1201 614058 state->distance_fn = vs_rabitq_distance;
1202 614058 state->distance_with_bound_fn = vs_rabitq_distance_with_bound;
1203 614058 state->error_multiplier = 1.0f;
1204 }
1205 614070 }
1206
1207 void
1208 6696 vs_rabitq_free_query(RaBitQQueryState *state)
1209 {
1210
2/2
✓ Branch 0 taken 2 times.
✓ Branch 1 taken 6694 times.
6696 if (state == NULL)
1211 2 return;
1212
1213
1/2
✓ Branch 0 taken 6694 times.
✗ Branch 1 not taken.
6694 if (state->transformed != NULL)
1214 6694 vs_free_aligned(state->transformed);
1215
1/2
✓ Branch 0 taken 6694 times.
✗ Branch 1 not taken.
6694 if (state->query_bits != NULL)
1216 6694 vs_free_aligned(state->query_bits);
1217
1218 6694 vs_free(state);
1219 }
1220
1221 /*
1222 * Distance computation
1223 */
1224
1225 Distance
1226 6872 vs_rabitq_distance(
1227 const RaBitQQueryState *query_state,
1228 const RaBitQData *data,
1229 Dimension dim)
1230 {
1231
4/4
✓ Branch 0 taken 6868 times.
✓ Branch 1 taken 4 times.
✓ Branch 2 taken 4 times.
✓ Branch 3 taken 6864 times.
6872 if (query_state == NULL || data == NULL)
1232 8 return -1.0f;
1233
1234
2/2
✓ Branch 0 taken 6 times.
✓ Branch 1 taken 6858 times.
6864 if (query_state->dim != dim)
1235 6 return -1.0f;
1236
1237 /* Compute binary inner product: sum of transformed[i] where bit[i] = 1 */
1238 ✗ float binary_ip =
1239 6858 rabitq_inner_product(query_state->transformed, data->bits, dim);
1240
1241 /* FAISS-style distance formula (better accuracy with non-quantized query):
1242 *
1243 * est_dist = ||v-c||² + ||q-c||² - 2 * dp_multiplier * final_dot
1244 *
1245 * where:
1246 * - ||v-c||² = f_add (stored per vector)
1247 * - ||q-c||² = g_add (stored in query state)
1248 * - dp_multiplier = ||v-c||² * sqrt(d) / ||v-c||_1 = f_rescale
1249 * - final_dot = (2 * binary_ip - sum_transformed) / sqrt(d)
1250 *
1251 * FAISS normalizes dp_multiplier by the L1 norm to correct for the
1252 * distribution of residual values. This gives better distance estimates.
1253 */
1254 6858 float final_dot = (2.0f * binary_ip - query_state->sum_transformed) *
1255 6858 query_state->inv_sqrt_d;
1256
1257 6858 float est_dist = data->f_add + query_state->g_add -
1258 6858 2.0f * data->f_rescale * final_dot;
1259
1260 6858 return est_dist;
1261 }
1262
1263 /*
1264 * Shared distance formula for batch functions
1265 *
1266 * Applies: distances[i] = f_add[i] + g_add - 2 * f_rescale[i] * final_dots[i]
1267 * with optional error bound derivation using qstate->error_multiplier.
1268 *
1269 * Callers convert mode-specific intermediates (IPs or Hamming distances)
1270 * into uniform final_dots[] before calling this. For symmetric mode,
1271 * g_scale is folded into final_dots so the formula is identical.
1272 *
1273 * Marked always_inline so the compiler sees the full loop body in each
1274 * caller, preserving auto-vectorization of the f_add[]/f_rescale[] math.
1275 */
1276 __attribute__((always_inline)) static inline void
1277 293847 rabitq_apply_distances(
1278 const RaBitQQueryState *qstate,
1279 const float *f_add,
1280 const float *f_rescale,
1281 const float *final_dots,
1282 uint32_t count,
1283 Dimension dim,
1284 Distance *distances,
1285 Distance *lower_bounds)
1286 {
1287 1258339 float g_add = qstate->g_add;
1288
1289
2/2
✓ Branch 0 taken 14809753 times.
✓ Branch 1 taken 1258339 times.
16068092 for (uint32_t i = 0; i < count; i++)
1290 {
1291 14809753 distances[i] = f_add[i] + g_add - 2.0f * f_rescale[i] * final_dots[i];
1292
1293
2/2
✓ Branch 0 taken 16 times.
✓ Branch 1 taken 14809737 times.
14809753 if (lower_bounds != NULL)
1294 {
1295 16 float f_error = rabitq_derive_f_error(
1296 16 f_add[i], f_rescale[i], qstate->c_error, dim);
1297 16 lower_bounds[i] = rabitq_lower_bound(
1298 16 distances[i],
1299 f_error,
1300 16 qstate->g_error,
1301 16 qstate->error_multiplier);
1302 }
1303 }
1304 964492 }
1305
1306 void
1307 46 vs_rabitq_distance_batch(
1308 const RaBitQQueryState *qstate,
1309 const float *f_add,
1310 const float *f_rescale,
1311 const uint8_t *bits,
1312 uint32_t count,
1313 Dimension dim,
1314 Distance *distances)
1315 {
1316
10/10
✓ Branch 0 taken 44 times.
✓ Branch 1 taken 2 times.
✓ Branch 2 taken 42 times.
✓ Branch 3 taken 2 times.
✓ Branch 4 taken 40 times.
✓ Branch 5 taken 2 times.
✓ Branch 6 taken 38 times.
✓ Branch 7 taken 2 times.
✓ Branch 8 taken 36 times.
✓ Branch 9 taken 2 times.
46 if (qstate == NULL || f_add == NULL || f_rescale == NULL || bits == NULL ||
1317
2/2
✓ Branch 0 taken 2 times.
✓ Branch 1 taken 34 times.
36 distances == NULL || count == 0)
1318 12 return;
1319
1320 34 uint32_t packed_bytes = VS_RABITQ_BYTES(dim);
1321 34 float g_add = qstate->g_add;
1322 34 float sum_t = qstate->sum_transformed;
1323 34 float inv_sqrt_d = qstate->inv_sqrt_d;
1324
1325
2/2
✓ Branch 0 taken 280 times.
✓ Branch 1 taken 34 times.
314 for (uint32_t i = 0; i < count; i++)
1326 {
1327 280 float ip = rabitq_inner_product(
1328 280 qstate->transformed, bits + (size_t)i * packed_bytes, dim);
1329 280 float final_dot = (2.0f * ip - sum_t) * inv_sqrt_d;
1330 280 distances[i] = f_add[i] + g_add - 2.0f * f_rescale[i] * final_dot;
1331 }
1332 }
1333
1334 void
1335 1286799 vs_rabitq_inner_product_multi(
1336 const float *transformed,
1337 const uint8_t *bits,
1338 uint32_t stride,
1339 Dimension dim,
1340 uint32_t count,
1341 float *results)
1342 {
1343
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 1286799 times.
1286799 if (vs_unlikely(!g_rabitq_initialized))
1344 ✗ vs_rabitq_init_simd();
1345 1286799 g_inner_product_multi_fn(transformed, bits, stride, dim, count, results);
1346 1286799 }
1347
1348 void
1349 34 vs_rabitq_distance_batch_multi(
1350 const RaBitQQueryState *qstate,
1351 const float *f_add,
1352 const float *f_rescale,
1353 const uint8_t *bits,
1354 uint32_t count,
1355 Dimension dim,
1356 Distance *distances)
1357 {
1358 34 float *scratch = vs_alloc(count * sizeof(float));
1359 34 vs_rabitq_distance_batch_multi_with_bound(
1360 qstate,
1361 f_add,
1362 f_rescale,
1363 bits,
1364 34 VS_RABITQ_BYTES(dim),
1365 count,
1366 dim,
1367 distances,
1368 NULL,
1369 scratch);
1370 34 vs_free(scratch);
1371 34 }
1372
1373 void
1374 1258339 vs_rabitq_distance_batch_multi_with_bound(
1375 const RaBitQQueryState *qstate,
1376 const float *f_add,
1377 const float *f_rescale,
1378 const uint8_t *bits,
1379 uint32_t stride,
1380 uint32_t count,
1381 Dimension dim,
1382 Distance *distances,
1383 Distance *lower_bounds,
1384 float *scratch)
1385 {
1386
5/10
✓ Branch 0 taken 1258339 times.
✗ Branch 1 not taken.
✓ Branch 2 taken 1258339 times.
✗ Branch 3 not taken.
✓ Branch 4 taken 964492 times.
✗ Branch 5 not taken.
✓ Branch 6 taken 964492 times.
✗ Branch 7 not taken.
✓ Branch 8 taken 964492 times.
✗ Branch 9 not taken.
1258339 if (qstate == NULL || f_add == NULL || f_rescale == NULL || bits == NULL ||
1387
2/2
✓ Branch 0 taken 293847 times.
✓ Branch 1 taken 964492 times.
1258339 distances == NULL || count == 0)
1388 ✗ return;
1389
1390 /* Compute all inner products in a single multi-candidate pass */
1391 1258339 vs_rabitq_inner_product_multi(
1392 1258339 qstate->transformed, bits, stride, dim, count, scratch);
1393
1394 /* Convert IPs to final_dots in-place */
1395 1258339 float sum_t = qstate->sum_transformed;
1396 1258339 float inv_sqrt_d = qstate->inv_sqrt_d;
1397
2/2
✓ Branch 0 taken 14809753 times.
✓ Branch 1 taken 1258339 times.
16068092 for (uint32_t i = 0; i < count; i++)
1398 14809753 scratch[i] = (2.0f * scratch[i] - sum_t) * inv_sqrt_d;
1399
1400 1258339 rabitq_apply_distances(
1401 qstate,
1402 f_add,
1403 f_rescale,
1404 scratch,
1405 count,
1406 dim,
1407 distances,
1408 lower_bounds);
1409 }
1410
1411 void
1412 12 vs_rabitq_distance_batch_symmetric_with_bound(
1413 const RaBitQQueryState *qstate,
1414 const float *f_add,
1415 const float *f_rescale,
1416 const uint8_t *bits,
1417 uint32_t stride,
1418 uint32_t count,
1419 Dimension dim,
1420 Distance *distances,
1421 Distance *lower_bounds,
1422 uint32_t *scratch)
1423 {
1424
5/10
✓ Branch 0 taken 12 times.
✗ Branch 1 not taken.
✓ Branch 2 taken 12 times.
✗ Branch 3 not taken.
✓ Branch 4 taken 12 times.
✗ Branch 5 not taken.
✓ Branch 6 taken 12 times.
✗ Branch 7 not taken.
✓ Branch 8 taken 12 times.
✗ Branch 9 not taken.
12 if (qstate == NULL || f_add == NULL || f_rescale == NULL || bits == NULL ||
1425
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 12 times.
12 distances == NULL || count == 0)
1426 ✗ return;
1427
1428 12 uint32_t packed_bytes = VS_RABITQ_BYTES(dim);
1429
1430 /* Compute all Hamming distances in a single multi-candidate pass */
1431 12 vs_rabitq_hamming_distance_multi(
1432 12 qstate->query_bits, bits, stride, packed_bytes, count, scratch);
1433
1434 /* Apply symmetric distance formula (+ optional error bounds) */
1435 12 float g_add = qstate->g_add;
1436 12 float g_scale = qstate->g_scale;
1437 12 float inv_sqrt_d = qstate->inv_sqrt_d;
1438
1439
2/2
✓ Branch 0 taken 66 times.
✓ Branch 1 taken 12 times.
78 for (uint32_t i = 0; i < count; i++)
1440 {
1441 66 float sym_dot = (float)((int32_t)dim - 2 * (int32_t)scratch[i]);
1442 66 float final_dot = sym_dot * inv_sqrt_d * g_scale;
1443
1444 66 distances[i] = f_add[i] + g_add - 2.0f * f_rescale[i] * final_dot;
1445
1446
2/2
✓ Branch 0 taken 16 times.
✓ Branch 1 taken 50 times.
66 if (lower_bounds != NULL)
1447 {
1448 16 float f_error = rabitq_derive_f_error(
1449 16 f_add[i], f_rescale[i], qstate->c_error, dim);
1450 16 lower_bounds[i] = rabitq_lower_bound(
1451 16 distances[i],
1452 f_error,
1453 16 qstate->g_error,
1454 16 qstate->error_multiplier);
1455 }
1456 }
1457 }
1458
1459 /*
1460 * Common scalar _with_bound implementation
1461 *
1462 * Uses qstate->distance_fn() (mode-dispatched) and
1463 * qstate->error_multiplier to compute both estimated distance
1464 * and lower bound. Both public _with_bound functions delegate here.
1465 */
1466 static void
1467 6736 rabitq_distance_with_bound_common(
1468 const RaBitQQueryState *qstate,
1469 const RaBitQData *data,
1470 Dimension dim,
1471 Distance *est_dist,
1472 Distance *lower_bound)
1473 {
1474
6/8
✓ Branch 0 taken 6732 times.
✓ Branch 1 taken 4 times.
✓ Branch 2 taken 6728 times.
✓ Branch 3 taken 4 times.
✓ Branch 4 taken 6728 times.
✗ Branch 5 not taken.
✗ Branch 6 not taken.
✓ Branch 7 taken 6728 times.
6736 if (qstate == NULL || data == NULL || est_dist == NULL ||
1475 ✗ lower_bound == NULL)
1476 {
1477
2/2
✓ Branch 0 taken 6 times.
✓ Branch 1 taken 2 times.
8 if (est_dist)
1478 6 *est_dist = -1.0f;
1479
2/2
✓ Branch 0 taken 6 times.
✓ Branch 1 taken 2 times.
8 if (lower_bound)
1480 6 *lower_bound = -1.0f;
1481 8 return;
1482 }
1483
1484 6728 *est_dist = qstate->distance_fn(qstate, data, dim);
1485
1486 6728 float f_error = rabitq_derive_f_error(
1487 6728 data->f_add, data->f_rescale, qstate->c_error, dim);
1488 6728 *lower_bound = rabitq_lower_bound(
1489 6728 *est_dist, f_error, qstate->g_error, qstate->error_multiplier);
1490 }
1491
1492 void
1493 6634 vs_rabitq_distance_with_bound(
1494 const RaBitQQueryState *query_state,
1495 const RaBitQData *data,
1496 Dimension dim,
1497 Distance *est_dist,
1498 Distance *lower_bound)
1499 {
1500 6634 rabitq_distance_with_bound_common(
1501 query_state, data, dim, est_dist, lower_bound);
1502 6634 }
1503
1504 /*
1505 * Hamming distance - public API
1506 */
1507
1508 uint32_t
1509 132 vs_rabitq_hamming_distance(
1510 const uint8_t *a, const uint8_t *b, uint32_t packed_bytes)
1511 {
1512
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 132 times.
132 if (vs_unlikely(!g_rabitq_initialized))
1513 ✗ vs_rabitq_init_simd();
1514 132 return g_hamming_fn(a, b, packed_bytes);
1515 }
1516
1517 void
1518 14 vs_rabitq_hamming_distance_multi(
1519 const uint8_t *query_bits,
1520 const uint8_t *data_bits,
1521 uint32_t stride,
1522 uint32_t packed_bytes,
1523 uint32_t count,
1524 uint32_t *results)
1525 {
1526
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 14 times.
14 if (vs_unlikely(!g_rabitq_initialized))
1527 ✗ vs_rabitq_init_simd();
1528 14 g_hamming_multi_fn(
1529 query_bits, data_bits, stride, packed_bytes, count, results);
1530 14 }
1531
1532 /*
1533 * Symmetric distance computation
1534 *
1535 * Both query and data are 1-bit quantized. Uses Hamming distance
1536 * (XOR + popcount) instead of asymmetric inner product (mask + add).
1537 * ~32x fewer inner loop iterations at the cost of additional query
1538 * quantization error.
1539 */
1540
1541 Distance
1542 50 vs_rabitq_distance_symmetric(
1543 const RaBitQQueryState *qstate, const RaBitQData *data, Dimension dim)
1544 {
1545
4/4
✓ Branch 0 taken 48 times.
✓ Branch 1 taken 2 times.
✓ Branch 2 taken 2 times.
✓ Branch 3 taken 46 times.
50 if (qstate == NULL || data == NULL)
1546 4 return -1.0f;
1547
1548
2/2
✓ Branch 0 taken 2 times.
✓ Branch 1 taken 44 times.
46 if (qstate->dim != dim)
1549 2 return -1.0f;
1550
1551 44 uint32_t packed_bytes = VS_RABITQ_BYTES(dim);
1552 44 uint32_t hamming = vs_rabitq_hamming_distance(
1553 44 qstate->query_bits, data->bits, packed_bytes);
1554
1555 /* sym_dot = dim - 2 * hamming (range: [-dim, dim]) */
1556 44 float sym_dot = (float)((int32_t)dim - 2 * (int32_t)hamming);
1557 44 float final_dot = sym_dot * qstate->inv_sqrt_d;
1558
1559 /* est_dist = f_add + g_add - 2 * f_rescale * g_scale * final_dot */
1560 44 return data->f_add + qstate->g_add -
1561 44 2.0f * data->f_rescale * qstate->g_scale * final_dot;
1562 }
1563
1564 void
1565 102 vs_rabitq_distance_symmetric_with_bound(
1566 const RaBitQQueryState *qstate,
1567 const RaBitQData *data,
1568 Dimension dim,
1569 Distance *est_dist,
1570 Distance *lower_bound)
1571 {
1572 102 rabitq_distance_with_bound_common(
1573 qstate, data, dim, est_dist, lower_bound);
1574 102 }
1575
1576 void
1577 4 vs_rabitq_distance_batch_symmetric(
1578 const RaBitQQueryState *qstate,
1579 const float *f_add,
1580 const float *f_rescale,
1581 const uint8_t *bits,
1582 uint32_t count,
1583 Dimension dim,
1584 Distance *distances)
1585 {
1586 4 uint32_t *scratch = vs_alloc(count * sizeof(uint32_t));
1587 4 vs_rabitq_distance_batch_symmetric_with_bound(
1588 qstate,
1589 f_add,
1590 f_rescale,
1591 bits,
1592 4 VS_RABITQ_BYTES(dim),
1593 count,
1594 dim,
1595 distances,
1596 NULL,
1597 scratch);
1598 4 vs_free(scratch);
1599 4 }
1600