GCC Code Coverage Report


Directory: src/
File: src/quant/rabitq.h
Date: 2026-09-30 11:11:31
Exec Total Coverage
Lines: 8 8 100.0%
Functions: 3 3 100.0%
Branches: 2 2 100.0%

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.h - RaBitQ (Randomized Binary Quantization) for ANN search
6 *
7 * RaBitQ compresses D-dimensional vectors to D bits (32x compression) while
8 * maintaining 95-99% recall through theoretical error bounds. Key insight:
9 * vertices of a hypercube (±1/√D per coordinate) are evenly spread on the
10 * unit hypersphere.
11 *
12 * Algorithm:
13 * 1. Transform input via random orthogonal matrix P (ensures isotropic
14 * distribution)
15 * 2. Compute residual: r = P^T * (data - centroid)
16 * 3. Extract sign bits: bits[i] = (r[i] > 0) ? 1 : 0
17 * 4. Compute factors for distance estimation: f_add, f_rescale
18 * 5. At query time: est_dist = f_add + f_rescale * inner_product(q', bits)
19 * 6. Lower bound: lower_bound = est_dist - f_error * g_error
20 *
21 * Reference: RaBitQ-Library (https://github.com/nmslib/RaBitQ-Library)
22 */
23
24 #ifndef VS_RABITQ_H
25 #define VS_RABITQ_H
26
27 #include <stdbool.h>
28 #include <stdint.h>
29
30 #include "core/types.h"
31 #include "quant/fast_rotate.h"
32
33 /*
34 * Error constant from RaBitQ paper (empirically tuned by authors).
35 * Controls the tightness of error bounds.
36 */
37 #define VS_RABITQ_EPSILON 1.9f
38
39 /*
40 * Rotation seed for every index build. The rotation matrix is derived from
41 * the seed at decode time too, so every site that creates build-time RaBitQ
42 * params MUST use this seed: an index encoded under one seed and read under
43 * another decodes garbage with no error.
44 */
45 #define VS_RABITQ_BUILD_SEED UINT64_C(42)
46
47 /*
48 * RaBitQVector - Quantized vector (PostgreSQL varlena-compatible)
49 *
50 * Stores D bits packed into ceil(D/8) bytes, plus two float factors
51 * for distance estimation. The data portion (f_add, f_rescale, bits[])
52 * is layout-compatible with RaBitQData for zero-copy access via
53 * VS_RABITQ_DATA().
54 *
55 * This struct is reserved for the future PostgreSQL SQL type. Encoding
56 * functions produce RaBitQData (compact form) instead.
57 *
58 * Total size: 16 bytes header + ceil(dim/8) bytes data
59 */
60 typedef struct RaBitQVector
61 {
62 int32_t vl_len_; /* varlena header (for PG compatibility) */
63 int16_t dim; /* dimensions = number of bits */
64 int16_t flags; /* reserved for future use */
65 float f_add; /* additive factor for distance estimation */
66 float f_rescale; /* scaling factor for distance estimation */
67 uint8_t bits[]; /* D/8 bytes, LSB-first bit packing (FAISS-compatible) */
68 } RaBitQVector;
69
70 /* Calculate size of presentation vector structure (for future PG type) */
71 #define VS_RABITQ_VECTOR_SIZE(dim) \
72 (offsetof(RaBitQVector, bits) + (((dim) + 7) / 8))
73
74 /* Number of bytes needed to store dim bits */
75 #define VS_RABITQ_BYTES(dim) (((dim) + 7) / 8)
76
77 /*
78 * RaBitQData - Compact quantized vector (primary encoding form)
79 *
80 * Stores only f_add and f_rescale. f_error is derived at query time:
81 * f_error = C_error * sqrt(f_rescale² - f_add)
82 * where C_error = 2 * VS_RABITQ_EPSILON / sqrt(dim - 1).
83 *
84 * Total size: 8 bytes header + ceil(dim/8) bytes data
85 */
86 typedef struct RaBitQData
87 {
88 float f_add; /* ||v-c||² - additive distance factor */
89 float f_rescale; /* dp_multiplier - scaling factor */
90 uint8_t bits[]; /* D/8 bytes, LSB-first bit packing */
91 } RaBitQData;
92
93 /* Calculate size of compact quantized vector (4-byte aligned for
94 * safe struct access when stored in posting pages) */
95 #define VS_RABITQ_DATA_SIZE(dim) \
96 (((offsetof(RaBitQData, bits) + VS_RABITQ_BYTES(dim)) + 3) & ~3u)
97
98 /* Access compact data portion of a presentation vector (zero-copy cast) */
99 #define VS_RABITQ_DATA(v) ((RaBitQData *)&(v)->f_add)
100
101 /*
102 * RaBitQBatch - Batch of encoded vectors in separate arrays
103 *
104 * Used for batch encoding where separate arrays for each field enable
105 * efficient SIMD processing and bulk page insertion.
106 */
107 typedef struct RaBitQBatch
108 {
109 uint16_t count; /* number of encoded vectors */
110 uint16_t packed_bytes; /* ceil(dim/8) per vector */
111 float *f_add; /* [count] */
112 float *f_rescale; /* [count] */
113 uint8_t *bits; /* [count * packed_bytes] */
114 } RaBitQBatch;
115
116 /*
117 * RaBitQParams - Quantizer parameters (shared per index)
118 *
119 * Contains the random orthogonal matrix and derived values. Generated
120 * once during index creation and shared across all vectors in the index.
121 * The matrix P ensures isotropic distribution of residuals, which is
122 * essential for RaBitQ's error bounds.
123 */
124 typedef struct RaBitQParams
125 {
126 Dimension dim; /* Vector dimension */
127 uint32_t packed_bytes; /* ceil(dim / 8) */
128 uint64_t seed; /* Seed for reproducibility */
129 /*
130 * Rotation. When use_fast_rotate is set (dim is a supported N*K
131 * shape) the orthonormal rotation P^T*x is computed by the O(d log d)
132 * Randomized Hadamard Transform in `fr`, and the dense matrix P is
133 * neither built nor stored (the trailing P[] is allocated only for
134 * unsupported dims, which fall back to the dense O(d^2) multiply).
135 */
136 bool use_fast_rotate;
137 VsFastRotateParams fr;
138 float P[FLEXIBLE_ARRAY_MEMBER]; /* dense P^T (dense path) */
139 } RaBitQParams;
140
141 /*
142 * Dense layout size: header + the dim*dim float matrix. Cast to a fixed
143 * 64-bit type before the multiply so the product is computed in 64 bits on
144 * every platform (plain int32 overflows past dim ~46340, and size_t is still
145 * 32-bit on ILP32). Used when an explicit dense matrix is stored
146 * (vs_rabitq_create_from_matrix) or when the fast rotation is unavailable.
147 */
148 #define VS_RABITQ_PARAMS_DENSE_SIZE(dim) \
149 (offsetof(RaBitQParams, P) + (uint64_t)(dim) * (dim) * sizeof(float))
150
151 /*
152 * Allocation size for seed-derived params. When the fast (Hadamard) rotation
153 * supports the dim, the dense matrix P is not built or stored, so only the
154 * header (which includes the small `fr` params) is needed -- saving the
155 * O(dim^2) per-backend matrix and its O(dim^3) build. Otherwise the full
156 * dense layout is allocated.
157 */
158 static inline uint64_t
159 878 vs_rabitq_params_size(Dimension dim)
160 {
161 878 return vs_fast_rotate_supported(dim) ? (uint64_t)offsetof(RaBitQParams, P)
162
2/2
✓ Branch 0 taken 220 times.
✓ Branch 1 taken 658 times.
878 : VS_RABITQ_PARAMS_DENSE_SIZE(dim);
163 }
164
165 #define VS_RABITQ_PARAMS_SIZE(dim) vs_rabitq_params_size(dim)
166
167 /*
168 * RaBitQScratch - Pre-allocated scratch buffers for encoding
169 *
170 * Avoids per-vector allocation in vs_rabitq_encode_into. Create once
171 * per builder/thread, reuse across all encode calls.
172 */
173 typedef struct RaBitQScratch
174 {
175 float *residual; /* [dim], 64-byte aligned */
176 float *transformed; /* [dim], 64-byte aligned */
177 float *xu_cb; /* [dim], 64-byte aligned */
178 } RaBitQScratch;
179
180 void vs_rabitq_scratch_init(RaBitQScratch *scratch, Dimension dim);
181 void vs_rabitq_scratch_cleanup(RaBitQScratch *scratch);
182
183 /*
184 * RaBitQQueryState - Query state (amortizes work across vectors)
185 *
186 * Precomputes query-specific values that are reused when comparing against
187 * multiple quantized vectors. The transformation and factors are computed
188 * once per query, then used for all distance calculations.
189 *
190 * Distance formula (FAISS-style for better accuracy):
191 * est_dist = g_add + f_add - 2 * f_rescale * final_dot
192 * where:
193 * final_dot = (2 * binary_ip - sum_transformed) * inv_sqrt_d
194 * binary_ip = sum of transformed[i] where bit[i] = 1
195 *
196 * The distance_fn and distance_with_bound_fn pointers are set at
197 * query preparation time based on the selected VsDistanceMode,
198 * enabling zero-branch dispatch in the hot loop.
199 */
200
201 /* Forward declaration for function pointer types */
202 struct RaBitQQueryState;
203
204 typedef Distance (*RaBitQDistanceFn)(
205 const struct RaBitQQueryState *qstate,
206 const RaBitQData *data,
207 Dimension dim);
208 typedef void (*RaBitQDistanceWithBoundFn)(
209 const struct RaBitQQueryState *qstate,
210 const RaBitQData *data,
211 Dimension dim,
212 Distance *est_dist,
213 Distance *lower_bound);
214
215 typedef struct RaBitQQueryState
216 {
217 float *transformed; /* P^T * (query - centroid) */
218 uint8_t *query_bits; /* sign(transformed), packed bits */
219 float g_add; /* ||query - centroid||² */
220 float g_error; /* sqrt(g_add) for error bound */
221 float g_scale; /* mean(|transformed|) for symmetric */
222 float sum_transformed; /* sum(transformed) for distance formula */
223 float inv_sqrt_d; /* 1 / sqrt(dim) */
224 float c_error; /* 2*ε/√(d-1), for deriving f_error */
225 float error_multiplier; /* 1.0 asymmetric, 3.0 symmetric */
226 Dimension dim;
227
228 /* Runtime dispatch (set by prepare_query_ex) */
229 VsDistanceMode mode;
230 RaBitQDistanceFn distance_fn;
231 RaBitQDistanceWithBoundFn distance_with_bound_fn;
232 } RaBitQQueryState;
233
234 /*
235 * Lifecycle - create/destroy quantizer parameters
236 */
237
238 /*
239 * Create new RaBitQ parameters with heap allocation.
240 *
241 * Generates a random orthogonal matrix of size dim x dim using the given
242 * seed. The matrix is created via QR decomposition of a random Gaussian
243 * matrix.
244 *
245 * Returns NULL on allocation failure.
246 */
247 /*
248 * Derive the per-vector error factor f_error from the encoded f_add /
249 * f_rescale (see the RaBitQData comment for the formula). Shared by the
250 * posting and centroid build paths so the formula lives in one place.
251 */
252 float vs_rabitq_derive_f_error(float f_add, float f_rescale, Dimension dim);
253
254 RaBitQParams *vs_rabitq_create(Dimension dim, uint64_t seed);
255
256 /*
257 * Create RaBitQ parameters from an existing rotation matrix.
258 * Copies the matrix into the new allocation.
259 */
260 RaBitQParams *
261 vs_rabitq_create_from_matrix(Dimension dim, uint64_t seed, const float *P);
262
263 /*
264 * Initialize RaBitQ parameters in pre-allocated memory.
265 *
266 * Same as vs_rabitq_create but uses caller-provided buffer.
267 * Buffer must be at least VS_RABITQ_PARAMS_SIZE(dim) bytes.
268 * Generates the rotation matrix P in-place.
269 *
270 * Returns 0 on success, -1 on failure.
271 */
272 int vs_rabitq_init(RaBitQParams *params, Dimension dim, uint64_t seed);
273
274 /*
275 * Free heap-allocated RaBitQ parameters.
276 */
277 void vs_rabitq_destroy(RaBitQParams *params);
278
279 /*
280 * Cleanup internal resources (no-op since P is now inline).
281 * Kept for API compatibility with stack-allocated usage.
282 */
283 void vs_rabitq_cleanup(RaBitQParams *params);
284
285 /*
286 * Encoding - convert full-precision vectors to binary codes
287 */
288
289 /*
290 * Encode a vector to compact RaBitQ format with heap allocation.
291 *
292 * Computes residual from centroid, transforms through P^T, extracts
293 * sign bits, and computes estimation factors.
294 *
295 * Returns NULL on failure.
296 */
297 RaBitQData *vs_rabitq_encode(
298 const RaBitQParams *params, Vec32Ref input, Vec32Ref centroid);
299
300 /*
301 * Encode a vector into pre-allocated output buffer.
302 *
303 * Output buffer must be at least VS_RABITQ_DATA_SIZE(dim) bytes.
304 *
305 * Returns 0 on success, -1 on failure.
306 */
307 int vs_rabitq_encode_into(
308 const RaBitQParams *params,
309 Vec32Ref input,
310 Vec32Ref centroid,
311 RaBitQData *output);
312
313 /*
314 * Encode with pre-allocated scratch buffers (zero per-call allocation).
315 */
316 int vs_rabitq_encode_into_ex(
317 const RaBitQParams *params,
318 Vec32Ref input,
319 Vec32Ref centroid,
320 RaBitQData *output,
321 RaBitQScratch *scratch);
322
323 /*
324 * Encode from an already-rotated residual: pt_residual = P^T * (input -
325 * centroid). The caller supplies the rotated residual directly, so this skips
326 * the residual subtraction and P^T multiply that encode_into_ex performs; it
327 * runs only the sign-extract + factor math. Used by the runtime insert path,
328 * which rotates the inserted vector once and subtracts the posting head's
329 * stored pt_centroid (P^T is linear: P^T*(v-c) = P^T*v - P^T*c), avoiding any
330 * dependency on the raw leaf centroid (unavailable for RABITQ/FASTSCAN
331 * centroid formats). Only scratch->xu_cb is used.
332 *
333 * pt_residual must be [params->dim] floats. Returns 0 on success, -1 on
334 * failure.
335 */
336 int vs_rabitq_encode_from_pt(
337 const RaBitQParams *params,
338 const float *pt_residual,
339 RaBitQData *output,
340 RaBitQScratch *scratch);
341
342 /*
343 * Batch encode multiple vectors into separate output arrays.
344 *
345 * More efficient than calling vs_rabitq_encode_into() repeatedly because:
346 * 1. Matrix P is loaded into cache once and reused for all vectors
347 * 2. Centroid rotation is computed once and reused
348 * 3. Batched matrix-vector multiplication enables better SIMD utilization
349 *
350 * Memory layout:
351 * vectors: count vectors, each dim elements, contiguous
352 * vec_type: element type (VS_VEC_F32, VS_VEC_F16, VS_VEC_F16C)
353 * f_add: count floats (output)
354 * f_rescale: count floats (output)
355 * bits: count * packed_bytes bytes (output)
356 *
357 * For non-f32 input, vectors are converted to float32 at entry.
358 * This is O(count × dim), negligible vs the O(count × dim²) rotation.
359 *
360 * Returns 0 on success, -1 on failure.
361 */
362 int vs_rabitq_encode_batch(
363 const RaBitQParams *params,
364 const void *vectors,
365 VecType vec_type,
366 Vec32Ref centroid,
367 float *f_add,
368 float *f_rescale,
369 uint8_t *bits,
370 uint16_t count);
371
372 /*
373 * Batch encode with heap-allocated RaBitQBatch output.
374 *
375 * Convenience wrapper that allocates a RaBitQBatch and calls
376 * vs_rabitq_encode_batch() with the batch's arrays.
377 *
378 * Returns NULL on failure.
379 */
380 RaBitQBatch *vs_rabitq_encode_batch_alloc(
381 const RaBitQParams *params,
382 const void *vectors,
383 VecType vec_type,
384 Vec32Ref centroid,
385 uint16_t count);
386
387 /*
388 * Create an empty RaBitQBatch with allocated arrays.
389 *
390 * Returns NULL on allocation failure.
391 */
392 RaBitQBatch *vs_rabitq_batch_create(uint16_t count, Dimension dim);
393
394 /*
395 * Free a RaBitQBatch and its arrays.
396 */
397 void vs_rabitq_batch_destroy(RaBitQBatch *batch);
398
399 /*
400 * Query preparation - precompute query-specific factors
401 */
402
403 /*
404 * Prepare query state for efficient distance computation.
405 *
406 * Transforms the query through P^T and precomputes factors that are
407 * reused when comparing against multiple quantized vectors.
408 *
409 * Returns NULL on failure.
410 */
411 RaBitQQueryState *vs_rabitq_prepare_query(
412 const RaBitQParams *params, Vec32Ref query, Vec32Ref centroid);
413
414 /*
415 * Free query state.
416 */
417 void vs_rabitq_free_query(RaBitQQueryState *state);
418
419 /*
420 * Prepare query state with explicit distance mode.
421 *
422 * Like vs_rabitq_prepare_query() but additionally sets function pointers
423 * for the selected mode, enabling zero-branch dispatch via the inline
424 * helpers below.
425 *
426 * Returns NULL on failure.
427 */
428 RaBitQQueryState *vs_rabitq_prepare_query_ex(
429 const RaBitQParams *params,
430 Vec32Ref query,
431 Vec32Ref centroid,
432 VsDistanceMode mode);
433
434 /*
435 * Pre-rotation API — eliminates per-cluster matrix multiply
436 *
437 * Instead of calling vs_rabitq_prepare_query_ex() per cluster
438 * (which does O(dim²) matrix multiply each time), precompute:
439 * - P^T * centroid at index build time (once per cluster)
440 * - P^T * query at query time (once per query)
441 * Then per cluster: transformed = pt_query - pt_centroid (O(dim))
442 *
443 * This is valid because P^T is linear:
444 * P^T * (query - centroid) = P^T * query - P^T * centroid
445 */
446
447 /*
448 * Rotate a vector through P^T into pre-allocated output buffer.
449 * Output must have space for dim floats, 64-byte aligned preferred.
450 */
451 void vs_rabitq_rotate(
452 const RaBitQParams *params, const float *input, float *output);
453
454 /*
455 * Initialize a pre-allocated query state from already-rotated vectors.
456 *
457 * Computes transformed = pt_query - pt_centroid (vector subtraction),
458 * then derives all scalar fields (g_add, g_error, query_bits, etc.).
459 * No matrix multiply — O(dim) instead of O(dim²).
460 *
461 * state->transformed and state->query_bits must be pre-allocated by
462 * the caller (dim floats and packed_bytes bytes respectively).
463 *
464 * Call vs_rabitq_init_query_constants() once at context creation to
465 * set dim-dependent constants (inv_sqrt_d, c_error) that don't change
466 * per cluster.
467 */
468 void vs_rabitq_init_query_constants(RaBitQQueryState *state, Dimension dim);
469
470 void vs_rabitq_init_query_state(
471 RaBitQQueryState *state,
472 const float *pt_query,
473 const float *pt_centroid,
474 Dimension dim,
475 VsDistanceMode mode);
476
477 /*
478 * Dispatch helpers - call through function pointers set at prepare time
479 */
480
481 static inline Distance
482 4 vs_rabitq_distance_dispatch(
483 const RaBitQQueryState *qstate, const RaBitQData *data, Dimension dim)
484 {
485 4 return qstate->distance_fn(qstate, data, dim);
486 }
487
488 static inline void
489 4 vs_rabitq_distance_dispatch_with_bound(
490 const RaBitQQueryState *qstate,
491 const RaBitQData *data,
492 Dimension dim,
493 Distance *est_dist,
494 Distance *lower_bound)
495 {
496 4 qstate->distance_with_bound_fn(qstate, data, dim, est_dist, lower_bound);
497 4 }
498
499 /*
500 * Distance computation - estimate L2 distance from quantized codes
501 */
502
503 /*
504 * Compute estimated L2 squared distance from compact data.
505 *
506 * Uses the precomputed query state and compact vector factors:
507 * est_dist = g_add + f_add - 2 * f_rescale * final_dot
508 */
509 Distance vs_rabitq_distance(
510 const RaBitQQueryState *query_state,
511 const RaBitQData *data,
512 Dimension dim);
513
514 /*
515 * Compute estimated distance with derived error bound from compact data.
516 *
517 * f_error is derived from f_add and f_rescale using c_error in query
518 * state. Returns both estimated distance and lower bound guaranteed
519 * to be <= true distance.
520 */
521 void vs_rabitq_distance_with_bound(
522 const RaBitQQueryState *query_state,
523 const RaBitQData *data,
524 Dimension dim,
525 Distance *est_dist,
526 Distance *lower_bound);
527
528 /*
529 * Batch distance on separate arrays
530 *
531 * Computes estimated L2 distances for 'count' vectors whose fields
532 * are stored in separate contiguous arrays: f_add[], f_rescale[],
533 * and bits[]. Used by benchmarks and the multi-candidate kernel.
534 *
535 * The scalar arithmetic on contiguous f_add[]/f_rescale[] arrays
536 * auto-vectorizes with the compiler.
537 */
538 void vs_rabitq_distance_batch(
539 const RaBitQQueryState *qstate,
540 const float *f_add,
541 const float *f_rescale,
542 const uint8_t *bits,
543 uint32_t count,
544 Dimension dim,
545 Distance *distances);
546
547 /*
548 * Internal SIMD dispatch (called automatically)
549 */
550
551 /*
552 * Initialize SIMD dispatch for RaBitQ operations.
553 * Called automatically on first use, but can be called explicitly
554 * for deterministic initialization timing.
555 *
556 * Returns 0 on success.
557 */
558 int vs_rabitq_init_simd(void);
559
560 /*
561 * Get name of the active SIMD implementation.
562 * Returns one of: "avx512", "avx2", "neon", "compiler"
563 */
564 const char *vs_rabitq_impl_name(void);
565
566 /*
567 * Force re-initialization of SIMD dispatch (for testing).
568 */
569 void vs_rabitq_force_reinit(void);
570
571 /*
572 * Hamming distance - XOR + popcount between two bit vectors
573 */
574
575 /*
576 * Compute Hamming distance between two packed bit vectors.
577 *
578 * Returns the number of bit positions where a and b differ.
579 */
580 uint32_t vs_rabitq_hamming_distance(
581 const uint8_t *a, const uint8_t *b, uint32_t packed_bytes);
582
583 /*
584 * Multi-candidate Hamming distance (vertical SIMD).
585 *
586 * Computes Hamming distances from query_bits to count data vectors.
587 * Candidate i's bits start at data_bits + i * stride.
588 */
589 void vs_rabitq_hamming_distance_multi(
590 const uint8_t *query_bits,
591 const uint8_t *data_bits,
592 uint32_t stride,
593 uint32_t packed_bytes,
594 uint32_t count,
595 uint32_t *results);
596
597 /*
598 * Symmetric distance - both query and data are 1-bit quantized
599 *
600 * Uses Hamming distance (XOR + popcount) instead of asymmetric inner
601 * product (mask + add). ~32x fewer iterations of the inner loop,
602 * but with additional query quantization error.
603 *
604 * Formula:
605 * sym_dot = dim - 2 * hamming(query_bits, data_bits)
606 * final_dot = sym_dot * inv_sqrt_d
607 * est_dist = f_add + g_add - 2 * f_rescale * g_scale * final_dot
608 */
609
610 Distance vs_rabitq_distance_symmetric(
611 const RaBitQQueryState *qstate, const RaBitQData *data, Dimension dim);
612
613 void vs_rabitq_distance_symmetric_with_bound(
614 const RaBitQQueryState *qstate,
615 const RaBitQData *data,
616 Dimension dim,
617 Distance *est_dist,
618 Distance *lower_bound);
619
620 void vs_rabitq_distance_batch_symmetric(
621 const RaBitQQueryState *qstate,
622 const float *f_add,
623 const float *f_rescale,
624 const uint8_t *bits,
625 uint32_t count,
626 Dimension dim,
627 Distance *distances);
628
629 /*
630 * Multi-candidate inner product (vertical SIMD)
631 *
632 * Computes inner products for multiple candidates in a single pass over
633 * transformed[]. Loads transformed[] once per dimension chunk and
634 * processes N candidates simultaneously.
635 *
636 * Candidate i's bits start at bits + i * stride.
637 */
638 typedef void (*InnerProductMultiFn)(
639 const float *transformed,
640 const uint8_t *bits,
641 uint32_t stride,
642 Dimension dim,
643 uint32_t count,
644 float *results);
645
646 void vs_rabitq_inner_product_multi(
647 const float *transformed,
648 const uint8_t *bits,
649 uint32_t stride,
650 Dimension dim,
651 uint32_t count,
652 float *results);
653
654 /*
655 * Batch distance using multi-candidate inner product
656 *
657 * Same interface as vs_rabitq_distance_batch() but uses the
658 * vertical SIMD inner product to process multiple candidates per
659 * pass over transformed[].
660 */
661 void vs_rabitq_distance_batch_multi(
662 const RaBitQQueryState *qstate,
663 const float *f_add,
664 const float *f_rescale,
665 const uint8_t *bits,
666 uint32_t count,
667 Dimension dim,
668 Distance *distances);
669
670 /*
671 * Batch distance with error bounds (multi-candidate inner product)
672 *
673 * Combines vs_rabitq_inner_product_multi with per-entry error bound
674 * derivation. Bits are accessed via stride (not packed_bytes), allowing
675 * direct use on interleaved page data where stride = data_size.
676 *
677 * scratch: caller-provided buffer of at least count floats, reusable
678 * across calls to avoid per-call allocation.
679 */
680 void vs_rabitq_distance_batch_multi_with_bound(
681 const RaBitQQueryState *qstate,
682 const float *f_add,
683 const float *f_rescale,
684 const uint8_t *bits,
685 uint32_t stride,
686 uint32_t count,
687 Dimension dim,
688 Distance *distances,
689 Distance *lower_bounds,
690 float *scratch);
691
692 /*
693 * Batch symmetric distance with error bounds
694 *
695 * Combines vs_rabitq_hamming_distance_multi with per-entry error
696 * bound derivation. Like the asymmetric variant, bits are accessed
697 * via stride for direct use on interleaved page data.
698 *
699 * scratch: caller-provided buffer of at least count uint32_t's,
700 * reusable across calls to avoid per-call allocation.
701 */
702 void vs_rabitq_distance_batch_symmetric_with_bound(
703 const RaBitQQueryState *qstate,
704 const float *f_add,
705 const float *f_rescale,
706 const uint8_t *bits,
707 uint32_t stride,
708 uint32_t count,
709 Dimension dim,
710 Distance *distances,
711 Distance *lower_bounds,
712 uint32_t *scratch);
713
714 /*
715 * Get name of the active Hamming SIMD implementation.
716 * Returns one of: "avx512-vpopcntdq", "avx2", "compiler"
717 */
718 const char *vs_rabitq_hamming_impl_name(void);
719
720 /*
721 * Hand-optimized SIMD implementations (simd=full only)
722 *
723 * These are resolved via function pointers in vs_rabitq_init_simd().
724 */
725 #ifdef VS_SIMD_FULL
726
727 #if defined(__x86_64__) || defined(_M_X64)
728 /* AVX-512 implementations */
729 float vs_rabitq_inner_product_avx512(
730 const float *transformed, const uint8_t *bits, Dimension dim);
731 void vs_rabitq_extract_signs_avx512(
732 const float *transformed, uint8_t *bits, Dimension dim);
733 void vs_rabitq_inner_product_multi_avx512(
734 const float *transformed,
735 const uint8_t *bits,
736 uint32_t stride,
737 Dimension dim,
738 uint32_t count,
739 float *results);
740
741 /* AVX-512 VPOPCNTDQ Hamming implementations */
742 uint32_t vs_rabitq_hamming_avx512(
743 const uint8_t *a, const uint8_t *b, uint32_t packed_bytes);
744 void vs_rabitq_hamming_multi_avx512(
745 const uint8_t *query_bits,
746 const uint8_t *data_bits,
747 uint32_t stride,
748 uint32_t packed_bytes,
749 uint32_t count,
750 uint32_t *results);
751
752 /* AVX2 implementations */
753 float vs_rabitq_inner_product_avx2(
754 const float *transformed, const uint8_t *bits, Dimension dim);
755 void vs_rabitq_extract_signs_avx2(
756 const float *transformed, uint8_t *bits, Dimension dim);
757 void vs_rabitq_inner_product_multi_avx2(
758 const float *transformed,
759 const uint8_t *bits,
760 uint32_t stride,
761 Dimension dim,
762 uint32_t count,
763 float *results);
764
765 /* AVX2 lookup-table Hamming implementations */
766 uint32_t vs_rabitq_hamming_avx2(
767 const uint8_t *a, const uint8_t *b, uint32_t packed_bytes);
768 void vs_rabitq_hamming_multi_avx2(
769 const uint8_t *query_bits,
770 const uint8_t *data_bits,
771 uint32_t stride,
772 uint32_t packed_bytes,
773 uint32_t count,
774 uint32_t *results);
775 #endif
776
777 #if defined(__aarch64__) || defined(_M_ARM64)
778 /* NEON implementations */
779 float vs_rabitq_inner_product_neon(
780 const float *transformed, const uint8_t *bits, Dimension dim);
781 void vs_rabitq_extract_signs_neon(
782 const float *transformed, uint8_t *bits, Dimension dim);
783 void vs_rabitq_inner_product_multi_neon(
784 const float *transformed,
785 const uint8_t *bits,
786 uint32_t stride,
787 Dimension dim,
788 uint32_t count,
789 float *results);
790 #endif
791
792 #endif /* VS_SIMD_FULL */
793
794 #endif /* VS_RABITQ_H */
795