Skip to content

Commit 76d511c

Browse files
committed
cpp(distances): improve Int8/UInt8 IP SIMD kernels and builder regression dispatch
- Add and optimize VNNI/SIMD kernels for Int8 and UInt8 InnerProduct - Refactor builder regression tests to dispatch via InstructionSet enum - Update distance throughput benchmarks and unit tests
1 parent 39748fc commit 76d511c

8 files changed

Lines changed: 736 additions & 243 deletions

File tree

‎cpp/deglib/include/deglib/distance/int8_ip.h‎

Lines changed: 85 additions & 46 deletions
Original file line numberDiff line numberDiff line change
@@ -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
// ---------------------------------------------------------------------------------------------------------------------
62115
template <ResidualMode Mode = ResidualMode::Full>
63116
class 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
// ---------------------------------------------------------------------------------------------------------------------
302349
template <ResidualMode Mode = ResidualMode::Full>
303350
class 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

Comments
 (0)