GCC Code Coverage Report


Directory: src/
File: src/quant/fast_rotate.c
Date: 2026-09-30 11:11:31
Exec Total Coverage
Lines: 137 254 53.9%
Functions: 7 11 63.6%
Branches: 89 176 50.6%

Line Branch Exec Source
1 /*
2 * Copyright (c) 2026 Tiger Data, Inc.
3 * Licensed under the PostgreSQL License. See LICENSE for details.
4 *
5 * fast_rotate.c - O(d log d) orthonormal rotation
6 *
7 * Generalised Walsh-Hadamard transform with random sign-flip prefix
8 * and a K×K mixer across K sub-blocks, scaled to be orthonormal. Used
9 * as a drop-in for the dense `P^T * x` rotation in RaBitQ.
10 *
11 * For dim = N (power of two) we set K=1 and apply the classic
12 * randomised Hadamard transform: D · H_N · scale.
13 *
14 * For dim = N · K with N power-of-two and K ≥ 2 we view x as a K×N
15 * matrix (K rows of length N), apply H_N to each row in place, then
16 * apply a K×K orthonormal matrix M across the rows at each column.
17 * The combined map M · (I_K ⊗ H_N) · D · scale is orthonormal because
18 * each factor is, and the Kronecker structure makes it O(d log N + d·K)
19 * to evaluate — for dim=768 = 256·3 that's ~9 k ops vs ~1.18 M ops
20 * for the dense sgemv it replaces.
21 */
22
23 #include <math.h>
24 #include <string.h>
25
26 #include "core/memory.h"
27 #include "core/platform.h"
28 #include "quant/fast_rotate.h"
29
30 #if defined(__x86_64__) || defined(_M_X64)
31 #include <immintrin.h>
32 #elif defined(__aarch64__) || defined(_M_ARM64)
33 #include <arm_neon.h>
34 #endif
35
36 bool
37 1776 vs_fast_rotate_supported(Dimension dim)
38 {
39
2/2
✓ Branch 0 taken 366 times.
✓ Branch 1 taken 1410 times.
1776 if (dim < 4)
40 6 return false;
41
42 /* Factor dim = N * K with N the largest power-of-two divisor.
43 * K must be ≤ VS_FAST_ROTATE_K_MAX, and N ≥ 4 so the inner
44 * FWHT has at least two butterfly stages. */
45 1382 uint32_t n = dim;
46 1382 uint32_t k = 1;
47
2/2
✓ Branch 0 taken 6496 times.
✓ Branch 1 taken 1382 times.
7878 while ((n & 1) == 0)
48 6496 n >>= 1;
49 /* Now n is the odd part, k = original/odd_part is the power-of-two
50 * part of dim. The decomposition we want is the *opposite*:
51 * fwht_n = power-of-two part, k = odd cofactor. */
52 1382 uint32_t fwht_n = dim / n;
53 1382 k = n;
54
2/2
✓ Branch 0 taken 388 times.
✓ Branch 1 taken 994 times.
1382 if (fwht_n < 4)
55 28 return false;
56
2/2
✓ Branch 0 taken 26 times.
✓ Branch 1 taken 1328 times.
1354 if (k > VS_FAST_ROTATE_K_MAX)
57 26 return false;
58 992 return true;
59 }
60
61 /* SplitMix64: deterministic PRNG used only for sign-vector / mixer
62 * init. Self-contained so init doesn't depend on the rest of the
63 * codebase. */
64 static inline uint64_t
65 125626 splitmix64(uint64_t *state)
66 {
67 125626 uint64_t z = (*state += 0x9E3779B97F4A7C15ULL);
68 125626 z = (z ^ (z >> 30)) * 0xBF58476D1CE4E5B9ULL;
69 125626 z = (z ^ (z >> 27)) * 0x94D049BB133111EBULL;
70 125626 return z ^ (z >> 31);
71 }
72
73 /* Build a K×K orthonormal matrix from PRNG state via Gram-Schmidt on
74 * K random vectors. Cheap because K is small (≤ K_MAX). */
75 static void
76 82 init_mixer(float *m, uint32_t k, uint64_t *prng_state)
77 {
78 22 float v[VS_FAST_ROTATE_K_MAX][VS_FAST_ROTATE_K_MAX];
79
80 /* Random vectors */
81
2/2
✓ Branch 0 taken 262 times.
✓ Branch 1 taken 82 times.
344 for (uint32_t i = 0; i < k; i++)
82
2/2
✓ Branch 0 taken 882 times.
✓ Branch 1 taken 262 times.
1144 for (uint32_t j = 0; j < k; j++)
83 {
84 882 uint64_t r = splitmix64(prng_state);
85 /* Standard normal via Box-Muller would be cleaner but the
86 * mixer just needs to be a non-degenerate basis to
87 * orthogonalise — uniform in [-1, 1] is fine. */
88 882 v[i][j] = (float)((double)r / (double)UINT64_MAX) * 2.0f - 1.0f;
89 }
90
91 /* Modified Gram-Schmidt */
92
2/2
✓ Branch 0 taken 262 times.
✓ Branch 1 taken 82 times.
344 for (uint32_t i = 0; i < k; i++)
93 {
94
2/2
✓ Branch 0 taken 310 times.
✓ Branch 1 taken 262 times.
572 for (uint32_t j = 0; j < i; j++)
95 {
96 180 float dot = 0.0f;
97
2/2
✓ Branch 0 taken 1178 times.
✓ Branch 1 taken 310 times.
1488 for (uint32_t c = 0; c < k; c++)
98 1178 dot += v[i][c] * v[j][c];
99
2/2
✓ Branch 0 taken 1178 times.
✓ Branch 1 taken 310 times.
1488 for (uint32_t c = 0; c < k; c++)
100 1178 v[i][c] -= dot * v[j][c];
101 }
102 180 float norm = 0.0f;
103
2/2
✓ Branch 0 taken 882 times.
✓ Branch 1 taken 262 times.
1144 for (uint32_t c = 0; c < k; c++)
104 882 norm += v[i][c] * v[i][c];
105 262 norm = sqrtf(norm);
106 /* Degenerate rows are vanishingly unlikely for random uniform
107 * inputs at K ≤ 8; if it happens, fall back to a canonical
108 * basis vector for that slot. */
109
2/2
✓ Branch 0 taken 82 times.
✓ Branch 1 taken 180 times.
262 if (norm < 1e-6f)
110 {
111 ✗ for (uint32_t c = 0; c < k; c++)
112 ✗ v[i][c] = (c == i) ? 1.0f : 0.0f;
113 }
114 else
115 {
116
2/2
✓ Branch 0 taken 882 times.
✓ Branch 1 taken 262 times.
1144 for (uint32_t c = 0; c < k; c++)
117 882 v[i][c] /= norm;
118 }
119 }
120
121
2/2
✓ Branch 0 taken 262 times.
✓ Branch 1 taken 82 times.
344 for (uint32_t i = 0; i < k; i++)
122
2/2
✓ Branch 0 taken 882 times.
✓ Branch 1 taken 262 times.
1144 for (uint32_t j = 0; j < k; j++)
123 882 m[i * k + j] = v[i][j];
124 82 }
125
126 void
127 682 vs_fast_rotate_init(VsFastRotateParams *p, Dimension dim, uint64_t seed)
128 {
129 682 p->dim = dim;
130 682 p->seed = seed;
131
132 /* Factor dim = fwht_n * k with fwht_n the power-of-two part. */
133 682 uint32_t odd = dim;
134
2/2
✓ Branch 0 taken 3342 times.
✓ Branch 1 taken 682 times.
4024 while ((odd & 1) == 0)
135 3342 odd >>= 1;
136 682 p->fwht_n = dim / odd;
137 682 p->k = odd;
138
139
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 168 times.
682 memset(p->signs1, 0, sizeof(p->signs1));
140 682 memset(p->signs2, 0, sizeof(p->signs2));
141
142
2/2
✓ Branch 0 taken 514 times.
✓ Branch 1 taken 168 times.
682 uint64_t s = seed ? seed : 0xDEADBEEFCAFEBABEULL;
143
2/2
✓ Branch 0 taken 62372 times.
✓ Branch 1 taken 682 times.
63054 for (Dimension i = 0; i < dim; i++)
144 {
145 62372 uint64_t r = splitmix64(&s);
146
2/2
✓ Branch 0 taken 30605 times.
✓ Branch 1 taken 31767 times.
62372 if (r & 1)
147 30605 p->signs1[i >> 3] |= (uint8_t)(1u << (i & 7));
148 }
149
2/2
✓ Branch 0 taken 62372 times.
✓ Branch 1 taken 682 times.
63054 for (Dimension i = 0; i < dim; i++)
150 {
151 62372 uint64_t r = splitmix64(&s);
152
2/2
✓ Branch 0 taken 31825 times.
✓ Branch 1 taken 30547 times.
62372 if (r & 1)
153 31825 p->signs2[i >> 3] |= (uint8_t)(1u << (i & 7));
154 }
155
156
2/2
✓ Branch 0 taken 146 times.
✓ Branch 1 taken 22 times.
682 memset(p->mixer, 0, sizeof(p->mixer));
157
2/2
✓ Branch 0 taken 600 times.
✓ Branch 1 taken 82 times.
682 if (p->k == 1)
158 600 p->mixer[0] = 1.0f;
159 else
160 82 init_mixer(p->mixer, p->k, &s);
161 682 }
162
163 /* In-place Walsh-Hadamard on a length-N block (N power of two). */
164 static inline void
165 ✗ fwht_inplace(float *a, Dimension n)
166 {
167 ✗ for (uint32_t h = 1; h < n; h <<= 1)
168 {
169 ✗ for (uint32_t i = 0; i < n; i += (h << 1))
170 {
171 ✗ for (uint32_t j = i; j < i + h; j++)
172 {
173 ✗ float x = a[j];
174 ✗ float y = a[j + h];
175 ✗ a[j] = x + y;
176 ✗ a[j + h] = x - y;
177 }
178 }
179 }
180 ✗ }
181
182 /* Apply the K×K mixer matrix M to the K-vector formed by taking
183 * `col` from each of K rows. Mixer rows = K, columns = K. */
184 static inline void
185 ✗ apply_mixer_column(
186 float *out, const float *m, uint32_t k, Dimension n, Dimension col)
187 {
188 ✗ float tmp[VS_FAST_ROTATE_K_MAX];
189 ✗ for (uint32_t i = 0; i < k; i++)
190 ✗ tmp[i] = out[i * n + col];
191
192 ✗ for (uint32_t i = 0; i < k; i++)
193 {
194 ✗ float s = 0.0f;
195 ✗ for (uint32_t j = 0; j < k; j++)
196 ✗ s += m[i * k + j] * tmp[j];
197 ✗ out[i * n + col] = s;
198 }
199 ✗ }
200
201 /* Reference scalar implementation. */
202 static void
203 ✗ fast_rotate_apply_scalar(
204 const VsFastRotateParams *p, const float *in, float *out)
205 {
206 ✗ Dimension dim = p->dim;
207 ✗ Dimension n = p->fwht_n;
208 ✗ uint32_t k = p->k;
209 /* Unscaled FWHT on a length-N block has Parseval factor N
210 * (||y||² = N · ||x||²). The K×K mixer is unit-norm. So pre-scale
211 * by 1/sqrt(N), NOT 1/sqrt(dim) — that gets the overall map to
212 * unit-norm regardless of K. */
213 ✗ float invsq = 1.0f / sqrtf((float)n);
214
215 /* Pre-FWHT sign flip (D1) + scale. Folding the 1/sqrt(N) scale
216 * into this pass means the butterflies and mixer can run on plain
217 * sums without a final scaling sweep. */
218 ✗ for (Dimension i = 0; i < dim; i++)
219 {
220 ✗ float sign = ((p->signs1[i >> 3] >> (i & 7)) & 1u) ? -1.0f : 1.0f;
221 ✗ out[i] = in[i] * sign * invsq;
222 }
223
224 /* FWHT on each of K sub-blocks of length N. When K==1 this is
225 * just the classic randomised Hadamard rotation. */
226 ✗ for (uint32_t b = 0; b < k; b++)
227 ✗ fwht_inplace(out + (size_t)b * n, n);
228
229 /* K×K mixer across the blocks. Skip when K==1: mixer is identity. */
230 ✗ if (k > 1)
231 {
232 ✗ for (Dimension col = 0; col < n; col++)
233 ✗ apply_mixer_column(out, p->mixer, k, n, col);
234 }
235
236 /* Post-FWHT sign flip (D2). This second round of randomisation
237 * tightens the concentration of the transform: with one round,
238 * worst-case Lipschitz can break RaBitQ's strict lower bound at
239 * small dim; two rounds match the FJLT analysis (Ailon-Chazelle)
240 * and restore the bound. Cost is one extra O(d) pass. */
241 ✗ for (Dimension i = 0; i < dim; i++)
242 {
243 ✗ float sign = ((p->signs2[i >> 3] >> (i & 7)) & 1u) ? -1.0f : 1.0f;
244 ✗ out[i] *= sign;
245 }
246 ✗ }
247
248 #if defined(__x86_64__) || defined(_M_X64)
249 /*
250 * AVX2 kernel. Vectorizes the parts that map cleanly to 8-wide SIMD:
251 * - D1/D2 sign flips: expand 8 packed sign bits to a ±1 vector, multiply.
252 * - FWHT stages with stride h >= 8: contiguous wide add/sub (bit-identical
253 * to the scalar butterfly). Small stages (h < 8) stay scalar -- they'd
254 * need in-register shuffle networks for little of the total work.
255 * - K x K mixer: vectorized across 8 columns (contiguous within a block).
256 */
257 __attribute__((target("avx2,fma"))) static void
258 701528 fast_rotate_apply_avx2(
259 const VsFastRotateParams *p, const float *in, float *out)
260 {
261 701528 Dimension dim = p->dim;
262 701528 Dimension n = p->fwht_n;
263 701528 uint32_t k = p->k;
264 701528 float invsq = 1.0f / sqrtf((float)n);
265 701528 const __m256i lanebits = _mm256_setr_epi32(1, 2, 4, 8, 16, 32, 64, 128);
266 701528 const __m256 vpos = _mm256_set1_ps(1.0f);
267 701528 const __m256 vneg = _mm256_set1_ps(-1.0f);
268
269 /* D1: sign flip + 1/sqrt(N) scale. */
270 701528 __m256 vscale = _mm256_set1_ps(invsq);
271 701528 Dimension i = 0;
272
2/2
✓ Branch 0 taken 4846458 times.
✓ Branch 1 taken 701528 times.
5547986 for (; i + 8 <= dim; i += 8)
273 {
274 5930542 __m256i bb = _mm256_and_si256(
275 4846458 _mm256_set1_epi32(p->signs1[i >> 3]), lanebits);
276 5930542 __m256 sign = _mm256_blendv_ps(
277 vpos,
278 vneg,
279 _mm256_castsi256_ps(_mm256_cmpeq_epi32(bb, lanebits)));
280 5930542 __m256 v = _mm256_mul_ps(
281 4846458 _mm256_mul_ps(_mm256_loadu_ps(in + i), sign), vscale);
282 4846458 _mm256_storeu_ps(out + i, v);
283 }
284
2/2
✓ Branch 0 taken 53580 times.
✓ Branch 1 taken 701528 times.
755108 for (; i < dim; i++)
285 {
286
2/2
✓ Branch 0 taken 26790 times.
✓ Branch 1 taken 26790 times.
53580 float s = ((p->signs1[i >> 3] >> (i & 7)) & 1u) ? -1.0f : 1.0f;
287 53580 out[i] = in[i] * s * invsq;
288 }
289
290 /* FWHT per block: wide add/sub for h >= 8, scalar for the small tail. */
291
2/2
✓ Branch 0 taken 915790 times.
✓ Branch 1 taken 701528 times.
1617318 for (uint32_t blk = 0; blk < k; blk++)
292 {
293 915790 float *a = out + (size_t)blk * n;
294
2/2
✓ Branch 0 taken 4119896 times.
✓ Branch 1 taken 915790 times.
5035686 for (uint32_t h = 1; h < n; h <<= 1)
295 {
296
2/2
✓ Branch 0 taken 1385921 times.
✓ Branch 1 taken 2733975 times.
4119896 if (h >= 8)
297 {
298
2/2
✓ Branch 0 taken 3944063 times.
✓ Branch 1 taken 1385921 times.
5329984 for (uint32_t s = 0; s < n; s += (h << 1))
299
2/2
✓ Branch 0 taken 8278298 times.
✓ Branch 1 taken 3944063 times.
12222361 for (uint32_t j = s; j < s + h; j += 8)
300 {
301 8278298 __m256 x = _mm256_loadu_ps(a + j);
302 9303860 __m256 y = _mm256_loadu_ps(a + j + h);
303 8278298 _mm256_storeu_ps(a + j, _mm256_add_ps(x, y));
304 8278298 _mm256_storeu_ps(a + j + h, _mm256_sub_ps(x, y));
305 }
306 }
307 else
308 {
309
2/2
✓ Branch 0 taken 33965391 times.
✓ Branch 1 taken 2733975 times.
36699366 for (uint32_t s = 0; s < n; s += (h << 1))
310
2/2
✓ Branch 0 taken 58211076 times.
✓ Branch 1 taken 33965391 times.
92176467 for (uint32_t j = s; j < s + h; j++)
311 {
312 58211076 float x = a[j];
313 58211076 float y = a[j + h];
314 58211076 a[j] = x + y;
315 58211076 a[j + h] = x - y;
316 }
317 }
318 }
319 }
320
321 /* K x K mixer across blocks, vectorized across columns. */
322
2/2
✓ Branch 0 taken 106045 times.
✓ Branch 1 taken 595483 times.
701528 if (k > 1)
323 {
324 91792 Dimension col = 0;
325
2/2
✓ Branch 0 taken 517431 times.
✓ Branch 1 taken 106045 times.
623476 for (; col + 8 <= n; col += 8)
326 {
327 __m256 blkv[VS_FAST_ROTATE_K_MAX];
328
2/2
✓ Branch 0 taken 1555003 times.
✓ Branch 1 taken 517431 times.
2072434 for (uint32_t j = 0; j < k; j++)
329 1863307 blkv[j] = _mm256_loadu_ps(out + (size_t)j * n + col);
330
2/2
✓ Branch 0 taken 1555003 times.
✓ Branch 1 taken 517431 times.
2072434 for (uint32_t r = 0; r < k; r++)
331 {
332 1555003 __m256 acc = _mm256_mul_ps(
333 1555003 _mm256_set1_ps(p->mixer[r * k]), blkv[0]);
334
2/2
✓ Branch 0 taken 3125740 times.
✓ Branch 1 taken 1555003 times.
4680743 for (uint32_t j = 1; j < k; j++)
335 3742348 acc = _mm256_fmadd_ps(
336 3125740 _mm256_set1_ps(p->mixer[r * k + j]), blkv[j], acc);
337 1555003 _mm256_storeu_ps(out + (size_t)r * n + col, acc);
338 }
339 }
340
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 106045 times.
106045 for (; col < n; col++)
341 {
342 float tmp[VS_FAST_ROTATE_K_MAX];
343 ✗ for (uint32_t j = 0; j < k; j++)
344 ✗ tmp[j] = out[(size_t)j * n + col];
345 ✗ for (uint32_t r = 0; r < k; r++)
346 {
347 ✗ float s = 0.0f;
348 ✗ for (uint32_t j = 0; j < k; j++)
349 ✗ s += p->mixer[r * k + j] * tmp[j];
350 ✗ out[(size_t)r * n + col] = s;
351 }
352 }
353 }
354
355 /* D2: post-FWHT sign flip. */
356 274250 i = 0;
357
2/2
✓ Branch 0 taken 4846458 times.
✓ Branch 1 taken 701528 times.
5547986 for (; i + 8 <= dim; i += 8)
358 {
359 5930542 __m256i bb = _mm256_and_si256(
360 4846458 _mm256_set1_epi32(p->signs2[i >> 3]), lanebits);
361 5930542 __m256 sign = _mm256_blendv_ps(
362 vpos,
363 vneg,
364 _mm256_castsi256_ps(_mm256_cmpeq_epi32(bb, lanebits)));
365 4846458 _mm256_storeu_ps(
366 5930542 out + i, _mm256_mul_ps(_mm256_loadu_ps(out + i), sign));
367 }
368
2/2
✓ Branch 0 taken 53580 times.
✓ Branch 1 taken 701528 times.
755108 for (; i < dim; i++)
369 {
370
2/2
✓ Branch 0 taken 40185 times.
✓ Branch 1 taken 13395 times.
53580 float s = ((p->signs2[i >> 3] >> (i & 7)) & 1u) ? -1.0f : 1.0f;
371 53580 out[i] *= s;
372 }
373 701528 }
374
375 /*
376 * AVX-512 kernel (16-wide). Same structure as the AVX2 kernel; the packed
377 * sign bits map directly onto a __mmask16 so D1/D2 need no compare trick.
378 * Uses only avx512f (float arithmetic + masked blend), so it is gated on
379 * SIMD_AVX512F alone. FWHT: 512-wide for h >= 16, 256-wide for h == 8,
380 * scalar below that.
381 */
382 __attribute__((target("avx512f"))) static void
383 ✗ fast_rotate_apply_avx512(
384 const VsFastRotateParams *p, const float *in, float *out)
385 {
386 ✗ Dimension dim = p->dim;
387 ✗ Dimension n = p->fwht_n;
388 ✗ uint32_t k = p->k;
389 ✗ float invsq = 1.0f / sqrtf((float)n);
390 ✗ const __m512 vpos = _mm512_set1_ps(1.0f);
391 ✗ const __m512 vneg = _mm512_set1_ps(-1.0f);
392 ✗ const __m512 vscale = _mm512_set1_ps(invsq);
393
394 /* D1: sign flip + scale. */
395 ✗ Dimension i = 0;
396 ✗ for (; i + 16 <= dim; i += 16)
397 {
398 ✗ __mmask16 m = (__mmask16)(*(const uint16_t *)(p->signs1 + (i >> 3)));
399 ✗ __m512 sign = _mm512_mask_blend_ps(m, vpos, vneg);
400 ✗ __m512 v = _mm512_mul_ps(
401 ✗ _mm512_mul_ps(_mm512_loadu_ps(in + i), sign), vscale);
402 ✗ _mm512_storeu_ps(out + i, v);
403 }
404 ✗ for (; i < dim; i++)
405 {
406 ✗ float s = ((p->signs1[i >> 3] >> (i & 7)) & 1u) ? -1.0f : 1.0f;
407 ✗ out[i] = in[i] * s * invsq;
408 }
409
410 /* FWHT per block. */
411 ✗ for (uint32_t blk = 0; blk < k; blk++)
412 {
413 ✗ float *a = out + (size_t)blk * n;
414 ✗ for (uint32_t h = 1; h < n; h <<= 1)
415 {
416 ✗ if (h >= 16)
417 {
418 ✗ for (uint32_t s = 0; s < n; s += (h << 1))
419 ✗ for (uint32_t j = s; j < s + h; j += 16)
420 {
421 ✗ __m512 x = _mm512_loadu_ps(a + j);
422 ✗ __m512 y = _mm512_loadu_ps(a + j + h);
423 ✗ _mm512_storeu_ps(a + j, _mm512_add_ps(x, y));
424 ✗ _mm512_storeu_ps(a + j + h, _mm512_sub_ps(x, y));
425 }
426 }
427 ✗ else if (h == 8)
428 {
429 ✗ for (uint32_t s = 0; s < n; s += 16)
430 {
431 ✗ __m256 x = _mm256_loadu_ps(a + s);
432 ✗ __m256 y = _mm256_loadu_ps(a + s + 8);
433 ✗ _mm256_storeu_ps(a + s, _mm256_add_ps(x, y));
434 ✗ _mm256_storeu_ps(a + s + 8, _mm256_sub_ps(x, y));
435 }
436 }
437 else
438 {
439 ✗ for (uint32_t s = 0; s < n; s += (h << 1))
440 ✗ for (uint32_t j = s; j < s + h; j++)
441 {
442 ✗ float x = a[j];
443 ✗ float y = a[j + h];
444 ✗ a[j] = x + y;
445 ✗ a[j + h] = x - y;
446 }
447 }
448 }
449 }
450
451 /* K x K mixer across blocks, vectorized across columns. */
452 ✗ if (k > 1)
453 {
454 ✗ Dimension col = 0;
455 ✗ for (; col + 16 <= n; col += 16)
456 {
457 __m512 blkv[VS_FAST_ROTATE_K_MAX];
458 ✗ for (uint32_t j = 0; j < k; j++)
459 ✗ blkv[j] = _mm512_loadu_ps(out + (size_t)j * n + col);
460 ✗ for (uint32_t r = 0; r < k; r++)
461 {
462 ✗ __m512 acc = _mm512_mul_ps(
463 ✗ _mm512_set1_ps(p->mixer[r * k]), blkv[0]);
464 ✗ for (uint32_t j = 1; j < k; j++)
465 ✗ acc = _mm512_fmadd_ps(
466 ✗ _mm512_set1_ps(p->mixer[r * k + j]), blkv[j], acc);
467 ✗ _mm512_storeu_ps(out + (size_t)r * n + col, acc);
468 }
469 }
470 ✗ for (; col < n; col++)
471 {
472 float tmp[VS_FAST_ROTATE_K_MAX];
473 ✗ for (uint32_t j = 0; j < k; j++)
474 ✗ tmp[j] = out[(size_t)j * n + col];
475 ✗ for (uint32_t r = 0; r < k; r++)
476 {
477 ✗ float s = 0.0f;
478 ✗ for (uint32_t j = 0; j < k; j++)
479 ✗ s += p->mixer[r * k + j] * tmp[j];
480 ✗ out[(size_t)r * n + col] = s;
481 }
482 }
483 }
484
485 /* D2: post-FWHT sign flip. */
486 ✗ i = 0;
487 ✗ for (; i + 16 <= dim; i += 16)
488 {
489 ✗ __mmask16 m = (__mmask16)(*(const uint16_t *)(p->signs2 + (i >> 3)));
490 ✗ __m512 sign = _mm512_mask_blend_ps(m, vpos, vneg);
491 ✗ _mm512_storeu_ps(
492 ✗ out + i, _mm512_mul_ps(_mm512_loadu_ps(out + i), sign));
493 }
494 ✗ for (; i < dim; i++)
495 {
496 ✗ float s = ((p->signs2[i >> 3] >> (i & 7)) & 1u) ? -1.0f : 1.0f;
497 ✗ out[i] *= s;
498 }
499 ✗ }
500 #endif /* x86_64 */
501
502 #if defined(__aarch64__) || defined(_M_ARM64)
503 /*
504 * NEON kernel (128-bit / 4-wide). Same structure as the AVX2 kernel: D1/D2
505 * expand packed sign bits to a +/-1 vector via vtstq/vbslq; FWHT stages with
506 * stride h >= 4 use contiguous wide add/sub (smaller stages stay scalar);
507 * the K x K mixer is vectorized across 4 columns.
508 */
509 static void
510 fast_rotate_apply_neon(
511 const VsFastRotateParams *p, const float *in, float *out)
512 {
513 Dimension dim = p->dim;
514 Dimension n = p->fwht_n;
515 uint32_t k = p->k;
516 float invsq = 1.0f / sqrtf((float)n);
517 const uint32_t lanebits_arr[4] = {1, 2, 4, 8};
518 const uint32x4_t lanebits = vld1q_u32(lanebits_arr);
519 const float32x4_t vpos = vdupq_n_f32(1.0f);
520 const float32x4_t vneg = vdupq_n_f32(-1.0f);
521
522 /* D1: sign flip + scale. */
523 float32x4_t vscale = vdupq_n_f32(invsq);
524 Dimension i = 0;
525 for (; i + 4 <= dim; i += 4)
526 {
527 uint32_t nib = (uint32_t)((p->signs1[i >> 3] >> (i & 7)) & 0xF);
528 uint32x4_t iss = vtstq_u32(vdupq_n_u32(nib), lanebits);
529 float32x4_t sign = vbslq_f32(iss, vneg, vpos);
530 float32x4_t v = vmulq_f32(vmulq_f32(vld1q_f32(in + i), sign), vscale);
531 vst1q_f32(out + i, v);
532 }
533 for (; i < dim; i++)
534 {
535 float s = ((p->signs1[i >> 3] >> (i & 7)) & 1u) ? -1.0f : 1.0f;
536 out[i] = in[i] * s * invsq;
537 }
538
539 /* FWHT per block: wide add/sub for h >= 4, scalar below. */
540 for (uint32_t blk = 0; blk < k; blk++)
541 {
542 float *a = out + (size_t)blk * n;
543 for (uint32_t h = 1; h < n; h <<= 1)
544 {
545 if (h >= 4)
546 {
547 for (uint32_t s = 0; s < n; s += (h << 1))
548 for (Dimension j = s; j < s + h; j += 4)
549 {
550 float32x4_t x = vld1q_f32(a + j);
551 float32x4_t y = vld1q_f32(a + j + h);
552 vst1q_f32(a + j, vaddq_f32(x, y));
553 vst1q_f32(a + j + h, vsubq_f32(x, y));
554 }
555 }
556 else
557 {
558 for (uint32_t s = 0; s < n; s += (h << 1))
559 for (uint32_t j = s; j < s + h; j++)
560 {
561 float x = a[j];
562 float y = a[j + h];
563 a[j] = x + y;
564 a[j + h] = x - y;
565 }
566 }
567 }
568 }
569
570 /* K x K mixer across blocks, vectorized across columns. */
571 if (k > 1)
572 {
573 Dimension col = 0;
574 for (; col + 4 <= n; col += 4)
575 {
576 float32x4_t blkv[VS_FAST_ROTATE_K_MAX];
577 for (uint32_t j = 0; j < k; j++)
578 blkv[j] = vld1q_f32(out + (size_t)j * n + col);
579 for (uint32_t r = 0; r < k; r++)
580 {
581 float32x4_t acc =
582 vmulq_f32(vdupq_n_f32(p->mixer[r * k]), blkv[0]);
583 for (uint32_t j = 1; j < k; j++)
584 acc = vfmaq_f32(
585 acc, vdupq_n_f32(p->mixer[r * k + j]), blkv[j]);
586 vst1q_f32(out + (size_t)r * n + col, acc);
587 }
588 }
589 for (; col < n; col++)
590 {
591 float tmp[VS_FAST_ROTATE_K_MAX];
592 for (uint32_t j = 0; j < k; j++)
593 tmp[j] = out[(size_t)j * n + col];
594 for (uint32_t r = 0; r < k; r++)
595 {
596 float s = 0.0f;
597 for (uint32_t j = 0; j < k; j++)
598 s += p->mixer[r * k + j] * tmp[j];
599 out[(size_t)r * n + col] = s;
600 }
601 }
602 }
603
604 /* D2: post-FWHT sign flip. */
605 i = 0;
606 for (; i + 4 <= dim; i += 4)
607 {
608 uint32_t nib = (uint32_t)((p->signs2[i >> 3] >> (i & 7)) & 0xF);
609 uint32x4_t iss = vtstq_u32(vdupq_n_u32(nib), lanebits);
610 float32x4_t sign = vbslq_f32(iss, vneg, vpos);
611 vst1q_f32(out + i, vmulq_f32(vld1q_f32(out + i), sign));
612 }
613 for (; i < dim; i++)
614 {
615 float s = ((p->signs2[i >> 3] >> (i & 7)) & 1u) ? -1.0f : 1.0f;
616 out[i] *= s;
617 }
618 }
619 #endif /* aarch64 */
620
621 /*
622 * Pick the best kernel for the detected CPU. Resolved once and cached in
623 * the function pointer below (mirroring the distance/fastscan dispatchers)
624 * so the hot path -- which runs per encoded vector at build and once per
625 * query -- pays no per-call capability check.
626 */
627 typedef void (*vs_fast_rotate_fn)(
628 const VsFastRotateParams *, const float *, float *);
629
630 static vs_fast_rotate_fn
631 96 resolve_fast_rotate(void)
632 {
633 #if defined(__x86_64__) || defined(_M_X64)
634
2/2
✓ Branch 1 taken 86 times.
✓ Branch 2 taken 10 times.
96 if (vs_has_simd(SIMD_AVX512F))
635 ✗ return fast_rotate_apply_avx512;
636
1/2
✓ Branch 1 taken 96 times.
✗ Branch 2 not taken.
96 if (vs_has_simd(SIMD_AVX2))
637 96 return fast_rotate_apply_avx2;
638 #elif defined(__aarch64__) || defined(_M_ARM64)
639 if (vs_has_simd(SIMD_NEON))
640 return fast_rotate_apply_neon;
641 #endif
642 ✗ return fast_rotate_apply_scalar;
643 }
644
645 void
646 701528 vs_fast_rotate_apply(const VsFastRotateParams *p, const float *in, float *out)
647 {
648 /* Benign race: concurrent resolvers all compute the same pointer, and a
649 * pointer store is atomic on supported platforms. */
650 427278 static vs_fast_rotate_fn fn = NULL;
651
2/2
✓ Branch 0 taken 96 times.
✓ Branch 1 taken 701432 times.
701528 if (vs_unlikely(fn == NULL))
652 96 fn = resolve_fast_rotate();
653 701528 fn(p, in, out);
654 701528 }
655