Skip to content

Commit 13f10a4

Browse files
committed
feat(kmeans): add kmeans clustering to find good entry point in the graph for the searcher
1 parent be37147 commit 13f10a4

6 files changed

Lines changed: 154 additions & 48 deletions

File tree

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

Lines changed: 107 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -28,6 +28,7 @@
2828
#include <span>
2929
#include <stdexcept>
3030
#include <type_traits>
31+
#include <utility>
3132
#include <vector>
3233

3334
namespace deglib::search {
@@ -40,22 +41,15 @@ concept DistanceCallable = requires(F f, const float* a, const float* b) {
4041
};
4142

4243
// Decodes one native graph feature vector into `dim` floats.
43-
inline void decode_feature_vector(
44-
const deglib::distances::FloatSpace& space, const std::byte* native, float* out
45-
) {
44+
inline void decode_feature_vector(const deglib::distances::FloatSpace& space, const std::byte* native, float* out) {
4645
const size_t dim = space.dim();
4746
using deglib::distances::MetricDataType;
4847
switch (space.metric().get_data_type()) {
4948
case MetricDataType::FP32:
50-
std::copy(
51-
reinterpret_cast<const float*>(native),
52-
reinterpret_cast<const float*>(native) + dim, out
53-
);
49+
std::copy(reinterpret_cast<const float*>(native), reinterpret_cast<const float*>(native) + dim, out);
5450
break;
5551
case MetricDataType::FP16:
56-
deglib::distances::fp16::fp16_to_floats(
57-
reinterpret_cast<const uint16_t*>(native), out, dim
58-
);
52+
deglib::distances::fp16::fp16_to_floats(reinterpret_cast<const uint16_t*>(native), out, dim);
5953
break;
6054
case MetricDataType::Uint8: {
6155
const auto* v = reinterpret_cast<const uint8_t*>(native);
@@ -89,19 +83,12 @@ inline void decode_feature_vector(
8983
}
9084
}
9185

92-
// Lloyd k-means over caller-provided rows. Centroids are normalized means
93-
// (cosine-style, matching the previous Python behavior); empty or degenerate
94-
// clusters keep their previous centroid. Returns medoid positions into `rows`.
86+
// Lloyd k-means over caller-provided rows. Centroids are normalized means.
87+
// empty or degenerate clusters keep their previous centroid.
88+
// Returns medoid positions (indices) into `rows`.
9589
template <DistanceCallable DistFn>
96-
[[nodiscard]] std::vector<uint32_t> kmeans_medoids(
97-
std::span<const float* const> rows,
98-
uint32_t dim,
99-
DistFn&& distance,
100-
uint32_t n_clusters,
101-
uint32_t n_iter,
102-
uint32_t seed,
103-
size_t threads
104-
) {
90+
[[nodiscard]] std::vector<uint32_t>
91+
kmeans_medoids(std::span<const float* const> rows, uint32_t dim, DistFn&& distance, uint32_t n_clusters, uint32_t n_iter, uint32_t seed, size_t threads) {
10592
if (rows.empty() || dim == 0) {
10693
throw std::invalid_argument("kmeans_medoids: empty input");
10794
}
@@ -128,7 +115,7 @@ template <DistanceCallable DistFn>
128115
std::vector<float> sums(static_cast<size_t>(n_clusters) * dim);
129116
std::vector<uint32_t> counts(n_clusters);
130117
for (uint32_t iter = 0; iter < n_iter; ++iter) {
131-
// Assignment: embarrassingly parallel over rows.
118+
// Assignment: parallel over rows.
132119
deglib::concurrent::parallel_for(0, rows.size(), threads, [&](size_t i, size_t) {
133120
const float* vec = rows[i];
134121
float best_dist = std::numeric_limits<float>::max();
@@ -203,8 +190,8 @@ template <DistanceCallable DistFn>
203190
uint32_t n_clusters,
204191
uint32_t n_iter,
205192
size_t sample_size,
206-
uint32_t seed,
207-
size_t threads
193+
uint32_t seed = 7,
194+
size_t threads = 1
208195
) {
209196
const size_t n_vectors = graph.size();
210197
if (n_vectors == 0) {
@@ -238,20 +225,14 @@ template <DistanceCallable DistFn>
238225
}
239226

240227
const bool is_l2 = space.metric().get_distance_kind() == deglib::distances::MetricDistanceKind::L2;
241-
const deglib::distances::FloatSpace float_space(
242-
dim, is_l2 ? deglib::distances::Metric::FP32_L2 : deglib::distances::Metric::FP32_InnerProduct
243-
);
228+
const deglib::distances::FloatSpace float_space(dim, is_l2 ? deglib::distances::Metric::FP32_L2 : deglib::distances::Metric::FP32_InnerProduct);
244229
const std::span<const float* const> row_view(rows.data(), rows.size());
245230
const uint32_t dim32 = static_cast<uint32_t>(dim);
246231
const std::vector<uint32_t> positions = float_space.compute([&](const auto& kernel) {
247232
// Static dispatch: the concrete SIMD kernel inlines into the loop.
248233
return kmeans_medoids(
249234
row_view, dim32,
250-
[&](const float* a, const float* b) -> float {
251-
return std::decay_t<decltype(kernel)>::compare(
252-
a, b, float_space.get_dist_func_param()
253-
);
254-
},
235+
[&](const float* a, const float* b) -> float { return std::decay_t<decltype(kernel)>::compare(a, b, float_space.get_dist_func_param()); },
255236
n_clusters, n_iter, seed, threads
256237
);
257238
});
@@ -264,4 +245,97 @@ template <DistanceCallable DistFn>
264245
return medoids;
265246
}
266247

248+
/**
249+
* Owns the k-means entry vertices for one graph.
250+
*
251+
* @param graph Graph the entry vertices refer to, must outlive this selector.
252+
*/
253+
class KMeansEntrySelector {
254+
public:
255+
explicit KMeansEntrySelector(const deglib::graph::InternalGraph& graph) : graph_(&graph) {}
256+
257+
/**
258+
* Select entry vertices via k-means medoids.
259+
*
260+
* @param n_clusters Number of entry vertices to select.
261+
* @param n_iter Number of k-means iterations.
262+
* @param sample_size Number of vertices sampled for clustering, 0 selects 3% of the graph size.
263+
* @param seed Random seed for sampling and centroid init.
264+
* @param threads Number of worker threads.
265+
*/
266+
void optimize(uint32_t n_clusters = 128, uint32_t n_iter = 15, size_t sample_size = 0, uint32_t seed = 7, size_t threads = 1) {
267+
if (sample_size == 0) {
268+
const size_t n = graph_->size();
269+
sample_size = n * 3 / 100;
270+
if (sample_size == 0 && n > 0) sample_size = n;
271+
}
272+
entries_ = graph_kmeans_medoids(*graph_, n_clusters, n_iter, sample_size, seed, threads);
273+
}
274+
275+
/**
276+
* Nearest entry vertices for a query, nearest first.
277+
*
278+
* Falls back to vertex 0 when no usable entry exists.
279+
*
280+
* @param query Native query bytes sized to the graph feature size.
281+
* @param count Maximum number of entries to return.
282+
*/
283+
[[nodiscard]] std::vector<uint32_t> top_entries(std::span<const std::byte> query, size_t count = 2) const {
284+
const uint32_t n = graph_->size();
285+
if (n == 0 || count == 0) return {};
286+
const auto dist_func = graph_->getFeatureSpace().get_dist_func();
287+
const auto dist_param = graph_->getFeatureSpace().get_dist_func_param();
288+
std::vector<uint32_t> valid;
289+
valid.reserve(entries_.size());
290+
for (auto ep : entries_) {
291+
if (ep < n) valid.push_back(ep);
292+
}
293+
if (valid.empty()) return {0};
294+
if (valid.size() <= count) return valid;
295+
if (count == 2) {
296+
uint32_t ep1 = valid[0], ep2 = valid[0];
297+
float dist1 = std::numeric_limits<float>::max(), dist2 = std::numeric_limits<float>::max();
298+
for (auto ep : valid) {
299+
float d = dist_func(query.data(), graph_->getFeatureVector(ep), dist_param);
300+
if (d < dist1) {
301+
dist2 = dist1;
302+
ep2 = ep1;
303+
dist1 = d;
304+
ep1 = ep;
305+
} else if (d < dist2) {
306+
dist2 = d;
307+
ep2 = ep;
308+
}
309+
}
310+
if (ep1 == ep2) return {ep1};
311+
return {ep1, ep2};
312+
}
313+
std::vector<std::pair<float, uint32_t>> scored;
314+
scored.reserve(valid.size());
315+
for (auto ep : valid) {
316+
scored.emplace_back(dist_func(query.data(), graph_->getFeatureVector(ep), dist_param), ep);
317+
}
318+
std::nth_element(scored.begin(), scored.begin() + count, scored.end(), [](const auto& a, const auto& b) { return a.first < b.first; });
319+
scored.resize(count);
320+
std::sort(scored.begin(), scored.end(), [](const auto& a, const auto& b) { return a.first < b.first; });
321+
std::vector<uint32_t> top;
322+
top.reserve(count);
323+
for (const auto& s : scored) top.push_back(s.second);
324+
return top;
325+
}
326+
327+
/** Stored entry vertices as graph internal indices. */
328+
[[nodiscard]] const std::vector<uint32_t>& entries() const noexcept { return entries_; }
329+
330+
/** True when optimize() has not produced entries yet. */
331+
[[nodiscard]] bool empty() const noexcept { return entries_.empty(); }
332+
333+
/** Number of stored entry vertices. */
334+
[[nodiscard]] size_t size() const noexcept { return entries_.size(); }
335+
336+
private:
337+
const deglib::graph::InternalGraph* graph_ = nullptr;
338+
std::vector<uint32_t> entries_;
339+
};
340+
267341
} // namespace deglib::search

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

Lines changed: 32 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -195,7 +195,16 @@ class SearcherBase {
195195
bool unsorted = false
196196
) const = 0;
197197

198-
virtual void optimize(uint32_t n_clusters = 128, uint32_t n_iter = 15, size_t sample_size = 30000, uint32_t seed = 42, size_t threads = 0) = 0;
198+
/**
199+
* Select entry vertices via k-means medoids.
200+
*
201+
* @param n_clusters Number of entry vertices to select.
202+
* @param n_iter Number of k-means iterations.
203+
* @param sample_size Number of vertices sampled for clustering, 0 selects 3% of the graph size.
204+
* @param seed Random seed for sampling and centroid init.
205+
* @param threads Number of worker threads, 0 selects a library default.
206+
*/
207+
virtual void optimize(uint32_t n_clusters = 128, uint32_t n_iter = 15, size_t sample_size = 0, uint32_t seed = 42, size_t threads = 0) = 0;
199208

200209
// --- Modern C++20 std::span and std::vector Convenience API ---
201210

@@ -357,14 +366,23 @@ class SearcherImpl : public SearcherBase {
357366
const deglib::graph::InternalGraph* graph_ = nullptr;
358367
QuantT quantizer_;
359368
RefinerT refiner_;
360-
std::vector<uint32_t> entry_vertex_indices_;
369+
KMeansEntrySelector entry_selector_;
361370

362371
public:
363372
SearcherImpl(const deglib::graph::InternalGraph& graph, QuantT quantizer, RefinerT refiner)
364-
: graph_(&graph), quantizer_(std::move(quantizer)), refiner_(std::move(refiner)) {}
373+
: graph_(&graph), quantizer_(std::move(quantizer)), refiner_(std::move(refiner)), entry_selector_(graph) {}
365374

366-
void optimize(uint32_t n_clusters = 128, uint32_t n_iter = 15, size_t sample_size = 30000, uint32_t seed = 42, size_t threads = 0) override {
367-
entry_vertex_indices_ = graph_kmeans_medoids(*graph_, n_clusters, n_iter, sample_size, seed, threads);
375+
/**
376+
* Select entry vertices via k-means medoids.
377+
*
378+
* @param n_clusters Number of entry vertices to select.
379+
* @param n_iter Number of k-means iterations.
380+
* @param sample_size Number of vertices sampled for clustering, 0 selects 3% of the graph size.
381+
* @param seed Random seed for sampling and centroid init.
382+
* @param threads Number of worker threads.
383+
*/
384+
void optimize(uint32_t n_clusters = 128, uint32_t n_iter = 15, size_t sample_size = 0, uint32_t seed = 7, size_t threads = 1) override {
385+
entry_selector_.optimize(n_clusters, n_iter, sample_size, seed, threads);
368386
}
369387

370388
template <typename QueryT>
@@ -381,6 +399,8 @@ class SearcherImpl : public SearcherBase {
381399
const uint32_t fetch_k = std::max(k, static_cast<uint32_t>(std::round(k * rerank_factor)));
382400
const size_t graph_feature_bytes = graph_->getFeatureSpace().get_data_size();
383401

402+
if (graph_->size() == 0) throw std::invalid_argument("Searcher: graph is empty");
403+
384404
// 1. Static Query Transformation (Zero-copy bypass for NoQuantizer)
385405
const std::byte* query_bytes = nullptr;
386406
alignas(64) std::byte stack_query_bytes[512];
@@ -404,15 +424,18 @@ class SearcherImpl : public SearcherBase {
404424
query_bytes = q_buf;
405425
}
406426

407-
// 2. Direct Graph Search (entries come solely from optimize())
427+
// 2. Entry selection via selector (falls back to vertex 0) + Direct Graph Search
408428
const std::span<const std::byte> query_span(query_bytes, graph_feature_bytes);
409-
deglib::graph::ResultSet result = entry_vertex_indices_.empty() ? graph_->search(query_span, fetch_k, eps, nullptr, 0)
410-
: graph_->search(query_span, entry_vertex_indices_, fetch_k, eps, nullptr, 0);
411-
const size_t found_count = result.size();
429+
const std::vector<uint32_t> entries = entry_selector_.top_entries(query_span, 2);
430+
deglib::graph::ResultSet result =
431+
entries.empty() ? graph_->search(query_span, fetch_k, eps, nullptr, 0) : graph_->search(query_span, entries, fetch_k, eps, nullptr, 0);
412432

413433
// 3. Static Compile-Time Reranker Check
414434
if constexpr (RefinerT::enabled) {
415435
if (fetch_k > k) {
436+
const size_t found_count = result.size();
437+
438+
// Stack fast path for up to 256 candidates, heap fallback beyond that without result limit.
416439
uint32_t stack_cand_indices[256];
417440
std::unique_ptr<uint32_t[]> heap_cand_indices;
418441
uint32_t* cands = stack_cand_indices;

‎examples/vibe/config.yml‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -32,4 +32,4 @@ float:
3232
prune_non_rng: [false]
3333
query_args:
3434
rerank_size_factor: [1.0, 1.15, 1.35]
35-
ef: [60, 80, 100, 150, 200, 250, 300, 350, 400, 500, 600, 700, 800, 845, 900, 1000]
35+
search_eps: [0.0, 0.005, 0.01, 0.02, 0.04, 0.06, 0.08, 0.1, 0.12, 0.15, 0.2, 0.25, 0.3]

‎examples/vibe/module.py‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -290,7 +290,7 @@ def fit(self, X: np.ndarray, cache_dir: Path | None = None):
290290

291291
# 6. Optimize entry vertices via K-Means cluster medoids computed in C++
292292
t_km = time.time()
293-
self.searcher.optimize()
293+
#self.searcher.optimize()
294294
print(f"K-Means 128 cluster medoids computed in {time.time() - t_km:.2f}s", flush=True)
295295

296296
def set_query_arguments(self, *args, **kwargs):
@@ -331,7 +331,7 @@ def query(self, v: np.ndarray, n: int) -> np.ndarray:
331331
threads=1,
332332
return_distances=False,
333333
unsorted=True,
334-
ef=self.ef,
334+
#ef=self.ef,
335335
)
336336

337337
def __str__(self) -> str:

‎python/src/deg_cpp/deglib_cpp.cpp‎

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -814,7 +814,7 @@ class SearcherPy {
814814
}
815815
}
816816

817-
void optimize(uint32_t n_clusters = 128, uint32_t n_iter = 15, size_t sample_size = 30000, uint32_t seed = 42, size_t num_threads = 0) {
817+
void optimize(uint32_t n_clusters = 128, uint32_t n_iter = 15, size_t sample_size = 0, uint32_t seed = 42, size_t num_threads = 0) {
818818
py::gil_scoped_release release;
819819
searcher_->optimize(n_clusters, n_iter, sample_size, seed, num_threads);
820820
}
@@ -1581,7 +1581,8 @@ PYBIND11_MODULE(deglib_cpp, m) {
15811581
.def(
15821582
"optimize", &SearcherPy::optimize,
15831583
py::arg("n_clusters") = 128, py::arg("n_iter") = 15,
1584-
py::arg("sample_size") = 30000, py::arg("seed") = 42, py::arg("num_threads") = 0
1584+
py::arg("sample_size") = 0, py::arg("seed") = 42, py::arg("num_threads") = 0,
1585+
"Select entry vertices via k-means medoids.\n\nArgs:\n n_clusters: number of entry vertices to select.\n n_iter: number of k-means iterations.\n sample_size: vertices sampled for clustering, 0 selects 3% of the graph size.\n seed: random seed for sampling and centroid init.\n num_threads: worker threads, 0 selects a library default."
15851586
);
15861587

15871588
// graphs

‎python/src/deglib/search.py‎

Lines changed: 9 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -127,10 +127,18 @@ def optimize(
127127
self,
128128
n_clusters: int = 128,
129129
n_iter: int = 15,
130-
sample_size: int = 30000,
130+
sample_size: int = 0,
131131
seed: int = 42,
132132
num_threads: int = 0,
133133
) -> None:
134+
"""Select entry vertices via k-means medoids.
135+
136+
:param n_clusters: Number of entry vertices to select.
137+
:param n_iter: Number of k-means iterations.
138+
:param sample_size: Number of vertices sampled for clustering, 0 selects 3% of the graph size.
139+
:param seed: Random seed for sampling and centroid init.
140+
:param num_threads: Number of worker threads, 0 selects a library default.
141+
"""
134142
self.searcher_cpp.optimize(int(n_clusters), int(n_iter), int(sample_size), int(seed), int(num_threads))
135143

136144
def search(

0 commit comments

Comments
 (0)