| 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_avx2.c - AVX2 optimized binary inner product for RaBitQ | ||
| 6 | * | ||
| 7 | * The hot path in RaBitQ distance computation is: | ||
| 8 | * sum = 0 | ||
| 9 | * for each bit i: | ||
| 10 | * if bits[i] == 1: | ||
| 11 | * sum += transformed[i] | ||
| 12 | * | ||
| 13 | * This AVX2 implementation processes 8 floats at a time using masked | ||
| 14 | * addition. The key optimization is expanding each byte of the bit array | ||
| 15 | * into 8 float masks. | ||
| 16 | */ | ||
| 17 | |||
| 18 | #include "vs_config.h" | ||
| 19 | |||
| 20 | #ifdef VS_SIMD_FULL | ||
| 21 | |||
| 22 | #if defined(__x86_64__) || defined(_M_X64) | ||
| 23 | |||
| 24 | #include <immintrin.h> | ||
| 25 | |||
| 26 | #include "algo/simd_utils.h" | ||
| 27 | #include "core/types.h" | ||
| 28 | #include "quant/rabitq.h" | ||
| 29 | |||
| 30 | /* | ||
| 31 | * Expand a byte to 8 float masks for AVX2. | ||
| 32 | * | ||
| 33 | * Input: byte with 8 bits (LSB-first) | ||
| 34 | * Output: __m256 where each lane is either all 1s (0xFFFFFFFF) or all 0s | ||
| 35 | * | ||
| 36 | * The bit order is LSB-first (bit 0 -> float 0, bit 7 -> float 7). | ||
| 37 | */ | ||
| 38 | VS_TARGET_AVX2 static inline __m256 | ||
| 39 | 47346951 | expand_byte_to_mask_avx2(uint8_t byte) | |
| 40 | { | ||
| 41 | /* Broadcast byte to all positions */ | ||
| 42 | 89990955 | __m256i vbyte = _mm256_set1_epi32(byte); | |
| 43 | |||
| 44 | /* Bit positions for each float (LSB-first: 0,1,2,3,4,5,6,7) */ | ||
| 45 | 47346951 | __m256i bit_positions = _mm256_setr_epi32( | |
| 46 | 1 << 0, 1 << 1, 1 << 2, 1 << 3, 1 << 4, 1 << 5, 1 << 6, 1 << 7); | ||
| 47 | |||
| 48 | /* Test each bit: (byte & bit_pos) != 0 */ | ||
| 49 | 47346951 | __m256i masked = _mm256_and_si256(vbyte, bit_positions); | |
| 50 | |||
| 51 | /* Compare to zero - creates all 1s for non-zero, all 0s for zero */ | ||
| 52 | 89990955 | __m256i cmp = _mm256_cmpeq_epi32(masked, _mm256_setzero_si256()); | |
| 53 | |||
| 54 | /* Invert: we want 1s where bit is set */ | ||
| 55 | 89990955 | __m256i result = _mm256_xor_si256(cmp, _mm256_set1_epi32(-1)); | |
| 56 | |||
| 57 | 47346951 | return _mm256_castsi256_ps(result); | |
| 58 | } | ||
| 59 | |||
| 60 | /* | ||
| 61 | * AVX2 binary inner product implementation. | ||
| 62 | * | ||
| 63 | * Processes 8 floats at a time using masked addition. | ||
| 64 | */ | ||
| 65 | VS_TARGET_AVX2 float | ||
| 66 | 546027 | vs_rabitq_inner_product_avx2( | |
| 67 | const float *transformed, const uint8_t *bits, Dimension dim) | ||
| 68 | { | ||
| 69 | 546027 | __m256 sum = _mm256_setzero_ps(); | |
| 70 | |||
| 71 | /* Main loop: process 8 floats (1 byte of bits) at a time */ | ||
| 72 | 546027 | Dimension i = 0; | |
| 73 |
2/2✓ Branch 0 taken 2661132 times.
✓ Branch 1 taken 546027 times.
|
3207159 | for (; i + 8 <= dim; i += 8) |
| 74 | { | ||
| 75 | /* Load 8 transformed values */ | ||
| 76 | 2661132 | __m256 t = _mm256_loadu_ps(transformed + i); | |
| 77 | |||
| 78 | /* Get the byte containing bits for these 8 floats */ | ||
| 79 | 2661132 | uint8_t byte = bits[i / 8]; | |
| 80 | |||
| 81 | /* Expand byte to mask */ | ||
| 82 | 2661132 | __m256 mask = expand_byte_to_mask_avx2(byte); | |
| 83 | |||
| 84 | /* Masked addition: only add where bit is set */ | ||
| 85 | 2661132 | __m256 masked_t = _mm256_and_ps(t, mask); | |
| 86 | 2661132 | sum = _mm256_add_ps(sum, masked_t); | |
| 87 | } | ||
| 88 | |||
| 89 | /* Horizontal sum */ | ||
| 90 | 546027 | float result = vs_horizontal_sum_avx2(sum); | |
| 91 | |||
| 92 | /* Handle tail elements (LSB-first) */ | ||
| 93 |
2/2✓ Branch 0 taken 55089 times.
✓ Branch 1 taken 546027 times.
|
601116 | for (; i < dim; i++) |
| 94 | { | ||
| 95 | 55089 | int byte_idx = i / 8; | |
| 96 | 55089 | int bit_idx = i % 8; | |
| 97 | 55089 | int bit = (bits[byte_idx] >> bit_idx) & 1; | |
| 98 |
2/2✓ Branch 0 taken 16319 times.
✓ Branch 1 taken 38770 times.
|
55089 | if (bit) |
| 99 | 16319 | result += transformed[i]; | |
| 100 | } | ||
| 101 | |||
| 102 | 546027 | return result; | |
| 103 | } | ||
| 104 | |||
| 105 | /* | ||
| 106 | * AVX2 sign extraction for RaBitQ encoding. | ||
| 107 | * | ||
| 108 | * Extracts sign bits from transformed floats into packed bytes. | ||
| 109 | * Uses AVX2 compare and movemask to generate 8 bits at a time. | ||
| 110 | */ | ||
| 111 | VS_TARGET_AVX2 void | ||
| 112 | 685136 | vs_rabitq_extract_signs_avx2( | |
| 113 | const float *transformed, uint8_t *bits, Dimension dim) | ||
| 114 | { | ||
| 115 | 685136 | __m256 zero = _mm256_setzero_ps(); | |
| 116 | |||
| 117 | /* Main loop: process 8 floats -> 1 byte at a time */ | ||
| 118 | 685136 | Dimension i = 0; | |
| 119 |
2/2✓ Branch 0 taken 4965014 times.
✓ Branch 1 taken 685136 times.
|
5650150 | for (; i + 8 <= dim; i += 8) |
| 120 | { | ||
| 121 | /* Load 8 floats */ | ||
| 122 | 4965014 | __m256 t = _mm256_loadu_ps(transformed + i); | |
| 123 | |||
| 124 | /* Compare > 0 */ | ||
| 125 | 4965014 | __m256 cmp = _mm256_cmp_ps(t, zero, _CMP_GT_OQ); | |
| 126 | |||
| 127 | /* Extract sign bits as 8-bit mask (LSB-first matches our bit order) */ | ||
| 128 | 4965014 | int mask = _mm256_movemask_ps(cmp); | |
| 129 | |||
| 130 | /* Store as single byte */ | ||
| 131 | 4965014 | bits[i / 8] = (uint8_t)mask; | |
| 132 | } | ||
| 133 | |||
| 134 | /* Handle tail elements */ | ||
| 135 |
2/2✓ Branch 0 taken 73164 times.
✓ Branch 1 taken 611972 times.
|
685136 | if (i < dim) |
| 136 | { | ||
| 137 | 73164 | int byte_idx = i / 8; | |
| 138 | 73164 | bits[byte_idx] = 0; | |
| 139 |
2/2✓ Branch 0 taken 232796 times.
✓ Branch 1 taken 73164 times.
|
305960 | for (; i < dim; i++) |
| 140 | { | ||
| 141 | 232796 | int bit_idx = i % 8; | |
| 142 |
2/2✓ Branch 0 taken 113349 times.
✓ Branch 1 taken 119447 times.
|
232796 | if (transformed[i] > 0) |
| 143 | 113349 | bits[byte_idx] |= (1 << bit_idx); | |
| 144 | } | ||
| 145 | } | ||
| 146 | 685136 | } | |
| 147 | |||
| 148 | /* | ||
| 149 | * AVX2 popcount helper using nibble lookup table (Mula's method). | ||
| 150 | * | ||
| 151 | * Uses _mm256_shuffle_epi8 as a parallel 4-bit lookup table to count | ||
| 152 | * bits in each byte. | ||
| 153 | */ | ||
| 154 | VS_TARGET_AVX2 static inline __m256i | ||
| 155 | 38 | popcount_avx2(__m256i v) | |
| 156 | { | ||
| 157 | 38 | const __m256i lut = _mm256_setr_epi8( | |
| 158 | 0, | ||
| 159 | 1, | ||
| 160 | 1, | ||
| 161 | 2, | ||
| 162 | 1, | ||
| 163 | 2, | ||
| 164 | 2, | ||
| 165 | 3, | ||
| 166 | 1, | ||
| 167 | 2, | ||
| 168 | 2, | ||
| 169 | 3, | ||
| 170 | 2, | ||
| 171 | 3, | ||
| 172 | 3, | ||
| 173 | 4, | ||
| 174 | 0, | ||
| 175 | 1, | ||
| 176 | 1, | ||
| 177 | 2, | ||
| 178 | 1, | ||
| 179 | 2, | ||
| 180 | 2, | ||
| 181 | 3, | ||
| 182 | 1, | ||
| 183 | 2, | ||
| 184 | 2, | ||
| 185 | 3, | ||
| 186 | 2, | ||
| 187 | 3, | ||
| 188 | 3, | ||
| 189 | 4); | ||
| 190 | 38 | const __m256i mask_lo = _mm256_set1_epi8(0x0F); | |
| 191 | |||
| 192 | 38 | __m256i lo = _mm256_and_si256(v, mask_lo); | |
| 193 | 76 | __m256i hi = _mm256_and_si256(_mm256_srli_epi16(v, 4), mask_lo); | |
| 194 | |||
| 195 | 114 | return _mm256_add_epi8( | |
| 196 | _mm256_shuffle_epi8(lut, lo), _mm256_shuffle_epi8(lut, hi)); | ||
| 197 | } | ||
| 198 | |||
| 199 | /* | ||
| 200 | * AVX2 Hamming distance using lookup-table popcount. | ||
| 201 | * | ||
| 202 | * Processes 32 bytes per iteration. Uses _mm256_sad_epu8 to | ||
| 203 | * horizontally sum byte-level popcounts into 64-bit accumulators. | ||
| 204 | */ | ||
| 205 | VS_TARGET_AVX2 uint32_t | ||
| 206 | 146 | vs_rabitq_hamming_avx2( | |
| 207 | const uint8_t *a, const uint8_t *b, uint32_t packed_bytes) | ||
| 208 | { | ||
| 209 | 146 | __m256i total = _mm256_setzero_si256(); | |
| 210 | |||
| 211 | 146 | uint32_t i = 0; | |
| 212 |
2/2✓ Branch 0 taken 38 times.
✓ Branch 1 taken 146 times.
|
184 | for (; i + 32 <= packed_bytes; i += 32) |
| 213 | { | ||
| 214 | 38 | __m256i va = _mm256_loadu_si256((const __m256i *)(a + i)); | |
| 215 | 76 | __m256i vb = _mm256_loadu_si256((const __m256i *)(b + i)); | |
| 216 | 38 | __m256i x = _mm256_xor_si256(va, vb); | |
| 217 | |||
| 218 | 38 | __m256i pc = popcount_avx2(x); | |
| 219 | |||
| 220 | /* Sum byte popcounts into 64-bit accumulators via SAD */ | ||
| 221 | 114 | total = _mm256_add_epi64( | |
| 222 | total, _mm256_sad_epu8(pc, _mm256_setzero_si256())); | ||
| 223 | } | ||
| 224 | |||
| 225 | 146 | uint64_t result = vs_horizontal_sum_epi64_avx2(total); | |
| 226 | |||
| 227 | /* Scalar tail */ | ||
| 228 |
2/2✓ Branch 0 taken 1170 times.
✓ Branch 1 taken 146 times.
|
1316 | for (; i < packed_bytes; i++) |
| 229 | 1170 | result += (uint64_t)__builtin_popcount(a[i] ^ b[i]); | |
| 230 | |||
| 231 | 146 | return (uint32_t)result; | |
| 232 | } | ||
| 233 | |||
| 234 | /* | ||
| 235 | * Multi-candidate AVX2 Hamming distance. | ||
| 236 | */ | ||
| 237 | VS_TARGET_AVX2 void | ||
| 238 | 14 | vs_rabitq_hamming_multi_avx2( | |
| 239 | const uint8_t *query_bits, | ||
| 240 | const uint8_t *data_bits, | ||
| 241 | uint32_t stride, | ||
| 242 | uint32_t packed_bytes, | ||
| 243 | uint32_t count, | ||
| 244 | uint32_t *results) | ||
| 245 | { | ||
| 246 |
2/2✓ Branch 0 taken 74 times.
✓ Branch 1 taken 14 times.
|
88 | for (uint32_t c = 0; c < count; c++) |
| 247 | { | ||
| 248 | 74 | results[c] = vs_rabitq_hamming_avx2( | |
| 249 | 74 | query_bits, data_bits + c * stride, packed_bytes); | |
| 250 | } | ||
| 251 | 14 | } | |
| 252 | |||
| 253 | /* | ||
| 254 | * AVX2 multi-candidate vertical inner product. | ||
| 255 | * | ||
| 256 | * Processes 4 candidates per dimension chunk. Each iteration: | ||
| 257 | * 1. Load 8 floats from transformed[] (1 ymm register) | ||
| 258 | * 2. For each of 4 candidates: expand 1 byte to mask, masked-and-add | ||
| 259 | * 3. After all dimensions: horizontal sum each accumulator | ||
| 260 | * | ||
| 261 | * Tail candidates (count % 4) use the single-candidate kernel. | ||
| 262 | */ | ||
| 263 | VS_TARGET_AVX2 void | ||
| 264 | 1286783 | vs_rabitq_inner_product_multi_avx2( | |
| 265 | const float *transformed, | ||
| 266 | const uint8_t *bits, | ||
| 267 | uint32_t stride, | ||
| 268 | Dimension dim, | ||
| 269 | uint32_t count, | ||
| 270 | float *results) | ||
| 271 | { | ||
| 272 | 1286783 | uint32_t groups = count / 4; | |
| 273 | 1286783 | uint32_t tail = count % 4; | |
| 274 | |||
| 275 |
2/2✓ Branch 0 taken 4279448 times.
✓ Branch 1 taken 1286783 times.
|
5566231 | for (uint32_t g = 0; g < groups; g++) |
| 276 | { | ||
| 277 | 4279448 | uint32_t base = g * 4; | |
| 278 | |||
| 279 | 4279448 | const uint8_t *b0 = bits + (size_t)base * stride; | |
| 280 | 4279448 | const uint8_t *b1 = bits + (size_t)(base + 1) * stride; | |
| 281 | 4279448 | const uint8_t *b2 = bits + (size_t)(base + 2) * stride; | |
| 282 | 4279448 | const uint8_t *b3 = bits + (size_t)(base + 3) * stride; | |
| 283 | |||
| 284 | 4279448 | __m256 sum0 = _mm256_setzero_ps(); | |
| 285 | 4279448 | __m256 sum1 = _mm256_setzero_ps(); | |
| 286 | 4279448 | __m256 sum2 = _mm256_setzero_ps(); | |
| 287 | 4279448 | __m256 sum3 = _mm256_setzero_ps(); | |
| 288 | |||
| 289 | /* Main loop: process 8 floats (1 byte of bits) at a time */ | ||
| 290 | 4279448 | Dimension i = 0; | |
| 291 |
2/2✓ Branch 0 taken 14034567 times.
✓ Branch 1 taken 4279448 times.
|
18314015 | for (; i + 8 <= dim; i += 8) |
| 292 | { | ||
| 293 | 14034567 | __m256 t = _mm256_loadu_ps(transformed + i); | |
| 294 | |||
| 295 | 14034567 | uint32_t bi = i / 8; | |
| 296 | |||
| 297 | 14034567 | __m256 mask0 = expand_byte_to_mask_avx2(b0[bi]); | |
| 298 | 14034567 | __m256 mask1 = expand_byte_to_mask_avx2(b1[bi]); | |
| 299 | 14034567 | __m256 mask2 = expand_byte_to_mask_avx2(b2[bi]); | |
| 300 | 14034567 | __m256 mask3 = expand_byte_to_mask_avx2(b3[bi]); | |
| 301 | |||
| 302 | 24251651 | sum0 = _mm256_add_ps(sum0, _mm256_and_ps(t, mask0)); | |
| 303 | 24251651 | sum1 = _mm256_add_ps(sum1, _mm256_and_ps(t, mask1)); | |
| 304 | 24251651 | sum2 = _mm256_add_ps(sum2, _mm256_and_ps(t, mask2)); | |
| 305 | 24251651 | sum3 = _mm256_add_ps(sum3, _mm256_and_ps(t, mask3)); | |
| 306 | } | ||
| 307 | |||
| 308 | 4279448 | results[base + 0] = vs_horizontal_sum_avx2(sum0); | |
| 309 | 4279448 | results[base + 1] = vs_horizontal_sum_avx2(sum1); | |
| 310 | 4279448 | results[base + 2] = vs_horizontal_sum_avx2(sum2); | |
| 311 | 4279448 | results[base + 3] = vs_horizontal_sum_avx2(sum3); | |
| 312 | |||
| 313 | /* Scalar tail for remaining dimensions */ | ||
| 314 |
2/2✓ Branch 0 taken 211307 times.
✓ Branch 1 taken 4279448 times.
|
4490755 | for (; i < dim; i++) |
| 315 | { | ||
| 316 | 211307 | int byte_idx = i / 8; | |
| 317 | 211307 | int bit_idx = i % 8; | |
| 318 | |||
| 319 |
2/2✓ Branch 0 taken 103942 times.
✓ Branch 1 taken 107365 times.
|
211307 | if ((b0[byte_idx] >> bit_idx) & 1) |
| 320 | 103942 | results[base + 0] += transformed[i]; | |
| 321 |
2/2✓ Branch 0 taken 127977 times.
✓ Branch 1 taken 83330 times.
|
211307 | if ((b1[byte_idx] >> bit_idx) & 1) |
| 322 | 127977 | results[base + 1] += transformed[i]; | |
| 323 |
2/2✓ Branch 0 taken 98863 times.
✓ Branch 1 taken 112444 times.
|
211307 | if ((b2[byte_idx] >> bit_idx) & 1) |
| 324 | 98863 | results[base + 2] += transformed[i]; | |
| 325 |
2/2✓ Branch 0 taken 99031 times.
✓ Branch 1 taken 112276 times.
|
211307 | if ((b3[byte_idx] >> bit_idx) & 1) |
| 326 | 99031 | results[base + 3] += transformed[i]; | |
| 327 | } | ||
| 328 | } | ||
| 329 | |||
| 330 | /* Handle remaining candidates with single-candidate kernel */ | ||
| 331 |
2/2✓ Branch 0 taken 539081 times.
✓ Branch 1 taken 1286783 times.
|
1825864 | for (uint32_t i = groups * 4; i < groups * 4 + tail; i++) |
| 332 | { | ||
| 333 | 539081 | results[i] = vs_rabitq_inner_product_avx2( | |
| 334 | 539081 | transformed, bits + (size_t)i * stride, dim); | |
| 335 | } | ||
| 336 | 1286783 | } | |
| 337 | |||
| 338 | #endif /* x86_64 */ | ||
| 339 | |||
| 340 | #endif /* VS_SIMD_FULL */ | ||
| 341 |