2828#include < span>
2929#include < stdexcept>
3030#include < type_traits>
31+ #include < utility>
3132#include < vector>
3233
3334namespace 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`.
9589template <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
0 commit comments