MCPcopy Create free account
hub / github.com/VectorDB-NTU/RaBitQ-Library / build

Method build

python_bindings/hnsw_bindings.cpp:44–82  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

42 )) {}
43
44 void build(
45 py::handle data,
46 py::handle centroids,
47 py::handle cluster_ids,
48 size_t num_threads = 1,
49 bool fast_quantization = false
50 ) {
51 auto data_array = ensure_2d_array<float>(data, "data");
52 auto centroids_array = ensure_2d_array<float>(centroids, "centroids");
53 auto cluster_ids_array = ensure_1d_array<rabitqlib::PID>(cluster_ids, "cluster_ids");
54
55 if (static_cast<size_t>(data_array.shape(1)) != dim_) {
56 throw std::invalid_argument("data dimension does not match index dim");
57 }
58 if (static_cast<size_t>(centroids_array.shape(1)) != dim_) {
59 throw std::invalid_argument("centroid dimension does not match index dim");
60 }
61 if (static_cast<size_t>(cluster_ids_array.shape(0)) != static_cast<size_t>(data_array.shape(0))) {
62 throw std::invalid_argument("cluster_ids length must match number of rows in data");
63 }
64
65 const size_t num_clusters = static_cast<size_t>(centroids_array.shape(0));
66 num_clusters_ = num_clusters;
67
68 // Ensure cluster_ids are writable for the C++ API by making a copy
69 std::vector<rabitqlib::PID> cluster_ids_vec(static_cast<size_t>(cluster_ids_array.shape(0)));
70 std::memcpy(cluster_ids_vec.data(), cluster_ids_array.data(), cluster_ids_vec.size() * sizeof(rabitqlib::PID));
71
72 index_->construct(
73 num_clusters,
74 centroids_array.data(),
75 static_cast<size_t>(data_array.shape(0)),
76 data_array.data(),
77 cluster_ids_vec.data(),
78 num_threads,
79 fast_quantization
80 );
81 built_ = true;
82 }
83
84 py::tuple search(py::handle queries, size_t k, size_t ef = 0, size_t num_threads = 1) {
85 auto query_array = ensure_2d_array<float>(queries, "queries");

Callers 1

mainFunction · 0.95

Calls 3

dataMethod · 0.45
sizeMethod · 0.45
constructMethod · 0.45

Tested by

no test coverage detected