| 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 |