| 14 | #include <utility> |
| 15 | #include <vector> |
| 16 | |
| 17 | #include "bindings_common.hpp" |
| 18 | #include "rabitqlib/defines.hpp" |
| 19 | #include "rabitqlib/index/hnsw/hnsw.hpp" |
| 20 | |
| 21 | namespace py = pybind11; |
| 22 | |
| 23 | namespace rabitqlib::python_bindings { |
| 24 | |
| 25 | class HnswIndex { |
| 26 | public: |
| 27 | HnswIndex( |
| 28 | size_t dim, |
| 29 | size_t max_elements, |
| 30 | size_t M, |
| 31 | size_t ef_construction, |
| 32 | size_t nbits, |
| 33 | const std::string& metric = "l2", |
| 34 | size_t random_seed = 100 |
| 35 | ) |
| 36 | : dim_(dim) |
| 37 | , max_elements_(max_elements) |
| 38 | , M_(M) |
| 39 | , ef_construction_(ef_construction) |
| 40 | , nbits_(nbits) |
| 41 | , metric_(metric_from_string(metric)) |
| 42 | , random_seed_(random_seed) |
| 43 | , index_(std::make_unique<rabitqlib::hnsw::HierarchicalNSW>( |
| 44 | max_elements, dim, nbits, M, ef_construction, random_seed, metric_ |
| 45 | )) {} |
| 46 | |
| 47 | void build( |
| 48 | py::handle data, |
| 49 | py::handle centroids, |
| 50 | py::handle cluster_ids, |
| 51 | size_t num_threads = 1, |
| 52 | bool fast_quantization = false |
| 53 | ) { |
| 54 | auto data_array = ensure_2d_array<float>(data, "data"); |
| 55 | auto centroids_array = ensure_2d_array<float>(centroids, "centroids"); |
| 56 | auto cluster_ids_array = |
| 57 | ensure_1d_array<rabitqlib::PID>(cluster_ids, "cluster_ids"); |
| 58 | |
| 59 | if (static_cast<size_t>(data_array.shape(1)) != dim_) { |
| 60 | throw std::invalid_argument("data dimension does not match index dim"); |
| 61 | } |
| 62 | if (data_array.shape(0) == 0 || |
| 63 | static_cast<size_t>(data_array.shape(0)) > max_elements_) { |
| 64 | throw std::invalid_argument( |
| 65 | "number of data rows must be between 1 and max_elements" |
| 66 | ); |
| 67 | } |
| 68 | if (static_cast<size_t>(centroids_array.shape(1)) != dim_) { |
| 69 | throw std::invalid_argument("centroid dimension does not match index dim"); |
| 70 | } |
| 71 | if (static_cast<size_t>(cluster_ids_array.shape(0)) != |
| 72 | static_cast<size_t>(data_array.shape(0))) { |
| 73 | throw std::invalid_argument( |