| Line | Branch | Exec | Source |
|---|---|---|---|
| 1 | /* | ||
| 2 | * Copyright (c) 2026 Tiger Data, Inc. | ||
| 3 | * Licensed under the PostgreSQL License. See LICENSE for details. | ||
| 4 | * | ||
| 5 | * fastscan.h - VPSHUFB-based fast scan for RaBitQ | ||
| 6 | * | ||
| 7 | * Repacks 1-bit RaBitQ sign codes into 4-bit nibble layout for | ||
| 8 | * VPSHUFB table-lookup distance computation. Processes 32 vectors | ||
| 9 | * per batch using the even/odd byte accumulation trick. | ||
| 10 | * | ||
| 11 | * Based on the approach from the RaBitQ Library (Gao & Long) and | ||
| 12 | * FAISS FastScan (Andre et al., "Cache locality is not enough"). | ||
| 13 | * | ||
| 14 | * Data layout per batch of 32 vectors: | ||
| 15 | * - Packed codes: dim/8 columns x 32 bytes, with kPerm0 | ||
| 16 | * vector interleaving for even/odd byte accumulation | ||
| 17 | * - LUT: dim/4 subquantizers x 16 uint8 entries | ||
| 18 | * | ||
| 19 | * Query LUT: | ||
| 20 | * - 16-entry uint8 table per subquantizer, precomputing | ||
| 21 | * partial inner products for all 16 sign combinations | ||
| 22 | * - Quantized from float with linear scale + bias | ||
| 23 | * - De-quantized after accumulation: ip = accum * delta + bias | ||
| 24 | */ | ||
| 25 | |||
| 26 | #ifndef VS_FASTSCAN_H | ||
| 27 | #define VS_FASTSCAN_H | ||
| 28 | |||
| 29 | #include <stdint.h> | ||
| 30 | |||
| 31 | #include "core/types.h" | ||
| 32 | |||
| 33 | /* Minimum LUT range to avoid division by near-zero */ | ||
| 34 | #define VS_FASTSCAN_MIN_RANGE 1e-10f | ||
| 35 | |||
| 36 | /* Dimensions per subquantizer */ | ||
| 37 | #define VS_FASTSCAN_SQ_DIM 4 | ||
| 38 | |||
| 39 | /* Vectors per VPSHUFB batch */ | ||
| 40 | #define VS_FASTSCAN_GROUP 32 | ||
| 41 | |||
| 42 | /* Number of subquantizers for a given dimension */ | ||
| 43 | #define VS_FASTSCAN_NSQ(dim) \ | ||
| 44 | (((dim) + VS_FASTSCAN_SQ_DIM - 1) / VS_FASTSCAN_SQ_DIM) | ||
| 45 | |||
| 46 | /* Number of subquantizer pairs (two nibbles per byte) */ | ||
| 47 | #define VS_FASTSCAN_NSQ_PAIRS(dim) ((VS_FASTSCAN_NSQ(dim) + 1) / 2) | ||
| 48 | |||
| 49 | /* Bytes per 32-vector group in packed fastscan layout. | ||
| 50 | * Each column (8 dims = 2 sq) uses 32 bytes. */ | ||
| 51 | #define VS_FASTSCAN_GROUP_BYTES(dim) \ | ||
| 52 | ((uint32_t)(((dim) + 7) / 8) * VS_FASTSCAN_GROUP) | ||
| 53 | |||
| 54 | /* Bytes for the query lookup table. | ||
| 55 | * nsq_pairs * 2 to include phantom sq when nsq is odd. */ | ||
| 56 | #define VS_FASTSCAN_LUT_BYTES(dim) \ | ||
| 57 | ((uint32_t)VS_FASTSCAN_NSQ_PAIRS(dim) * 2 * 16) | ||
| 58 | |||
| 59 | /* | ||
| 60 | * Build a uint8 lookup table per subquantizer from the | ||
| 61 | * transformed query vector. | ||
| 62 | * | ||
| 63 | * lut_out: output buffer, VS_FASTSCAN_LUT_BYTES(dim) bytes | ||
| 64 | * delta_out: quantization step size for de-quantization | ||
| 65 | * bias_out: LUT min value (sum_vl = vl * nsq) | ||
| 66 | */ | ||
| 67 | void vs_fastscan_build_lut( | ||
| 68 | const float *transformed, | ||
| 69 | Dimension dim, | ||
| 70 | uint8_t *lut_out, | ||
| 71 | float *delta_out, | ||
| 72 | float *bias_out); | ||
| 73 | |||
| 74 | /* | ||
| 75 | * Repack 1-bit RaBitQ codes into fastscan layout with kPerm0 | ||
| 76 | * vector interleaving for even/odd byte accumulation. | ||
| 77 | * | ||
| 78 | * Input: bits_1bit[count * packed_bytes] | ||
| 79 | * Output: codes_out[ngroups * group_bytes] | ||
| 80 | * | ||
| 81 | * Returns the number of 32-vector groups. | ||
| 82 | */ | ||
| 83 | uint32_t vs_fastscan_pack_codes( | ||
| 84 | const uint8_t *bits_1bit, | ||
| 85 | uint32_t count, | ||
| 86 | Dimension dim, | ||
| 87 | uint8_t *codes_out); | ||
| 88 | |||
| 89 | /* | ||
| 90 | * Inverse of vs_fastscan_pack_codes: reconstruct per-vector 1-bit codes | ||
| 91 | * from the packed fastscan layout. bits_out[count * packed_bytes]. | ||
| 92 | */ | ||
| 93 | void vs_fastscan_unpack_codes( | ||
| 94 | const uint8_t *codes, | ||
| 95 | uint32_t count, | ||
| 96 | Dimension dim, | ||
| 97 | uint8_t *bits_out); | ||
| 98 | |||
| 99 | /* | ||
| 100 | * Required output buffer size for packed fastscan codes. | ||
| 101 | */ | ||
| 102 | static inline uint32_t | ||
| 103 | 26 | vs_fastscan_codes_size(uint32_t count, Dimension dim) | |
| 104 | { | ||
| 105 | 26 | uint32_t ngroups = (count + VS_FASTSCAN_GROUP - 1) / VS_FASTSCAN_GROUP; | |
| 106 | 26 | uint32_t group_bytes = VS_FASTSCAN_GROUP_BYTES(dim); | |
| 107 | 26 | return ngroups * group_bytes; | |
| 108 | } | ||
| 109 | |||
| 110 | /* | ||
| 111 | * VPSHUFB accumulate kernel -- process one 32-vector group. | ||
| 112 | * | ||
| 113 | * Uses even/odd byte accumulation: 4 uint16 accumulators | ||
| 114 | * track vectors {0-7, 8-15, 16-23, 24-31} via kPerm0 | ||
| 115 | * interleaving. No explicit uint8->uint16 widening in the | ||
| 116 | * inner loop. | ||
| 117 | * | ||
| 118 | * codes: packed codes for one group | ||
| 119 | * lut: uint8 lookup table | ||
| 120 | * accum: output, 32 uint16 accumulated distances | ||
| 121 | * dim: vector dimension (determines loop count) | ||
| 122 | */ | ||
| 123 | void vs_fastscan_accumulate( | ||
| 124 | const uint8_t *codes, | ||
| 125 | const uint8_t *lut, | ||
| 126 | uint16_t *accum, | ||
| 127 | Dimension dim); | ||
| 128 | |||
| 129 | /* ---------------------------------------------------------------- | ||
| 130 | * High-accuracy mode (uint16 LUT → int32 accumulators) | ||
| 131 | * | ||
| 132 | * Splits each uint16 LUT entry into lo/hi byte tables and does | ||
| 133 | * two VPSHUFB passes per code block. ~2× kernel cost but | ||
| 134 | * eliminates uint8 quantization error → matches float recall. | ||
| 135 | * | ||
| 136 | * LUT layout: [lo_64B, hi_64B] interleaved per 4 subquantizers. | ||
| 137 | * Total LUT size: 2 × nsq × 16 bytes. | ||
| 138 | * ---------------------------------------------------------------- */ | ||
| 139 | |||
| 140 | /* | ||
| 141 | * LUT bytes for high-accuracy (16-bit) mode. | ||
| 142 | * | ||
| 143 | * Entries are split into lo/hi bytes laid out in 128-byte blocks, one block | ||
| 144 | * per group of 4 subquantizers: [4×16 lo][4×16 hi]. The build and accumulate | ||
| 145 | * kernels address it as lut + (sq/4)*128 + (sq%4)*16 with the hi half at +64, | ||
| 146 | * so the buffer must cover ceil(nsq/4) whole 128-byte blocks. Sizing it from | ||
| 147 | * NSQ_PAIRS (ceil(nsq/2)*64) under-allocates by a block whenever nsq mod 4 is | ||
| 148 | * 1 or 2 (e.g. dim ≤ 4 → nsq = 1), letting the hi-half store run 64 bytes past | ||
| 149 | * the end and corrupt the following allocation. | ||
| 150 | */ | ||
| 151 | #define VS_FASTSCAN_LUT_HACC_BYTES(dim) \ | ||
| 152 | ((uint32_t)(((VS_FASTSCAN_NSQ(dim) + 3) / 4) * 128)) | ||
| 153 | |||
| 154 | /* | ||
| 155 | * Build a uint16-precision LUT, split into lo/hi byte tables. | ||
| 156 | * Same interface as build_lut but lut_out must be | ||
| 157 | * VS_FASTSCAN_LUT_HACC_BYTES(dim) bytes. | ||
| 158 | */ | ||
| 159 | void vs_fastscan_build_lut_hacc( | ||
| 160 | const float *transformed, | ||
| 161 | Dimension dim, | ||
| 162 | uint8_t *lut_out, | ||
| 163 | float *delta_out, | ||
| 164 | float *bias_out); | ||
| 165 | |||
| 166 | /* | ||
| 167 | * High-accuracy accumulate: two VPSHUFB passes per code block. | ||
| 168 | * Output is 32 int32 values (not uint16). | ||
| 169 | */ | ||
| 170 | /* | ||
| 171 | * Function-pointer type of the 16-bit (hacc) accumulate kernel, and a | ||
| 172 | * resolver that returns the SIMD variant selected for this CPU | ||
| 173 | * (initializing dispatch if needed). Hot per-group callers cache the | ||
| 174 | * pointer per scan instead of paying the guarded wrapper's atomic load | ||
| 175 | * and double indirection on every group. | ||
| 176 | */ | ||
| 177 | typedef void (*VsFastscanAccumulateHaccFn)( | ||
| 178 | const uint8_t *codes, | ||
| 179 | const uint8_t *lut, | ||
| 180 | int32_t *accum, | ||
| 181 | Dimension dim); | ||
| 182 | |||
| 183 | VsFastscanAccumulateHaccFn vs_fastscan_get_accumulate_hacc(void); | ||
| 184 | |||
| 185 | void vs_fastscan_accumulate_hacc( | ||
| 186 | const uint8_t *codes, | ||
| 187 | const uint8_t *lut, | ||
| 188 | int32_t *accum, | ||
| 189 | Dimension dim); | ||
| 190 | |||
| 191 | /* | ||
| 192 | * Batch distance computation using fastscan. | ||
| 193 | */ | ||
| 194 | struct RaBitQQueryState; | ||
| 195 | |||
| 196 | void vs_fastscan_distance_batch( | ||
| 197 | const struct RaBitQQueryState *qstate, | ||
| 198 | const float *f_add, | ||
| 199 | const float *f_rescale, | ||
| 200 | const uint8_t *codes, | ||
| 201 | uint32_t ngroups, | ||
| 202 | uint32_t count, | ||
| 203 | Dimension dim, | ||
| 204 | float *distances, | ||
| 205 | uint8_t *lut_buf, | ||
| 206 | uint16_t *accum_buf); | ||
| 207 | |||
| 208 | void vs_fastscan_init_simd(void); | ||
| 209 | void vs_fastscan_reset_simd(void); | ||
| 210 | const char *vs_fastscan_impl_name(void); | ||
| 211 | |||
| 212 | /* SIMD implementations */ | ||
| 213 | #ifdef VS_SIMD_FULL | ||
| 214 | #if defined(__x86_64__) || defined(_M_X64) | ||
| 215 | |||
| 216 | void vs_fastscan_accumulate_avx2( | ||
| 217 | const uint8_t *codes, | ||
| 218 | const uint8_t *lut, | ||
| 219 | uint16_t *accum, | ||
| 220 | Dimension dim); | ||
| 221 | |||
| 222 | void vs_fastscan_accumulate_hacc_avx2( | ||
| 223 | const uint8_t *codes, | ||
| 224 | const uint8_t *lut, | ||
| 225 | int32_t *accum, | ||
| 226 | Dimension dim); | ||
| 227 | |||
| 228 | void vs_fastscan_accumulate_avx512( | ||
| 229 | const uint8_t *codes, | ||
| 230 | const uint8_t *lut, | ||
| 231 | uint16_t *accum, | ||
| 232 | Dimension dim); | ||
| 233 | |||
| 234 | void vs_fastscan_build_lut_avx512( | ||
| 235 | const float *transformed, | ||
| 236 | Dimension dim, | ||
| 237 | uint8_t *lut_out, | ||
| 238 | float *delta_out, | ||
| 239 | float *bias_out); | ||
| 240 | |||
| 241 | void vs_fastscan_accumulate_hacc_avx512( | ||
| 242 | const uint8_t *codes, | ||
| 243 | const uint8_t *lut, | ||
| 244 | int32_t *accum, | ||
| 245 | Dimension dim); | ||
| 246 | |||
| 247 | void vs_fastscan_build_lut_hacc_avx512( | ||
| 248 | const float *transformed, | ||
| 249 | Dimension dim, | ||
| 250 | uint8_t *lut_out, | ||
| 251 | float *delta_out, | ||
| 252 | float *bias_out); | ||
| 253 | |||
| 254 | #elif defined(__aarch64__) || defined(_M_ARM64) | ||
| 255 | |||
| 256 | void vs_fastscan_accumulate_neon( | ||
| 257 | const uint8_t *codes, | ||
| 258 | const uint8_t *lut, | ||
| 259 | uint16_t *accum, | ||
| 260 | Dimension dim); | ||
| 261 | |||
| 262 | void vs_fastscan_accumulate_hacc_neon( | ||
| 263 | const uint8_t *codes, | ||
| 264 | const uint8_t *lut, | ||
| 265 | int32_t *accum, | ||
| 266 | Dimension dim); | ||
| 267 | |||
| 268 | #endif | ||
| 269 | #endif | ||
| 270 | |||
| 271 | #endif /* VS_FASTSCAN_H */ | ||
| 272 |