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
51 changes: 33 additions & 18 deletions src/simd/space_excode_avx2.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -2,13 +2,28 @@

#include <array>
#include <cstdint>
#include <cstdlib>
#include <iostream>
#include <cstring>

#include "rabitqlib/utils/space.hpp"
#include "rabitqlib/simd/space_dispatch.hpp"

namespace rabitqlib::simd::excode_ipimpl {

namespace {
[[nodiscard]] inline uint64_t load_u64(const uint8_t* data) noexcept {
uint64_t value = 0;
std::memcpy(&value, data, sizeof(value));
return value;
}

[[nodiscard]] inline __m128i set_u64x(uint64_t high, uint64_t low) noexcept {
int64_t signed_high = 0;
int64_t signed_low = 0;
std::memcpy(&signed_high, &high, sizeof(signed_high));
std::memcpy(&signed_low, &low, sizeof(signed_low));
return _mm_set_epi64x(signed_high, signed_low);
}
} // namespace

// helper function for AVX2 inner product
inline void contribute_ip(__m128i vec, const float* __restrict__ query, __m256& sum) {
__m256 q = _mm256_loadu_ps(query);
Expand Down Expand Up @@ -113,7 +128,7 @@ float ip64_fxu3_avx2(
__m128i compact2 = _mm_loadu_si128(reinterpret_cast<const __m128i*>(compact_code));
compact_code += 16;

int64_t top_bit = *reinterpret_cast<const int64_t*>(compact_code);
const uint64_t top_bit = load_u64(compact_code);
compact_code += 8;

__m128i vec_00_to_15 = _mm_and_si128(compact2, mask);
Expand All @@ -122,13 +137,13 @@ float ip64_fxu3_avx2(
__m128i vec_48_to_63 = _mm_and_si128(_mm_srli_epi16(compact2, 6), mask);

__m128i top_00_to_15 =
_mm_and_si128(_mm_set_epi64x(top_bit << 1, top_bit << 2), top_mask);
_mm_and_si128(set_u64x(top_bit << 1, top_bit << 2), top_mask);
__m128i top_16_to_31 =
_mm_and_si128(_mm_set_epi64x(top_bit >> 1, top_bit >> 0), top_mask);
_mm_and_si128(set_u64x(top_bit >> 1, top_bit >> 0), top_mask);
__m128i top_32_to_47 =
_mm_and_si128(_mm_set_epi64x(top_bit >> 3, top_bit >> 2), top_mask);
_mm_and_si128(set_u64x(top_bit >> 3, top_bit >> 2), top_mask);
__m128i top_48_to_63 =
_mm_and_si128(_mm_set_epi64x(top_bit >> 5, top_bit >> 4), top_mask);
_mm_and_si128(set_u64x(top_bit >> 5, top_bit >> 4), top_mask);

vec_00_to_15 = _mm_or_si128(top_00_to_15, vec_00_to_15);
vec_16_to_31 = _mm_or_si128(top_16_to_31, vec_16_to_31);
Expand Down Expand Up @@ -183,7 +198,7 @@ float ip64_fxu5_avx2(
_mm_loadu_si128(reinterpret_cast<const __m128i*>(compact_code + 16));
compact_code += 32;

int64_t top_bit = *reinterpret_cast<const int64_t*>(compact_code);
const uint64_t top_bit = load_u64(compact_code);
compact_code += 8;

__m128i vec_00_to_15 = _mm_and_si128(compact4_1, mask);
Expand All @@ -192,13 +207,13 @@ float ip64_fxu5_avx2(
__m128i vec_48_to_63 = _mm_and_si128(_mm_srli_epi16(compact4_2, 4), mask);

__m128i top_00_to_15 =
_mm_and_si128(_mm_set_epi64x(top_bit << 3, top_bit << 4), top_mask);
_mm_and_si128(set_u64x(top_bit << 3, top_bit << 4), top_mask);
__m128i top_16_to_31 =
_mm_and_si128(_mm_set_epi64x(top_bit << 1, top_bit << 2), top_mask);
_mm_and_si128(set_u64x(top_bit << 1, top_bit << 2), top_mask);
__m128i top_32_to_47 =
_mm_and_si128(_mm_set_epi64x(top_bit >> 1, top_bit >> 0), top_mask);
_mm_and_si128(set_u64x(top_bit >> 1, top_bit >> 0), top_mask);
__m128i top_48_to_63 =
_mm_and_si128(_mm_set_epi64x(top_bit >> 3, top_bit >> 2), top_mask);
_mm_and_si128(set_u64x(top_bit >> 3, top_bit >> 2), top_mask);

vec_00_to_15 = _mm_or_si128(top_00_to_15, vec_00_to_15);
vec_16_to_31 = _mm_or_si128(top_16_to_31, vec_16_to_31);
Expand Down Expand Up @@ -279,17 +294,17 @@ float ip64_fxu7_avx2(
_mm_srli_epi16(_mm_and_si128(cpt3, mask2), 2)
);

int64_t top_bit = *reinterpret_cast<const int64_t*>(compact_code);
const uint64_t top_bit = load_u64(compact_code);
compact_code += 8;

__m128i top_00_to_15 =
_mm_and_si128(_mm_set_epi64x(top_bit << 5, top_bit << 6), top_mask);
_mm_and_si128(set_u64x(top_bit << 5, top_bit << 6), top_mask);
__m128i top_16_to_31 =
_mm_and_si128(_mm_set_epi64x(top_bit << 3, top_bit << 4), top_mask);
_mm_and_si128(set_u64x(top_bit << 3, top_bit << 4), top_mask);
__m128i top_32_to_47 =
_mm_and_si128(_mm_set_epi64x(top_bit << 1, top_bit << 2), top_mask);
_mm_and_si128(set_u64x(top_bit << 1, top_bit << 2), top_mask);
__m128i top_48_to_63 =
_mm_and_si128(_mm_set_epi64x(top_bit >> 1, top_bit << 0), top_mask);
_mm_and_si128(set_u64x(top_bit >> 1, top_bit << 0), top_mask);

vec_00_to_15 = _mm_or_si128(top_00_to_15, vec_00_to_15);
vec_16_to_31 = _mm_or_si128(top_16_to_31, vec_16_to_31);
Expand Down
52 changes: 33 additions & 19 deletions src/simd/space_excode_avx512.cpp
Original file line number Diff line number Diff line change
@@ -1,14 +1,28 @@
#include <immintrin.h>

#include <array>
#include <cstdint>
#include <cstdlib>
#include <iostream>
#include <cstring>

#include "rabitqlib/utils/space.hpp"
#include "rabitqlib/simd/space_dispatch.hpp"

namespace rabitqlib::simd::excode_ipimpl {

namespace {
[[nodiscard]] inline uint64_t load_u64(const uint8_t* data) noexcept {
uint64_t value = 0;
std::memcpy(&value, data, sizeof(value));
return value;
}

[[nodiscard]] inline __m128i set_u64x(uint64_t high, uint64_t low) noexcept {
int64_t signed_high = 0;
int64_t signed_low = 0;
std::memcpy(&signed_high, &high, sizeof(signed_high));
std::memcpy(&signed_low, &low, sizeof(signed_low));
return _mm_set_epi64x(signed_high, signed_low);
}
} // namespace

// ip16: this function is used to compute inner product of
// vectors padded to multiple of 16
// fxu1: the inner product is computed between float and 1-bit unsigned int (lay out can be
Expand Down Expand Up @@ -89,7 +103,7 @@ float ip64_fxu3_avx512(
__m128i compact2 = _mm_loadu_si128(reinterpret_cast<const __m128i*>(compact_code));
compact_code += 16;

int64_t top_bit = *reinterpret_cast<const int64_t*>(compact_code);
const uint64_t top_bit = load_u64(compact_code);
compact_code += 8;

__m128i vec_00_to_15 = _mm_and_si128(compact2, mask);
Expand All @@ -98,13 +112,13 @@ float ip64_fxu3_avx512(
__m128i vec_48_to_63 = _mm_and_si128(_mm_srli_epi16(compact2, 6), mask);

__m128i top_00_to_15 =
_mm_and_si128(_mm_set_epi64x(top_bit << 1, top_bit << 2), top_mask);
_mm_and_si128(set_u64x(top_bit << 1, top_bit << 2), top_mask);
__m128i top_16_to_31 =
_mm_and_si128(_mm_set_epi64x(top_bit >> 1, top_bit >> 0), top_mask);
_mm_and_si128(set_u64x(top_bit >> 1, top_bit >> 0), top_mask);
__m128i top_32_to_47 =
_mm_and_si128(_mm_set_epi64x(top_bit >> 3, top_bit >> 2), top_mask);
_mm_and_si128(set_u64x(top_bit >> 3, top_bit >> 2), top_mask);
__m128i top_48_to_63 =
_mm_and_si128(_mm_set_epi64x(top_bit >> 5, top_bit >> 4), top_mask);
_mm_and_si128(set_u64x(top_bit >> 5, top_bit >> 4), top_mask);

vec_00_to_15 = _mm_or_si128(top_00_to_15, vec_00_to_15);
vec_16_to_31 = _mm_or_si128(top_16_to_31, vec_16_to_31);
Expand Down Expand Up @@ -175,7 +189,7 @@ float ip64_fxu5_avx512(
_mm_loadu_si128(reinterpret_cast<const __m128i*>(compact_code + 16));
compact_code += 32;

int64_t top_bit = *reinterpret_cast<const int64_t*>(compact_code);
const uint64_t top_bit = load_u64(compact_code);
compact_code += 8;

__m128i vec_00_to_15 = _mm_and_si128(compact4_1, mask);
Expand All @@ -184,13 +198,13 @@ float ip64_fxu5_avx512(
__m128i vec_48_to_63 = _mm_and_si128(_mm_srli_epi16(compact4_2, 4), mask);

__m128i top_00_to_15 =
_mm_and_si128(_mm_set_epi64x(top_bit << 3, top_bit << 4), top_mask);
_mm_and_si128(set_u64x(top_bit << 3, top_bit << 4), top_mask);
__m128i top_16_to_31 =
_mm_and_si128(_mm_set_epi64x(top_bit << 1, top_bit << 2), top_mask);
_mm_and_si128(set_u64x(top_bit << 1, top_bit << 2), top_mask);
__m128i top_32_to_47 =
_mm_and_si128(_mm_set_epi64x(top_bit >> 1, top_bit >> 0), top_mask);
_mm_and_si128(set_u64x(top_bit >> 1, top_bit >> 0), top_mask);
__m128i top_48_to_63 =
_mm_and_si128(_mm_set_epi64x(top_bit >> 3, top_bit >> 2), top_mask);
_mm_and_si128(set_u64x(top_bit >> 3, top_bit >> 2), top_mask);

vec_00_to_15 = _mm_or_si128(top_00_to_15, vec_00_to_15);
vec_16_to_31 = _mm_or_si128(top_16_to_31, vec_16_to_31);
Expand Down Expand Up @@ -299,17 +313,17 @@ float ip64_fxu7_avx512(
_mm_srli_epi16(_mm_and_si128(cpt3, mask2), 2)
);

int64_t top_bit = *reinterpret_cast<const int64_t*>(compact_code);
const uint64_t top_bit = load_u64(compact_code);
compact_code += 8;

__m128i top_00_to_15 =
_mm_and_si128(_mm_set_epi64x(top_bit << 5, top_bit << 6), top_mask);
_mm_and_si128(set_u64x(top_bit << 5, top_bit << 6), top_mask);
__m128i top_16_to_31 =
_mm_and_si128(_mm_set_epi64x(top_bit << 3, top_bit << 4), top_mask);
_mm_and_si128(set_u64x(top_bit << 3, top_bit << 4), top_mask);
__m128i top_32_to_47 =
_mm_and_si128(_mm_set_epi64x(top_bit << 1, top_bit << 2), top_mask);
_mm_and_si128(set_u64x(top_bit << 1, top_bit << 2), top_mask);
__m128i top_48_to_63 =
_mm_and_si128(_mm_set_epi64x(top_bit >> 1, top_bit << 0), top_mask);
_mm_and_si128(set_u64x(top_bit >> 1, top_bit << 0), top_mask);

vec_00_to_15 = _mm_or_si128(top_00_to_15, vec_00_to_15);
vec_16_to_31 = _mm_or_si128(top_16_to_31, vec_16_to_31);
Expand Down
80 changes: 74 additions & 6 deletions tests/unit/rabitqlib/utils/space_test.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -2,17 +2,16 @@

#include <gtest/gtest.h>

#include <array>
#include <cmath>
#include <cstdlib>
#include <vector>

#include "rabitqlib/defines.hpp"
#include "rabitqlib/simd/pack_excode_dispatch.hpp"
#include "rabitqlib/simd/space_dispatch.hpp"
#include "rabitqlib/utils/cpu_features.hpp"
#include "test_data.hpp"
#include "test_helpers.hpp"

using namespace rabitqlib;
using namespace rabitq_test;

TEST(Select_IP_Func, returns_stable_function_pointer) {
auto ip_func = select_excode_ipfunc(0);
Expand Down Expand Up @@ -98,7 +97,7 @@ TEST(ScalarQuantize, Uint16MatchesRoundedScalar) {

TEST(ip16_fxu1_avx, ip_works) {
srand(42);
size_t dim = 64;
constexpr size_t dim = 64;
float query[dim];
uint8_t codes[dim / 8];

Expand All @@ -117,7 +116,7 @@ TEST(ip16_fxu1_avx, ip_works) {

TEST(ip64_fxu2_avx, ip_works) {
srand(42);
size_t dim = 64 * 4;
constexpr size_t dim = 64 * 4;
float query[dim];
uint8_t codes[dim / 4];

Expand All @@ -133,6 +132,75 @@ TEST(ip64_fxu2_avx, ip_works) {
);
}

TEST(OddBitExcodeIp, MatchesScalarInnerProduct) {
constexpr size_t dim = 64 * 4;
std::vector<float> query(dim);
std::vector<uint8_t> codes(dim);

for (size_t bits : std::array<size_t, 3>{3, 5, 7}) {
const uint8_t max_code = static_cast<uint8_t>((1U << bits) - 1U);
for (size_t i = 0; i < dim; ++i) {
query[i] = static_cast<float>(static_cast<int>(i % 23) - 11) / 7.0F;
codes[i] = static_cast<uint8_t>((i * 37U + 19U) & max_code);
}
// Exercise the high bit of each packed 64-value block, including bit 63 of the
// scalar word used by the SIMD unpacking path.
for (size_t i = 63; i < dim; i += 64) {
codes[i] = max_code;
}

std::vector<uint8_t> compact(dim * bits / 8);
if (bits == 3) {
simd::packing_3bit_excode(codes.data(), compact.data(), dim);
} else if (bits == 5) {
simd::packing_5bit_excode(codes.data(), compact.data(), dim);
} else {
simd::packing_7bit_excode(codes.data(), compact.data(), dim);
}

double expected = 0.0;
for (size_t i = 0; i < dim; ++i) {
expected += static_cast<double>(query[i]) * static_cast<double>(codes[i]);
}
const float expected_float = static_cast<float>(expected);

if (cpu::has_avx2()) {
const std::array<ex_ipfunc, 8> avx2_functions{
nullptr,
simd::excode_ipimpl::ip16_fxu1_avx2,
simd::excode_ipimpl::ip64_fxu2_avx2,
simd::excode_ipimpl::ip64_fxu3_avx2,
simd::excode_ipimpl::ip16_fxu4_avx2,
simd::excode_ipimpl::ip64_fxu5_avx2,
simd::excode_ipimpl::ip64_fxu6_avx2,
simd::excode_ipimpl::ip64_fxu7_avx2,
};
ASSERT_NEAR(
avx2_functions[bits](query.data(), compact.data(), dim),
expected_float,
0.1F
);
}
if (cpu::has_avx512_core()) {
const std::array<ex_ipfunc, 8> avx512_functions{
nullptr,
simd::excode_ipimpl::ip16_fxu1_avx512,
simd::excode_ipimpl::ip64_fxu2_avx512,
simd::excode_ipimpl::ip64_fxu3_avx512,
simd::excode_ipimpl::ip16_fxu4_avx512,
simd::excode_ipimpl::ip64_fxu5_avx512,
simd::excode_ipimpl::ip64_fxu6_avx512,
simd::excode_ipimpl::ip64_fxu7_avx512,
};
ASSERT_NEAR(
avx512_functions[bits](query.data(), compact.data(), dim),
expected_float,
0.1F
);
}
}
}

TEST(ip_fxu8_avx, ip_works) {
constexpr size_t dim = 1024;
std::vector<float> query(dim);
Expand Down
Loading