@@ -55,9 +55,62 @@ DEGLIB_TARGET_AVX512 inline static int64_t int8_ip_hsum512(__m512i s) {
5555 return int8_ip_hsum256 (sum256);
5656 #endif
5757}
58+ template <size_t BATCH_SIZE = 8 >
59+ inline static void process_batch_tail (
60+ const int8_t * query,
61+ const void * const * db,
62+ size_t offset,
63+ size_t dim,
64+ float * out_dists
65+ ) {
66+ const size_t tail_len = dim - offset;
67+ if (tail_len >= 8 ) {
68+ uint64_t q64;
69+ std::memcpy (&q64, query + offset, sizeof (uint64_t ));
70+ const int8_t * q_bytes = reinterpret_cast <const int8_t *>(&q64);
71+ for (size_t j = 0 ; j < BATCH_SIZE ; ++j) {
72+ const int8_t * db_ptr = static_cast <const int8_t *>(db[j]);
73+ uint64_t db64;
74+ std::memcpy (&db64, db_ptr + offset, sizeof (uint64_t ));
75+ const int8_t * db_bytes = reinterpret_cast <const int8_t *>(&db64);
76+ int32_t tsum = 0 ;
77+ for (size_t k = 0 ; k < 8 ; ++k) {
78+ tsum += int32_t (q_bytes[k]) * int32_t (db_bytes[k]);
79+ }
80+ out_dists[j] -= static_cast <float >(tsum);
81+ }
82+ offset += 8 ;
83+ }
84+ for (size_t k = offset; k < dim; ++k) {
85+ const int32_t q_val = static_cast <int32_t >(query[k]);
86+ for (size_t j = 0 ; j < BATCH_SIZE ; ++j) {
87+ const int8_t * db_ptr = static_cast <const int8_t *>(db[j]);
88+ out_dists[j] -= static_cast <float >(q_val * static_cast <int32_t >(db_ptr[k]));
89+ }
90+ }
91+ }
5892
5993// ---------------------------------------------------------------------------------------------------------------------
6094// AVX512-VNNI (vpdpbusd with signed transform)
95+ //
96+ // Mathematical Background for Signed x Signed dot product via vpdpbusd:
97+ // Hardware vpdpbusd natively computes unsigned int8 (u) * signed int8 (s):
98+ // dot(u, s) = sum(u_i * s_i)
99+ // To compute signed dot product sum(a_i * b_i):
100+ // 1. Map b_i (signed) to unsigned range [0, 255] by adding +128 (via XOR 0x80):
101+ // u_b_i = b_i + 128
102+ // 2. Execute vpdpbusd(u_b, a):
103+ // sum((b_i + 128) * a_i) = sum(a_i * b_i) + 128 * sum(a_i)
104+ // 3. Subtract the query bias compensation (q_correction = 128 * sum(a_i)):
105+ // sum(a_i * b_i) = vpdpbusd_result - 128 * sum(a_i)
106+ //
107+ // In compare_batch(), q_correction depends solely on the query vector and is
108+ // precomputed once across SIMD chunks, avoiding redundant calculations per DB vector.
109+ //
110+ // Note on newer ISAs (AVX-VNNI-INT8 / AVX10.2):
111+ // Newer architectures (e.g., Intel Granite Rapids / Arrow Lake / Lunar Lake) introduce
112+ // native signed-signed instructions `vpdpbssd` / `vpdpbssds`, which compute signed x signed
113+ // directly in one instruction without requiring XOR 0x80 or q_correction.
61114// ---------------------------------------------------------------------------------------------------------------------
62115template <ResidualMode Mode = ResidualMode::Full>
63116class InnerProductInt8_AVX512_VNNI {
@@ -135,22 +188,31 @@ class InnerProductInt8_AVX512_VNNI {
135188 const int8_t * query = static_cast <const int8_t *>(query_ptr);
136189 const size_t dim = *((const size_t *)qty_ptr);
137190
138- auto batch_impl = [query, dim](const void * const * db, float * out_dists) DEGLIB_TARGET_AVX512_VNNI {
191+ // Precalculate q_correction once for the entire batch of DB vectors
192+ int64_t q_correction = 0 ;
193+ if constexpr (HasDualSimd || HasSimd) {
194+ __m512i q_comp = _mm512_setzero_si512 ();
195+ const size_t simd_limit = dim >= 64 ? (dim - 63 ) : 0 ;
196+ for (size_t offset = 0 ; offset < simd_limit; offset += 64 ) {
197+ __m512i q_raw = _mm512_loadu_si512 (reinterpret_cast <const __m512i*>(query + offset));
198+ q_comp = _mm512_add_epi32 (q_comp, _mm512_madd_epi16 (_mm512_cvtepi8_epi16 (_mm512_castsi512_si256 (q_raw)), _mm512_set1_epi16 (1 )));
199+ q_comp = _mm512_add_epi32 (q_comp, _mm512_madd_epi16 (_mm512_cvtepi8_epi16 (_mm512_extracti64x4_epi64 (q_raw, 1 )), _mm512_set1_epi16 (1 )));
200+ }
201+ q_correction = int8_ip_hsum512 (q_comp) * 128 ;
202+ }
203+
204+ auto batch_impl = [query, dim, q_correction](const void * const * db, float * out_dists) DEGLIB_TARGET_AVX512_VNNI {
139205 const __m512i xor_mask = _mm512_set1_epi8 (static_cast <char >(0x80 ));
140206 __m512i acc[BATCH_SIZE ];
141207 for (size_t j = 0 ; j < BATCH_SIZE ; ++j) {
142208 acc[j] = _mm512_setzero_si512 ();
143209 }
144- __m512i q_comp = _mm512_setzero_si512 ();
145210
146211 size_t offset = 0 ;
147212 if constexpr (HasDualSimd || HasSimd) {
148213 const size_t simd_limit = dim >= 64 ? (dim - 63 ) : 0 ;
149214 for (; offset < simd_limit; offset += 64 ) {
150215 __m512i q_raw = _mm512_loadu_si512 (reinterpret_cast <const __m512i*>(query + offset));
151- q_comp = _mm512_add_epi32 (q_comp, _mm512_madd_epi16 (_mm512_cvtepi8_epi16 (_mm512_castsi512_si256 (q_raw)), _mm512_set1_epi16 (1 )));
152- q_comp = _mm512_add_epi32 (q_comp, _mm512_madd_epi16 (_mm512_cvtepi8_epi16 (_mm512_extracti64x4_epi64 (q_raw, 1 )), _mm512_set1_epi16 (1 )));
153-
154216 for (size_t j = 0 ; j < BATCH_SIZE ; ++j) {
155217 const int8_t * db_ptr = static_cast <const int8_t *>(db[j]);
156218 __m512i r_raw = _mm512_loadu_si512 (reinterpret_cast <const __m512i*>(db_ptr + offset));
@@ -160,22 +222,13 @@ class InnerProductInt8_AVX512_VNNI {
160222 }
161223 }
162224
163- const int64_t q_correction = int8_ip_hsum512 (q_comp) * 128 ;
164-
165225 for (size_t j = 0 ; j < BATCH_SIZE ; ++j) {
166226 int64_t total = int8_ip_hsum512 (acc[j]) - q_correction;
167227 out_dists[j] = -static_cast <float >(total);
168228 }
169229
170230 if constexpr (HasTail) {
171- for (size_t j = 0 ; j < BATCH_SIZE ; ++j) {
172- const int8_t * db_ptr = static_cast <const int8_t *>(db[j]);
173- int64_t tail_sum = 0 ;
174- for (size_t k = offset; k < dim; ++k) {
175- tail_sum += int64_t (query[k]) * int64_t (db_ptr[k]);
176- }
177- out_dists[j] -= static_cast <float >(tail_sum);
178- }
231+ process_batch_tail<BATCH_SIZE >(query, db, offset, dim, out_dists);
179232 }
180233 };
181234
@@ -275,14 +328,7 @@ class InnerProductInt8_AVX512 {
275328 }
276329
277330 if constexpr (HasTail) {
278- for (size_t j = 0 ; j < BATCH_SIZE ; ++j) {
279- const int8_t * db_ptr = static_cast <const int8_t *>(db[j]);
280- int64_t tail_sum = 0 ;
281- for (size_t k = offset; k < dim; ++k) {
282- tail_sum += int64_t (query[k]) * int64_t (db_ptr[k]);
283- }
284- out_dists[j] -= static_cast <float >(tail_sum);
285- }
331+ process_batch_tail<BATCH_SIZE >(query, db, offset, dim, out_dists);
286332 }
287333 };
288334
@@ -298,6 +344,7 @@ class InnerProductInt8_AVX512 {
298344
299345// ---------------------------------------------------------------------------------------------------------------------
300346// AVX2-VNNI (AVX_VNNI: _mm256_dpbusd_epi32)
347+ // Uses the same signed x signed mathematical transform as AVX512_VNNI (see comments above).
301348// ---------------------------------------------------------------------------------------------------------------------
302349template <ResidualMode Mode = ResidualMode::Full>
303350class InnerProductInt8_AVX2_VNNI {
@@ -374,23 +421,31 @@ class InnerProductInt8_AVX2_VNNI {
374421 static constexpr size_t BATCH_SIZE = 8 ;
375422 const int8_t * query = static_cast <const int8_t *>(query_ptr);
376423 const size_t dim = *((const size_t *)qty_ptr);
424+ // Precalculate q_correction once for the entire batch of DB vectors
425+ int64_t q_correction = 0 ;
426+ if constexpr (HasDualSimd || HasSimd) {
427+ __m256i q_comp = _mm256_setzero_si256 ();
428+ const size_t simd_limit = dim >= 32 ? (dim - 31 ) : 0 ;
429+ for (size_t offset = 0 ; offset < simd_limit; offset += 32 ) {
430+ __m256i q_raw = _mm256_loadu_si256 (reinterpret_cast <const __m256i*>(query + offset));
431+ q_comp = _mm256_add_epi32 (q_comp, _mm256_madd_epi16 (_mm256_cvtepi8_epi16 (_mm256_castsi256_si128 (q_raw)), _mm256_set1_epi16 (1 )));
432+ q_comp = _mm256_add_epi32 (q_comp, _mm256_madd_epi16 (_mm256_cvtepi8_epi16 (_mm256_extracti128_si256 (q_raw, 1 )), _mm256_set1_epi16 (1 )));
433+ }
434+ q_correction = int8_ip_hsum256 (q_comp) * 128 ;
435+ }
377436
378- auto batch_impl = [query, dim](const void * const * db, float * out_dists) DEGLIB_TARGET_AVX2_VNNI {
437+ auto batch_impl = [query, dim, q_correction ](const void * const * db, float * out_dists) DEGLIB_TARGET_AVX2_VNNI {
379438 const __m256i xor_mask = _mm256_set1_epi8 (static_cast <char >(0x80 ));
380439 __m256i acc[BATCH_SIZE ];
381440 for (size_t j = 0 ; j < BATCH_SIZE ; ++j) {
382441 acc[j] = _mm256_setzero_si256 ();
383442 }
384- __m256i q_comp = _mm256_setzero_si256 ();
385443
386444 size_t offset = 0 ;
387445 if constexpr (HasDualSimd || HasSimd) {
388446 const size_t simd_limit = dim >= 32 ? (dim - 31 ) : 0 ;
389447 for (; offset < simd_limit; offset += 32 ) {
390448 __m256i q_raw = _mm256_loadu_si256 (reinterpret_cast <const __m256i*>(query + offset));
391- q_comp = _mm256_add_epi32 (q_comp, _mm256_madd_epi16 (_mm256_cvtepi8_epi16 (_mm256_castsi256_si128 (q_raw)), _mm256_set1_epi16 (1 )));
392- q_comp = _mm256_add_epi32 (q_comp, _mm256_madd_epi16 (_mm256_cvtepi8_epi16 (_mm256_extracti128_si256 (q_raw, 1 )), _mm256_set1_epi16 (1 )));
393-
394449 for (size_t j = 0 ; j < BATCH_SIZE ; ++j) {
395450 const int8_t * db_ptr = static_cast <const int8_t *>(db[j]);
396451 __m256i r_raw = _mm256_loadu_si256 (reinterpret_cast <const __m256i*>(db_ptr + offset));
@@ -400,22 +455,13 @@ class InnerProductInt8_AVX2_VNNI {
400455 }
401456 }
402457
403- const int64_t q_correction = int8_ip_hsum256 (q_comp) * 128 ;
404-
405458 for (size_t j = 0 ; j < BATCH_SIZE ; ++j) {
406459 int64_t total = int8_ip_hsum256 (acc[j]) - q_correction;
407460 out_dists[j] = -static_cast <float >(total);
408461 }
409462
410463 if constexpr (HasTail) {
411- for (size_t j = 0 ; j < BATCH_SIZE ; ++j) {
412- const int8_t * db_ptr = static_cast <const int8_t *>(db[j]);
413- int64_t tail_sum = 0 ;
414- for (size_t k = offset; k < dim; ++k) {
415- tail_sum += int64_t (query[k]) * int64_t (db_ptr[k]);
416- }
417- out_dists[j] -= static_cast <float >(tail_sum);
418- }
464+ process_batch_tail<BATCH_SIZE >(query, db, offset, dim, out_dists);
419465 }
420466 };
421467
@@ -515,14 +561,7 @@ class InnerProductInt8_AVX2 {
515561 }
516562
517563 if constexpr (HasTail) {
518- for (size_t j = 0 ; j < BATCH_SIZE ; ++j) {
519- const int8_t * db_ptr = static_cast <const int8_t *>(db[j]);
520- int64_t tail_sum = 0 ;
521- for (size_t k = offset; k < dim; ++k) {
522- tail_sum += int64_t (query[k]) * int64_t (db_ptr[k]);
523- }
524- out_dists[j] -= static_cast <float >(tail_sum);
525- }
564+ process_batch_tail<BATCH_SIZE >(query, db, offset, dim, out_dists);
526565 }
527566 };
528567
0 commit comments