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
88 changes: 59 additions & 29 deletions include/rabitqlib/index/hnsw/hnsw.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -458,6 +458,9 @@ inline HierarchicalNSW::HierarchicalNSW(
, element_levels_(max_elements)
, raw_dist_func_((metric_type == METRIC_IP) ? dot_product_dis<float> : euclidean_sqr<float>) {
validate_metric_type(metric_type);
if (M < 2) {
throw std::invalid_argument("HNSW M must be at least 2");
}
max_elements_ = max_elements;
dim_ = dim;
rotator_.reset(choose_rotator<float>(
Expand Down Expand Up @@ -652,7 +655,7 @@ inline void HierarchicalNSW::load(const char* filename) {
loaded.padded_dim_ != round_up_to_multiple(loaded.dim_, 64) ||
loaded.ex_bits_ > 8 ||
(loaded.metric_type_ != METRIC_L2 && loaded.metric_type_ != METRIC_IP) ||
loaded.M_ == 0 || loaded.M_ > 10000 || loaded.maxM_ != loaded.M_ ||
loaded.M_ < 2 || loaded.M_ > 10000 || loaded.maxM_ != loaded.M_ ||
loaded.maxM0_ != 2 * loaded.M_ || loaded.ef_construction_ < loaded.M_ ||
(element_count == 0 &&
(loaded.maxlevel_ != -1 || loaded.enterpoint_node_ != kPidMax)) ||
Expand All @@ -676,8 +679,7 @@ inline void HierarchicalNSW::load(const char* filename) {
loaded.offsetExData_ != expected_ex_offset ||
loaded.size_data_per_element_ != expected_stride ||
loaded.size_links_per_element_ != (loaded.maxM_ + 1) * sizeof(PID) ||
(loaded.M_ > 1 && (!std::isfinite(loaded.mult_) || loaded.mult_ <= 0.0)) ||
(loaded.M_ == 1 && !std::isinf(loaded.mult_))) {
!std::isfinite(loaded.mult_) || loaded.mult_ <= 0.0) {
invalid_file();
}

Expand Down Expand Up @@ -898,6 +900,9 @@ inline void HierarchicalNSW::construct(
size_t num_threads = 0,
bool faster = false
) {
if (num_cluster_ != 0 || cur_element_count_ != 0) {
throw std::logic_error("HNSW index is already constructed or loaded");
}
if (cluster_num == 0 || cluster_num > buffer::kSearchBufferMaxPointCount) {
throw std::invalid_argument("HNSW cluster count is out of range");
}
Expand All @@ -917,32 +922,55 @@ inline void HierarchicalNSW::construct(
}
}

num_cluster_ = cluster_num;
const size_t centroids_bytes = num_cluster_ * padded_dim_ * sizeof(float);
centroids_memory_ = memory::huge_page_allocate<char>(centroids_bytes);
if (centroids_memory_ == nullptr) {
throw std::runtime_error("Not enough memory: HNSW failed to allocate centroids");
}
try {
num_cluster_ = cluster_num;
const size_t centroids_bytes = num_cluster_ * padded_dim_ * sizeof(float);
centroids_memory_ = memory::huge_page_allocate<char>(centroids_bytes);
if (centroids_memory_ == nullptr) {
throw std::runtime_error("Not enough memory: HNSW failed to allocate centroids"
);
}

for (size_t i = 0; i < cluster_num; ++i) {
this->rotator_->rotate(
centroids + (i * dim_),
reinterpret_cast<float*>(centroids_memory_) + (i * padded_dim_)
);
}
for (size_t i = 0; i < cluster_num; ++i) {
this->rotator_->rotate(
centroids + (i * dim_),
reinterpret_cast<float*>(centroids_memory_) + (i * padded_dim_)
);
}

quant::RabitqConfig config;
if (faster) {
config = quant::faster_config(padded_dim_, ex_bits_ + 1);
}
quant::RabitqConfig config;
if (faster) {
config = quant::faster_config(padded_dim_, ex_bits_ + 1);
}

rawDataPtr_ = data;
rabitqlib::ivf::parallel_for(
0,
data_num,
num_threads,
[&](size_t idx, size_t /*threadId*/) { add_point(idx, cluster_ids[idx], config); }
);
rawDataPtr_ = data;
rabitqlib::ivf::parallel_for(
0,
data_num,
num_threads,
[&](size_t idx, size_t /*threadId*/) {
add_point(idx, cluster_ids[idx], config);
}
);
} catch (...) {
// parallel_for joins every worker before propagating an exception.
// Discard the partial graph, retaining capacity and rotation for a retry.
for (size_t i = 0; i < cur_element_count_; ++i) {
std::free(linkLists_[i]);
linkLists_[i] = nullptr;
element_levels_[i] = 0;
}
cur_element_count_ = 0;
label_lookup_.clear();
enterpoint_node_ = kPidMax;
maxlevel_ = -1;
memory::aligned_deallocate(centroids_memory_);
centroids_memory_ = nullptr;
num_cluster_ = 0;
rawDataPtr_ = nullptr;
throw;
}
rawDataPtr_ = nullptr;
}

inline void HierarchicalNSW::add_point(
Expand All @@ -967,6 +995,7 @@ inline void HierarchicalNSW::add_point(
}

cur_c = cur_element_count_;
linkLists_[cur_c] = nullptr;
cur_element_count_++;
label_lookup_[label] = cur_c;
curlevel = get_random_level(mult_);
Expand All @@ -980,10 +1009,10 @@ inline void HierarchicalNSW::add_point(
element_levels_[cur_c] = curlevel;
std::unique_lock<std::mutex> templock(global_);
int maxlevelcopy = maxlevel_;
PID curr_obj = enterpoint_node_;
if (curlevel <= maxlevelcopy) {
templock.unlock();
}
PID curr_obj = enterpoint_node_;

// initialize the current memory.
memset(
Expand Down Expand Up @@ -1070,7 +1099,7 @@ inline void HierarchicalNSW::add_point(
inline maxheap<std::pair<float, PID>> HierarchicalNSW::search_base_layer(
PID ep_id, PID cur_c, int layer
) {
VisitedSet* vl = visited_list_pool_->get_free_vislist();
std::unique_ptr<VisitedSet> vl(visited_list_pool_->get_free_vislist());

maxheap<std::pair<float, PID>> top_candidates;
minheap<std::pair<float, PID>> candidate_set;
Expand Down Expand Up @@ -1143,7 +1172,8 @@ inline maxheap<std::pair<float, PID>> HierarchicalNSW::search_base_layer(
}
}
}
visited_list_pool_->release_vis_list(vl);
visited_list_pool_->release_vis_list(vl.get());
vl.release();
return top_candidates;
}

Expand Down
4 changes: 2 additions & 2 deletions include/rabitqlib/index/hnsw/hnsw_quant.hpp
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
// Insertion into a built index: the construct path with get_quant_dist in place
// of get_data_dist, because rawDataPtr_ dangles once construct returns.
// of get_data_dist, because raw input is borrowed only during construct.
//
// add routes, draws every level, allocates every link list and quantizes every
// point before reserving one block of slots and linking in parallel. Levels come
Expand Down Expand Up @@ -356,10 +356,10 @@ inline void HierarchicalNSW::add_point_quant(

std::unique_lock<std::mutex> templock(global_);
int maxlevelcopy = maxlevel_;
PID curr_obj = enterpoint_node_;
if (curlevel <= maxlevelcopy) {
templock.unlock();
}
PID curr_obj = enterpoint_node_;
const PID previous_entry_point = curr_obj;

linkLists_[cur_c] = link_list.release();
Expand Down
3 changes: 1 addition & 2 deletions python_bindings/hnsw_bindings.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -84,8 +84,6 @@ class HnswIndex {
throw std::invalid_argument("cluster_ids contains an out-of-range value");
}
}
num_clusters_ = num_clusters;

// Ensure cluster_ids are writable for the C++ API by making a copy
std::vector<rabitqlib::PID> cluster_ids_vec(
static_cast<size_t>(cluster_ids_array.shape(0))
Expand All @@ -105,6 +103,7 @@ class HnswIndex {
num_threads,
fast_quantization
);
num_clusters_ = num_clusters;
built_ = true;
}

Expand Down
57 changes: 57 additions & 0 deletions tests/failures/hnsw_add_allocation_test.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -86,6 +86,63 @@ struct Fixture {
}
};

TEST(HnswConstructionAllocationTest, EveryAllocationFailureLeavesEmptyReusableIndex) {
Fixture fixture;
for (const size_t threads : {1U, 4U}) {
SCOPED_TRACE(threads);
size_t failures = 0;
bool succeeded = false;
for (size_t allocation = 0; allocation < 512; ++allocation) {
SCOPED_TRACE(allocation);
HierarchicalNSW index(kCapacity, kDim, 4, 2, 10, 1);
auto clusters = fixture.clusters;
bool failed = false;
{
AllocationFailure failure(allocation);
try {
index.construct(
1,
fixture.centroid.data(),
kAdded,
fixture.data.data(),
clusters.data(),
threads,
false
);
} catch (const std::bad_alloc&) { failed = true; }
}
if (!failed) {
succeeded = true;
break;
}
++failures;
EXPECT_EQ(index.num_points(), 0U);
EXPECT_EQ(index.num_clusters(), 0U);
const auto empty = index.search(fixture.data.data(), 1, 1, kCapacity, 1);
ASSERT_EQ(empty.size(), 1U);
EXPECT_TRUE(empty[0].empty());
ASSERT_NO_THROW(index.construct(
1,
fixture.centroid.data(),
kAdded,
fixture.data.data(),
clusters.data(),
threads,
false
));
EXPECT_EQ(index.num_points(), kAdded);
EXPECT_EQ(index.num_clusters(), 1U);
const auto results = index.search(fixture.data.data(), kAdded, 1, kCapacity, 1);
for (size_t i = 0; i < kAdded; ++i) {
ASSERT_EQ(results[i].size(), 1U);
EXPECT_EQ(results[i][0].second, i);
}
}
EXPECT_GT(failures, 0U);
EXPECT_TRUE(succeeded);
}
}

TEST(HnswAddAllocationTest, AllocatesReturnedIdsBeforeMutation) {
Fixture fixture;
HierarchicalNSW index(kCapacity, kDim, 4, 2, 10, 1);
Expand Down
27 changes: 27 additions & 0 deletions tests/python/test_hnsw.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,33 @@ def test_is_built(built_hnsw):
assert built_hnsw.is_built


@pytest.mark.parametrize("m", [0, 1])
def test_rejects_too_small_m(m):
with pytest.raises(ValueError, match="^HNSW M must be at least 2$"):
HnswIndex(DIM, N_VECTORS, M=m)


@pytest.mark.parametrize("reload", [False, True])
def test_rejected_rebuild_preserves_index(base_data, clusters, tmp_path, reload):
idx = HnswIndex(DIM, N_VECTORS, M=8, ef_construction=50, nbits=4)
centroids, cluster_ids = clusters
idx.build(base_data, centroids, cluster_ids)
if reload:
path = str(tmp_path / "hnsw.index")
idx.save(path)
idx = HnswIndex.load(path)
before = idx.search(base_data[:10], k=3, ef=_EF)
with pytest.raises(
RuntimeError, match="^HNSW index is already constructed or loaded$"
):
idx.build(base_data, centroids[:1], np.zeros(N_VECTORS, dtype=np.uint32))
assert idx.is_built
assert idx.num_clusters == N_CLUSTERS
after = idx.search(base_data[:10], k=3, ef=_EF)
for actual, expected in zip(after, before, strict=True):
np.testing.assert_array_equal(actual, expected)


def test_properties(built_hnsw):
assert built_hnsw.dim == DIM
assert built_hnsw.nbits == 4
Expand Down
47 changes: 47 additions & 0 deletions tests/unit/rabitqlib/index/hnsw_test.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -645,6 +645,19 @@ TEST(HnswAddTest, AddIntoAnEmptyIndexBuildsFromScratch) {
EXPECT_EQ(exact, kCount);
}

TEST(HnswConstructionTest, RejectsTooSmallM) {
for (const size_t M : {0U, 1U}) {
SCOPED_TRACE(M);
try {
HierarchicalNSW index(8, 64, 4, M, 10);
FAIL() << "M below 2 must be rejected";
} catch (const std::invalid_argument& error) {
EXPECT_STREQ(error.what(), "HNSW M must be at least 2");
}
}
EXPECT_NO_THROW(HierarchicalNSW(8, 64, 4, 2, 10));
}

TEST(HnswConstructionTest, RejectsInvalidInputsBeforeChangingIndex) {
constexpr size_t kDim = 64;
std::vector<float> data(kDim, 1.0F);
Expand Down Expand Up @@ -705,6 +718,40 @@ TEST_F(HnswSaveTest, RejectsUnopenableDestination) {
}
}

TEST_F(HnswSaveTest, RejectedRebuildPreservesConstructedAndLoadedIndexes) {
// Multiple original centroids expose a rejected rebuild replacing the
// centroid buffer with a smaller one while retaining the old cluster IDs.
std::vector<float> centroids(2 * kDim, 0.0F);
centroids[kDim] = 1.0F;
for (size_t i = 0; i < kCount; ++i) {
cluster_ids_[i] = i % 2;
}
HierarchicalNSW original(kCount, kDim, 4, 4, 10);
original.construct(2, centroids.data(), kCount, data_.data(), cluster_ids_.data(), 1);
original.save(path_.c_str());
HierarchicalNSW loaded;
loaded.load(path_.c_str());
std::vector<PID> new_clusters(kCount, 0);
for (auto* index : {&original, &loaded}) {
const auto before = index->search(data_.data(), kCount, 2, kCount, 1);
try {
index->construct(
1, centroid_.data(), kCount, data_.data(), new_clusters.data(), 1
);
FAIL() << "A second construction must be rejected";
} catch (const std::logic_error& error) {
EXPECT_STREQ(error.what(), "HNSW index is already constructed or loaded");
}
EXPECT_EQ(index->num_points(), kCount);
EXPECT_EQ(index->num_clusters(), 2U);
EXPECT_EQ(index->search(data_.data(), kCount, 2, kCount, 1), before);
index->save(path_.c_str());
HierarchicalNSW roundtrip;
roundtrip.load(path_.c_str());
EXPECT_EQ(roundtrip.search(data_.data(), kCount, 2, kCount, 1), before);
}
}

TEST_F(HnswSaveTest, ReportsWriteOrCloseFailure) {
if (!std::filesystem::exists("/dev/full")) {
GTEST_SKIP() << "/dev/full is unavailable";
Expand Down
Loading