Skip to content

Commit 40a493a

Browse files
committed
feat(cpp/search): return std::vector<ResultSet> from rerank and optimize heap updates
- Refactor deglib::search::rerank signature to return std::vector<ResultSet> - Add PQV::replace_top returning const ObjectType& for O(log k) single-pass sift-down root replacement - Avoid heap.top() overhead in rerank loop by tracking max_dist locally - Add docstrings documenting unsorted heap ordering contract of ResultSet - Update Python bindings (float_space_rerank) and FloatSpace.rerank API with return_distances support - Update C++ unit tests in test_search.cpp
1 parent e7f0e71 commit 40a493a

8 files changed

Lines changed: 434 additions & 111 deletions

File tree

‎cpp/deglib/include/deglib/graph/internal_graph.h‎

Lines changed: 39 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -87,20 +87,53 @@ class PQV : public std::vector<ObjectType> {
8787
this->pop_back();
8888
}
8989

90-
/// Replaces the top element in O(log k) time without vector reallocation.
91-
void replace_top(ObjectType&& x) {
92-
std::pop_heap(this->begin(), this->end(), comp);
93-
this->back() = std::move(x);
94-
std::push_heap(this->begin(), this->end(), comp);
90+
/// Replaces the top element in-place, sifts it down in single-pass O(log k) time,
91+
/// and returns a const reference to the new top element.
92+
template <class... Args>
93+
const ObjectType& replace_top(Args&&... args) {
94+
ObjectType val(std::forward<Args>(args)...);
95+
size_t len = this->size();
96+
size_t parent = 0;
97+
size_t child = 1;
98+
99+
while (child < len) {
100+
// Find the larger child (comp(a, b) means a < b for max-heap)
101+
if (child + 1 < len && comp((*this)[child], (*this)[child + 1])) {
102+
child++;
103+
}
104+
// If val is not less than the larger child, heap property holds
105+
if (!comp(val, (*this)[child])) {
106+
break;
107+
}
108+
(*this)[parent] = std::move((*this)[child]);
109+
parent = child;
110+
child = 2 * parent + 1;
111+
}
112+
(*this)[parent] = std::move(val);
113+
return this->front();
95114
}
96115

97116
/// Re-establishes the heap order using the internal comparator.
98117
void heapify() {
99118
std::make_heap(this->begin(), this->end(), comp);
100119
}
120+
121+
/// Sorts the heap elements in ascending order using the internal comparator,
122+
/// consuming the heap property. After calling this, top() is no longer valid
123+
/// until heapify() is called again.
124+
void sort() {
125+
std::sort_heap(this->begin(), this->end(), comp);
126+
}
101127
};
102128

103-
// Search result set (Max-Heap: top() returns the entry with the LARGEST distance)
129+
/**
130+
* Search result set (Max-Heap: top() returns the entry with the LARGEST distance).
131+
*
132+
* Note: Elements are maintained in max-heap order within the underlying vector.
133+
* Iterating or indexing the container directly will yield elements in arbitrary (unsorted) heap order.
134+
* To access elements in ordered sequence (closest distance first), extract them using top() and pop(),
135+
* or call sort() explicitly to sort the underlying storage in ascending distance order.
136+
*/
104137
using ResultSet = PQV<std::less<ObjectDistance>, ObjectDistance>;
105138

106139
// Unchecked search candidates (Min-Heap: top() returns the entry with the SMALLEST distance)

‎cpp/deglib/include/deglib/graph/visited_list_pool.h‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@
55
#include <vector>
66
#include <deque>
77
#include <cstdint>
8+
#include <memory>
89

910
/**
1011
* Ref https://raw.githubusercontent.com/nmslib/hnswlib/master/hnswlib/visited_list_pool.h

‎cpp/deglib/include/deglib/search.h‎

Lines changed: 34 additions & 28 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,13 @@
11
#pragma once
22

33
#include <cstddef>
4+
#include <cstdint>
5+
#include <limits>
46
#include <queue>
57
#include <span>
68
#include <stdexcept>
79
#include <string>
10+
#include <vector>
811
#include "deglib/distances.h"
912
#include "deglib/filter.h"
1013
#include "deglib/concurrent.h"
@@ -18,31 +21,34 @@ using ResultSet = deglib::graph::ResultSet;
1821
/**
1922
* Rerank candidate neighbor indices for queries using exact FloatSpace distances.
2023
*
24+
* Returns a vector of ResultSet objects (one per query). Each ResultSet contains the top-k nearest
25+
* candidates maintained in a max-heap property. Note that ResultSet elements are in heap order (unsorted);
26+
* use top() and pop(), or call sort() on the ResultSet if sorted order is required.
27+
*
2128
* @param space FloatSpace distance calculator instance
2229
* @param queries Pointer to [num_queries x dim] query vectors
2330
* @param num_queries Number of query vectors
2431
* @param base_vectors Pointer to [num_base_vectors x dim] target/base vectors (if null, queries are used as targets)
2532
* @param num_base_vectors Number of base vectors
2633
* @param base_candidates Pointer to [num_queries x candidates_per_query] candidate indices
2734
* @param candidates_per_query Number of candidate indices provided per query
28-
* @param k_top Number of top candidates to output per query (0 = all)
35+
* @param k_top Number of top candidates to keep per query (0 = all)
2936
* @param num_threads Number of worker threads (0 = auto-detect)
30-
* @param out_result_indices Output pointer to [num_queries x k_top] uint32_t indices
37+
* @return std::vector<ResultSet> containing the top-k result sets per query (unsorted heap order)
3138
*/
32-
inline void rerank(
39+
inline std::vector<ResultSet> rerank(
3340
const deglib::distances::FloatSpace& space,
3441
const void* queries,
3542
size_t num_queries,
3643
const void* base_vectors,
3744
size_t num_base_vectors,
3845
const uint32_t* base_candidates,
3946
size_t candidates_per_query,
40-
size_t k_top,
41-
size_t num_threads,
42-
uint32_t* out_result_indices
47+
size_t k_top = 0,
48+
size_t num_threads = 0
4349
) {
44-
if (queries == nullptr || base_candidates == nullptr || out_result_indices == nullptr) {
45-
throw std::invalid_argument("rerank: queries, base_candidates, and out_result_indices must not be null");
50+
if (queries == nullptr || base_candidates == nullptr) {
51+
throw std::invalid_argument("rerank: queries and base_candidates must not be null");
4652
}
4753

4854
if (k_top == 0 || k_top > candidates_per_query) {
@@ -59,15 +65,18 @@ inline void rerank(
5965
const uint8_t* t_ptr = static_cast<const uint8_t*>(target_vectors);
6066

6167
const auto param = space.get_dist_func_param();
62-
// Resolve variant type ONCE outside loops via compile-time static dispatch
68+
69+
std::vector<ResultSet> results(num_queries);
70+
6371
space.compute([&](const auto& dist_func_obj) {
6472
using DistType = std::decay_t<decltype(dist_func_obj)>;
6573
deglib::concurrent::parallel_for(0, num_queries, num_threads, [&](size_t i, size_t) {
6674
const uint8_t* query_ptr = q_ptr + i * byte_stride_query;
6775
const uint32_t* cand_row = base_candidates + i * candidates_per_query;
6876

69-
std::vector<std::pair<float, uint32_t>> heap;
70-
heap.reserve(k_top + 1);
77+
ResultSet heap;
78+
heap.reserve(k_top);
79+
float max_dist = std::numeric_limits<float>::max();
7180

7281
for (size_t j = 0; j < candidates_per_query; ++j) {
7382
uint32_t cand_idx = cand_row[j];
@@ -78,29 +87,26 @@ inline void rerank(
7887
const uint8_t* cand_ptr = t_ptr + cand_idx * byte_stride_target;
7988
float dist = DistType::compare(query_ptr, cand_ptr, param);
8089

90+
// Keep candidates in max-heap of capacity k_top.
91+
// Until the heap reaches k_top elements, add all valid candidates.
8192
if (heap.size() < k_top) {
82-
heap.push_back({dist, cand_idx});
83-
std::push_heap(heap.begin(), heap.end());
84-
} else if (dist < heap.front().first) {
85-
std::pop_heap(heap.begin(), heap.end());
86-
heap.back() = {dist, cand_idx};
87-
std::push_heap(heap.begin(), heap.end());
93+
heap.emplace(cand_idx, dist);
94+
if (heap.size() == k_top) {
95+
max_dist = heap.top().getDistance();
96+
}
97+
}
98+
// Once full, only consider candidates strictly closer than the worst (max_dist).
99+
// replace_top replaces the root in O(log k) and returns a reference to the NEW root element (new max distance).
100+
else if (dist < max_dist) {
101+
max_dist = heap.replace_top(cand_idx, dist).getDistance();
88102
}
89103
}
90104

91-
std::sort_heap(heap.begin(), heap.end());
92-
93-
uint32_t* out_row = out_result_indices + i * k_top;
94-
size_t actual_k = heap.size();
95-
for (size_t k = 0; k < actual_k; ++k) {
96-
out_row[k] = heap[k].second;
97-
}
98-
for (size_t k = actual_k; k < k_top; ++k) {
99-
out_row[k] = static_cast<uint32_t>(i);
100-
}
105+
results[i] = std::move(heap);
101106
});
102107
});
108+
109+
return results;
103110
}
104111

105112
} // namespace deglib::search
106-

‎cpp/test/src/unit/graph/test_internal_graph.cpp‎

Lines changed: 65 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -133,6 +133,71 @@ TEST(UncheckedSet, PopRemovesMin) {
133133
EXPECT_NEAR(us.top().getDistance(), 3.0f, 1e-6f);
134134
}
135135

136+
// ---------------------------------------------------------------------------
137+
// PQV emplace / pop / heapify (heap invariant tests)
138+
// ---------------------------------------------------------------------------
139+
140+
TEST(PQV, EmplaceRestoresHeap) {
141+
deglib::graph::ResultSet rs;
142+
rs.emplace(1, 1.0f);
143+
rs.emplace(2, 5.0f);
144+
rs.emplace(3, 3.0f);
145+
146+
// ResultSet is a max-heap: top() must be the largest distance
147+
EXPECT_EQ(rs.size(), 3u);
148+
EXPECT_NEAR(rs.top().getDistance(), 5.0f, 1e-6f);
149+
EXPECT_EQ(rs.top().getIdentifier(), 2u);
150+
151+
// Verify heap property holds for all elements
152+
for (size_t i = 0; i < rs.size(); ++i) {
153+
EXPECT_GE(rs[i].getDistance(), 0.0f);
154+
}
155+
}
156+
157+
TEST(PQV, Heapify) {
158+
deglib::graph::ResultSet rs;
159+
// Insert elements without maintaining heap order (bypass emplace)
160+
rs.reserve(5);
161+
static_cast<std::vector<deglib::graph::ObjectDistance>&>(rs).emplace_back(1, 1.0f);
162+
static_cast<std::vector<deglib::graph::ObjectDistance>&>(rs).emplace_back(2, 5.0f);
163+
static_cast<std::vector<deglib::graph::ObjectDistance>&>(rs).emplace_back(3, 3.0f);
164+
165+
// Before heapify, top() is front() which may not be the max
166+
rs.heapify();
167+
EXPECT_NEAR(rs.top().getDistance(), 5.0f, 1e-6f);
168+
EXPECT_EQ(rs.top().getIdentifier(), 2u);
169+
}
170+
171+
TEST(PQV, PopMaintainsHeap) {
172+
deglib::graph::ResultSet rs;
173+
rs.emplace(1, 1.0f);
174+
rs.emplace(2, 5.0f);
175+
rs.emplace(3, 3.0f);
176+
177+
rs.pop(); // remove max (5.0)
178+
EXPECT_EQ(rs.size(), 2u);
179+
EXPECT_NEAR(rs.top().getDistance(), 3.0f, 1e-6f);
180+
181+
rs.pop(); // remove max (3.0)
182+
EXPECT_EQ(rs.size(), 1u);
183+
EXPECT_NEAR(rs.top().getDistance(), 1.0f, 1e-6f);
184+
}
185+
186+
// ---------------------------------------------------------------------------
187+
// ObjectDistance POD traits
188+
// ---------------------------------------------------------------------------
189+
190+
TEST(ObjectDistance, POD_Traits) {
191+
// ObjectDistance should be trivially default-constructible, copyable, and movable
192+
static_assert(std::is_trivially_default_constructible_v<deglib::graph::ObjectDistance>);
193+
static_assert(std::is_trivially_copyable_v<deglib::graph::ObjectDistance>);
194+
static_assert(std::is_trivially_destructible_v<deglib::graph::ObjectDistance>);
195+
static_assert(std::is_standard_layout_v<deglib::graph::ObjectDistance>);
196+
197+
// Verify size: uint32_t + float = 8 bytes
198+
EXPECT_EQ(sizeof(deglib::graph::ObjectDistance), 8u);
199+
}
200+
136201
// ---------------------------------------------------------------------------
137202
// Mock implementation of InternalGraph for testing the interface
138203
// ---------------------------------------------------------------------------

0 commit comments

Comments
 (0)