Skip to content
Open
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
5 changes: 5 additions & 0 deletions include/filter_match_proxy.h
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
#pragma once
#include "label_bitmask.h"
#include "integer_label_vector.h"
#include <xmmintrin.h> // _mm_prefetch

namespace diskann
{
Expand All @@ -9,6 +10,7 @@ namespace diskann
{
public:
virtual bool contain_filtered_label(uint32_t id) = 0;
virtual void prefetch_bitmask(uint32_t id) = 0;
};

template <typename LabelT>
Expand All @@ -21,6 +23,7 @@ namespace diskann
LabelT unv_label);

virtual bool contain_filtered_label(uint32_t id) override;
virtual void prefetch_bitmask(uint32_t id) override;

private:
simple_bitmask_buf& _bitmask_filters;
Expand All @@ -37,6 +40,7 @@ namespace diskann
LabelT unv_label);

virtual bool contain_filtered_label(uint32_t id) override;
virtual void prefetch_bitmask(uint32_t id) override;

private:
integer_label_vector& _label_vector;
Expand All @@ -56,6 +60,7 @@ class label_filter_match_holder : public filter_match_proxy
bool use_integer_labels);

virtual bool contain_filtered_label(uint32_t id) override;
virtual void prefetch_bitmask(uint32_t id) override;

private:
bitmask_filter_match<LabelT> _bitmask_filter_match;
Expand Down
32 changes: 30 additions & 2 deletions src/filter_match_proxy.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -15,7 +15,9 @@ bitmask_filter_match<LabelT>::bitmask_filter_match(
// _bitmask_size == 0 means no filter is set
if (_bitmask_filters._bitmask_size > 0)
{
query_bitmask_buf.resize(_bitmask_filters._bitmask_size, 0);
// Pad to at least 4 words (32 bytes) for safe AVX2 256-bit loads
size_t padded_size = std::max(_bitmask_filters._bitmask_size, (std::uint64_t)4);
query_bitmask_buf.resize(padded_size, 0);
_bitmask_full_val._mask = query_bitmask_buf.data();

for (const auto& filter_label : filter_labels)
Expand All @@ -38,6 +40,16 @@ bool bitmask_filter_match<LabelT>::contain_filtered_label(uint32_t id)
return bm.test_full_mask_val(_bitmask_full_val);
}

template <typename LabelT>
void bitmask_filter_match<LabelT>::prefetch_bitmask(uint32_t id)
{
// Prefetch the bitmask for id to L1 cache to hide DRAM latency
if (_bitmask_filters._bitmask_size > 0)
{
_mm_prefetch(reinterpret_cast<const char*>(_bitmask_filters.get_bitmask(id)), _MM_HINT_T0);
}
}

template <typename LabelT>
integer_label_filter_match<LabelT>::integer_label_filter_match(
integer_label_vector& label_vector,
Expand All @@ -53,10 +65,17 @@ template <typename LabelT>
bool integer_label_filter_match<LabelT>::contain_filtered_label(uint32_t id)
{
// if unv isn't set, it will be default value 0, and there will be no match
return _label_vector.check_label_exists(id, _filter_labels)
return _label_vector.check_label_exists(id, _filter_labels)
|| _label_vector.check_label_exists(id, _unv_label);
}

template <typename LabelT>
void integer_label_filter_match<LabelT>::prefetch_bitmask(uint32_t id)
{
// No-op for integer labels (no bitmask to prefetch)
(void)id;
}

template <typename LabelT>
label_filter_match_holder<LabelT>::label_filter_match_holder(simple_bitmask_buf& bitmask_filters,
std::vector<std::uint64_t>& query_bitmask_buf,
Expand All @@ -83,6 +102,15 @@ bool label_filter_match_holder<LabelT>::contain_filtered_label(uint32_t id)
}
}

template <typename LabelT>
void label_filter_match_holder<LabelT>::prefetch_bitmask(uint32_t id)
{
if (!_use_integer_labels)
{
_bitmask_filter_match.prefetch_bitmask(id);
}
}

template class bitmask_filter_match<uint16_t>;
template class bitmask_filter_match<uint32_t>;
template class integer_label_filter_match<uint16_t>;
Expand Down
49 changes: 46 additions & 3 deletions src/index.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -1057,10 +1057,31 @@ std::pair<uint32_t, uint32_t> Index<T, TagT, LabelT>::iterate_to_fixed_point(
{
LockGuard guard(_locks[n]);
auto neighbour_list = _graph_store->get_neighbours(n);
for (auto id : neighbour_list)
const location_t* neighbour_data = neighbour_list.data();
const size_t nbrs_count = neighbour_list.size();
constexpr size_t BITMASK_PREFETCH_K = 8;

// Pre-prefetch bitmasks for first K neighbors (only if filtering)
if (use_filter)
{
const size_t prefetch_init = std::min(BITMASK_PREFETCH_K, nbrs_count);
for (size_t p = 0; p < prefetch_init; ++p)
{
match_proxy.prefetch_bitmask(neighbour_data[p]);
}
}

for (size_t i = 0; i < nbrs_count; ++i)
{
auto id = neighbour_data[i];
assert(id < _max_points);

// Prefetch bitmask K steps ahead (sliding window)
if (use_filter && i + BITMASK_PREFETCH_K < nbrs_count)
{
match_proxy.prefetch_bitmask(neighbour_data[i + BITMASK_PREFETCH_K]);
}

if (!is_not_visited(id))
{
continue;
Expand Down Expand Up @@ -1094,10 +1115,31 @@ std::pair<uint32_t, uint32_t> Index<T, TagT, LabelT>::iterate_to_fixed_point(
// mark visited and collect unvisited into id_scratch
_locks[n].lock_shared();
auto nbrs = _graph_store->get_neighbours(n);
for (auto id : nbrs)
const location_t* nbrs_data = nbrs.data();
const size_t nbrs_count = nbrs.size();
constexpr size_t BITMASK_PREFETCH_K = 8;

// Pre-prefetch bitmasks for first K neighbors (only if filtering)
if (use_filter)
{
const size_t prefetch_init = std::min(BITMASK_PREFETCH_K, nbrs_count);
for (size_t p = 0; p < prefetch_init; ++p)
{
match_proxy.prefetch_bitmask(nbrs_data[p]);
}
}

for (size_t i = 0; i < nbrs_count; ++i)
{
auto id = nbrs_data[i];
assert(id < _max_points);

// Prefetch bitmask K steps ahead (sliding window)
if (use_filter && i + BITMASK_PREFETCH_K < nbrs_count)
{
match_proxy.prefetch_bitmask(nbrs_data[i + BITMASK_PREFETCH_K]);
}

if (!is_not_visited(id))
{
continue;
Expand Down Expand Up @@ -2296,7 +2338,8 @@ template <typename T, typename TagT, typename LabelT>
void Index<T, TagT, LabelT>::convert_pts_label_to_bitmask(std::vector<std::vector<LabelT>>& pts_to_labels, simple_bitmask_buf& bitmask_buf, size_t num_labels)
{
_bitmask_buf._bitmask_size = simple_bitmask::get_bitmask_size(num_labels + 1);
_bitmask_buf._buf.resize(pts_to_labels.size() * _bitmask_buf._bitmask_size, 0);
// Add 4 extra uint64 words at end for safe AVX2 256-bit loads on last node
_bitmask_buf._buf.resize(pts_to_labels.size() * _bitmask_buf._bitmask_size + 4, 0);

for (size_t i = 0; i < pts_to_labels.size(); i++)
{
Expand Down
42 changes: 41 additions & 1 deletion src/label_bitmask.cpp
Original file line number Diff line number Diff line change
@@ -1,5 +1,9 @@
#include "label_bitmask.h"

#ifdef _WINDOWS
#include <immintrin.h>
#endif

namespace diskann
{

Expand Down Expand Up @@ -42,15 +46,51 @@ bool simple_bitmask::test_mask_val(const simple_bitmask_val& bitmask_val) const

bool simple_bitmask::test_full_mask_val(const simple_bitmask_full_val& bitmask_full_val) const
{
#if defined(_WINDOWS) && defined(USE_AVX2)
// AVX2 branchless bitmask intersection test.
// Eliminates per-word branches that cause misprediction overhead.
// Handles up to 4 uint64 words (256 bits) in a single SIMD operation.
const std::uint64_t* query = bitmask_full_val._mask;
const std::uint64_t* node = _bitsets;

if (_bitmask_size <= 4)
{
// Fast path: load up to 256 bits, AND, test if any bit set.
// _mm256_testz_si256 returns 1 if (a & b) == 0, so we negate.
__m256i q = _mm256_loadu_si256(reinterpret_cast<const __m256i*>(query));
__m256i n = _mm256_loadu_si256(reinterpret_cast<const __m256i*>(node));
return !_mm256_testz_si256(q, n);
}
else
{
// Large bitmask: process 4 words (256 bits) at a time
size_t i = 0;
for (; i + 4 <= _bitmask_size; i += 4)
{
__m256i q = _mm256_loadu_si256(reinterpret_cast<const __m256i*>(query + i));
__m256i n = _mm256_loadu_si256(reinterpret_cast<const __m256i*>(node + i));
if (!_mm256_testz_si256(q, n))
return true;
}
// Tail: remaining words (0-3)
for (; i < _bitmask_size; i++)
{
if ((query[i] & node[i]) != 0)
return true;
}
return false;
}
#else
// Scalar fallback for non-AVX2 builds
for (size_t i = 0; i < _bitmask_size; i++)
{
if ((bitmask_full_val._mask[i] & _bitsets[i]) != 0)
{
return true;
}
}

return false;
#endif
}

bool simple_bitmask::test_full_mask_contain(const simple_bitmask& bitmask_full_val) const
Expand Down
2 changes: 1 addition & 1 deletion tests/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,7 @@ if (NOT Boost_FOUND)
endif()


set(DISKANN_UNIT_TEST_SOURCES main.cpp index_write_parameters_builder_tests.cpp)
set(DISKANN_UNIT_TEST_SOURCES main.cpp index_write_parameters_builder_tests.cpp filter_match_proxy_tests.cpp)

add_executable(${PROJECT_NAME}_unit_tests ${DISKANN_SOURCES} ${DISKANN_UNIT_TEST_SOURCES})
target_link_libraries(${PROJECT_NAME}_unit_tests ${PROJECT_NAME} ${DISKANN_TOOLS_TCMALLOC_LINK_OPTIONS} Boost::unit_test_framework)
Expand Down
Loading
Loading