GCC Code Coverage Report


Directory: src/
File: src/quant/rabitq_avx2.c
Date: 2026-09-30 11:11:31
Exec Total Coverage
Lines: 108 108 100.0%
Functions: 7 7 100.0%
Branches: 36 36 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_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