Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -47,6 +47,7 @@ endfunction()
set(RABITQ_COMMON_SOURCES
src/utils/cpu_features.cpp
src/simd/dispatch.cpp
src/simd/space_float.cpp
)

set(RABITQ_AVX2_SOURCES
Expand Down
21 changes: 21 additions & 0 deletions include/rabitqlib/simd/space_dispatch.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,27 @@
#include <cstdint>

namespace rabitqlib::simd {
// Raw float32 distances accept arbitrary dimensions and unaligned inputs.
float euclidean_sqr(const float* a, const float* b, size_t dim);
float dot_product(const float* a, const float* b, size_t dim);
float dot_product_dis(const float* a, const float* b, size_t dim);
float l2norm_sqr(const float* a, size_t dim);

float euclidean_sqr_generic(const float* a, const float* b, size_t dim);
float dot_product_generic(const float* a, const float* b, size_t dim);
float dot_product_dis_generic(const float* a, const float* b, size_t dim);
float l2norm_sqr_generic(const float* a, size_t dim);

float euclidean_sqr_avx2(const float* a, const float* b, size_t dim);
float dot_product_avx2(const float* a, const float* b, size_t dim);
float dot_product_dis_avx2(const float* a, const float* b, size_t dim);
float l2norm_sqr_avx2(const float* a, size_t dim);

float euclidean_sqr_avx512(const float* a, const float* b, size_t dim);
float dot_product_avx512(const float* a, const float* b, size_t dim);
float dot_product_dis_avx512(const float* a, const float* b, size_t dim);
float l2norm_sqr_avx512(const float* a, size_t dim);

namespace excode_ipimpl {

float ip16_fxu1_avx2(
Expand Down
38 changes: 27 additions & 11 deletions include/rabitqlib/utils/space.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -83,31 +83,47 @@ inline void vec_rescale(T* data, size_t dim, T val) {

template <typename T>
inline T euclidean_sqr(const T* __restrict__ vec0, const T* __restrict__ vec1, size_t dim) {
ConstVectorMap<T> v0(vec0, dim);
ConstVectorMap<T> v1(vec1, dim);
return (v0 - v1).dot(v0 - v1);
if constexpr (std::is_same_v<T, float>) {
return simd::euclidean_sqr(vec0, vec1, dim);
} else {
ConstVectorMap<T> v0(vec0, dim);
ConstVectorMap<T> v1(vec1, dim);
return (v0 - v1).dot(v0 - v1);
}
}

template <typename T>
inline T dot_product_dis(
const T* __restrict__ vec0, const T* __restrict__ vec1, size_t dim
) {
ConstVectorMap<T> v0(vec0, dim);
ConstVectorMap<T> v1(vec1, dim);
return 1 - v0.dot(v1);
if constexpr (std::is_same_v<T, float>) {
return simd::dot_product_dis(vec0, vec1, dim);
} else {
ConstVectorMap<T> v0(vec0, dim);
ConstVectorMap<T> v1(vec1, dim);
return 1 - v0.dot(v1);
}
}

template <typename T>
inline T l2norm_sqr(const T* __restrict__ vec0, size_t dim) {
ConstVectorMap<T> v0(vec0, dim);
return v0.dot(v0);
if constexpr (std::is_same_v<T, float>) {
return simd::l2norm_sqr(vec0, dim);
} else {
ConstVectorMap<T> v0(vec0, dim);
return v0.dot(v0);
}
}

template <typename T>
inline T dot_product(const T* __restrict__ vec0, const T* __restrict__ vec1, size_t dim) {
ConstVectorMap<T> v0(vec0, dim);
ConstVectorMap<T> v1(vec1, dim);
return v0.dot(v1);
if constexpr (std::is_same_v<T, float>) {
return simd::dot_product(vec0, vec1, dim);
} else {
ConstVectorMap<T> v0(vec0, dim);
ConstVectorMap<T> v1(vec1, dim);
return v0.dot(v1);
}
}

template <typename T>
Expand Down
60 changes: 33 additions & 27 deletions src/index/hnsw_search_avx2_kernels.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@

#include <cstddef>
#include <cstdint>
#include <cstring>

#include "rabitqlib/index/query.hpp"
#include "rabitqlib/simd/space_dispatch.hpp"
Expand Down Expand Up @@ -68,40 +69,45 @@ static inline float hnsw_mask_ip_x0_q_avx2(
const float* query, const uint64_t* data, size_t padded_dim
) {
const size_t num_blk = padded_dim / 64;
const uint8_t* it_data = reinterpret_cast<const uint8_t*>(data);
const auto* it_data = reinterpret_cast<const uint8_t*>(data);
const float* it_query = query;

__m256 sum = _mm256_setzero_ps();
__m256i bit_checker = _mm256_set_epi32(0x80, 0x40, 0x20, 0x10, 0x08, 0x04, 0x02, 0x01);
const __m256i shifts0 = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7);
const __m256i shifts1 = _mm256_setr_epi32(8, 9, 10, 11, 12, 13, 14, 15);
const __m256i shifts2 = _mm256_setr_epi32(16, 17, 18, 19, 20, 21, 22, 23);
const __m256i shifts3 = _mm256_setr_epi32(24, 25, 26, 27, 28, 29, 30, 31);
__m256 sum0 = _mm256_setzero_ps();
__m256 sum1 = _mm256_setzero_ps();
__m256 sum2 = _mm256_setzero_ps();
__m256 sum3 = _mm256_setzero_ps();

for (size_t i = 0; i < num_blk; ++i) {
uint64_t bits = rabitqlib::reverse_bits_u64(rabitqlib::load_unaligned_u64(it_data));

// 64 bits / 8 floats = 8 iterations
for (int j = 0; j < 8; ++j) {
uint8_t current_byte = static_cast<uint8_t>(bits >> (j * 8));
__m256i v_byte = _mm256_set1_epi32(current_byte);
__m256i masked_bits = _mm256_and_si256(v_byte, bit_checker);
__m256i mask = _mm256_cmpgt_epi32(masked_bits, _mm256_setzero_si256());

__m256 q_vals = _mm256_loadu_ps(it_query);
__m256 masked = _mm256_and_ps(q_vals, _mm256_castsi256_ps(mask));

sum = _mm256_add_ps(sum, masked);

it_query += 8;
// Stored coordinates run from bit 63 to bit 0: high word first,
// with each selected bit shifted into a maskload lane's sign bit.
for (size_t half = 0; half < 2; ++half) {
int32_t word;
std::memcpy(&word, it_data + (1 - half) * sizeof(word), sizeof(word));
const __m256i bits = _mm256_set1_epi32(word);
sum0 = _mm256_add_ps(
sum0, _mm256_maskload_ps(it_query, _mm256_sllv_epi32(bits, shifts0))
);
sum1 = _mm256_add_ps(
sum1, _mm256_maskload_ps(it_query + 8, _mm256_sllv_epi32(bits, shifts1))
);
sum2 = _mm256_add_ps(
sum2, _mm256_maskload_ps(it_query + 16, _mm256_sllv_epi32(bits, shifts2))
);
sum3 = _mm256_add_ps(
sum3, _mm256_maskload_ps(it_query + 24, _mm256_sllv_epi32(bits, shifts3))
);
it_query += 32;
}
it_data += sizeof(uint64_t);
}

alignas(32) float lanes[8];
_mm256_store_ps(lanes, sum);

float result = 0.0f;
for (float lane : lanes) {
result += lane;
}
return result;
const __m256 sum = _mm256_add_ps(_mm256_add_ps(sum0, sum1), _mm256_add_ps(sum2, sum3));
__m128 lanes = _mm_add_ps(_mm256_castps256_ps128(sum), _mm256_extractf128_ps(sum, 1));
lanes = _mm_add_ps(lanes, _mm_movehl_ps(lanes, lanes));
return _mm_cvtss_f32(_mm_add_ss(lanes, _mm_movehdup_ps(lanes)));
}

static inline float hnsw_warmup_ip_x0_q_512_avx2(
Expand Down
27 changes: 27 additions & 0 deletions src/simd/dispatch.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,33 @@

namespace rabitqlib::simd {

const auto kEuclideanSqrFn = cpu::has_avx512_core() ? euclidean_sqr_avx512
: cpu::has_avx2() ? euclidean_sqr_avx2
: euclidean_sqr_generic;
const auto kDotProductFn = cpu::has_avx512_core() ? dot_product_avx512
: cpu::has_avx2() ? dot_product_avx2
: dot_product_generic;
const auto kDotProductDisFn = cpu::has_avx512_core() ? dot_product_dis_avx512
: cpu::has_avx2() ? dot_product_dis_avx2
: dot_product_dis_generic;
const auto kL2normSqrFn = cpu::has_avx512_core() ? l2norm_sqr_avx512
: cpu::has_avx2() ? l2norm_sqr_avx2
: l2norm_sqr_generic;

float euclidean_sqr(const float* a, const float* b, size_t dim) {
return kEuclideanSqrFn(a, b, dim);
}

float dot_product(const float* a, const float* b, size_t dim) {
return kDotProductFn(a, b, dim);
}

float dot_product_dis(const float* a, const float* b, size_t dim) {
return kDotProductDisFn(a, b, dim);
}

float l2norm_sqr(const float* a, size_t dim) { return kL2normSqrFn(a, dim); }

namespace detail {

RescaleScratch& get_thread_local_rescale_scratch(size_t dim) {
Expand Down
75 changes: 50 additions & 25 deletions src/simd/space_avx2.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -2,11 +2,29 @@

#include <cmath>
#include <cstdint>
#include <cstring>

#include "rabitqlib/utils/space.hpp"
#include "space_float_kernels.hpp"

namespace rabitqlib::simd {

float euclidean_sqr_avx2(const float* a, const float* b, size_t dim) {
return raw_float<FloatOperation::SquaredL2>(a, b, dim);
}

float dot_product_avx2(const float* a, const float* b, size_t dim) {
return raw_float<FloatOperation::Dot>(a, b, dim);
}

float dot_product_dis_avx2(const float* a, const float* b, size_t dim) {
return raw_float<FloatOperation::InnerProductDistance>(a, b, dim);
}

float l2norm_sqr_avx2(const float* a, size_t dim) {
return raw_float<FloatOperation::SquaredNorm>(a, a, dim);
}

void scalar_quantize_uint8_avx2(
uint8_t* result, const float* vec0, size_t dim, float lo, float delta
) {
Expand Down Expand Up @@ -153,38 +171,45 @@ void new_transpose_bin_512_avx2(

float mask_ip_x0_q_avx2(const float* query, const uint64_t* data, size_t padded_dim) {
const size_t num_blk = padded_dim / 64;
const uint8_t* it_data = reinterpret_cast<const uint8_t*>(data);
const auto* it_data = reinterpret_cast<const uint8_t*>(data);
const float* it_query = query;

__m256 sum = _mm256_setzero_ps();

__m256i bit_checker = _mm256_set_epi32(0x80, 0x40, 0x20, 0x10, 0x08, 0x04, 0x02, 0x01);
const __m256i shifts0 = _mm256_setr_epi32(0, 1, 2, 3, 4, 5, 6, 7);
const __m256i shifts1 = _mm256_setr_epi32(8, 9, 10, 11, 12, 13, 14, 15);
const __m256i shifts2 = _mm256_setr_epi32(16, 17, 18, 19, 20, 21, 22, 23);
const __m256i shifts3 = _mm256_setr_epi32(24, 25, 26, 27, 28, 29, 30, 31);
__m256 sum0 = _mm256_setzero_ps();
__m256 sum1 = _mm256_setzero_ps();
__m256 sum2 = _mm256_setzero_ps();
__m256 sum3 = _mm256_setzero_ps();

for (size_t i = 0; i < num_blk; ++i) {
uint64_t bits = reverse_bits_u64(load_unaligned_u64(it_data));

// 64 bits / 8 floats = 8 iterations
for (int j = 0; j < 8; ++j) {
uint8_t current_byte = static_cast<uint8_t>(bits >> (j * 8));
__m256i v_byte = _mm256_set1_epi32(current_byte);
__m256i masked_bits = _mm256_and_si256(v_byte, bit_checker);
__m256i mask = _mm256_cmpgt_epi32(masked_bits, _mm256_setzero_si256());

__m256 q_vals = _mm256_loadu_ps(it_query);
__m256 masked = _mm256_and_ps(q_vals, _mm256_castsi256_ps(mask));

sum = _mm256_add_ps(sum, masked);

it_query += 8;
// Stored coordinates run from bit 63 to bit 0: high word first,
// with each selected bit shifted into a maskload lane's sign bit.
for (size_t half = 0; half < 2; ++half) {
int32_t word;
std::memcpy(&word, it_data + (1 - half) * sizeof(word), sizeof(word));
const __m256i bits = _mm256_set1_epi32(word);
sum0 = _mm256_add_ps(
sum0, _mm256_maskload_ps(it_query, _mm256_sllv_epi32(bits, shifts0))
);
sum1 = _mm256_add_ps(
sum1, _mm256_maskload_ps(it_query + 8, _mm256_sllv_epi32(bits, shifts1))
);
sum2 = _mm256_add_ps(
sum2, _mm256_maskload_ps(it_query + 16, _mm256_sllv_epi32(bits, shifts2))
);
sum3 = _mm256_add_ps(
sum3, _mm256_maskload_ps(it_query + 24, _mm256_sllv_epi32(bits, shifts3))
);
it_query += 32;
}
it_data += sizeof(uint64_t);
}

float result = 0.0f;
for (int i = 0; i < 8; ++i) {
result += reinterpret_cast<float*>(&sum)[i];
}
return result;
const __m256 sum = _mm256_add_ps(_mm256_add_ps(sum0, sum1), _mm256_add_ps(sum2, sum3));
__m128 lanes = _mm_add_ps(_mm256_castps256_ps128(sum), _mm256_extractf128_ps(sum, 1));
lanes = _mm_add_ps(lanes, _mm_movehl_ps(lanes, lanes));
return _mm_cvtss_f32(_mm_add_ss(lanes, _mm_movehdup_ps(lanes)));
}

} // namespace rabitqlib::simd
17 changes: 17 additions & 0 deletions src/simd/space_avx512.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -4,9 +4,26 @@
#include <cstdint>

#include "rabitqlib/utils/space.hpp"
#include "space_float_kernels.hpp"

namespace rabitqlib::simd {

float euclidean_sqr_avx512(const float* a, const float* b, size_t dim) {
return raw_float<FloatOperation::SquaredL2>(a, b, dim);
}

float dot_product_avx512(const float* a, const float* b, size_t dim) {
return raw_float<FloatOperation::Dot>(a, b, dim);
}

float dot_product_dis_avx512(const float* a, const float* b, size_t dim) {
return raw_float<FloatOperation::InnerProductDistance>(a, b, dim);
}

float l2norm_sqr_avx512(const float* a, size_t dim) {
return raw_float<FloatOperation::SquaredNorm>(a, a, dim);
}

void scalar_quantize_uint8_avx512(
uint8_t* result, const float* vec0, size_t dim, float lo, float delta
) {
Expand Down
Loading