@@ -819,7 +819,7 @@ class SearcherPy {
819819 searcher_->optimize (n_clusters, n_iter, sample_size, seed, num_threads);
820820 }
821821
822- py::object search (py::array query, uint32_t k, float eps = 0 .1f , float rerank_factor = 1 .0f , bool return_distances = false , bool unsorted = false , uint32_t ef = 0 ) {
822+ py::object search (py::array query, uint32_t k, float eps_or_ef = 0 .1f , float rerank_factor = 1 .0f , bool return_distances = false , bool unsorted = false ) {
823823 auto buf = query.request ();
824824 if (buf.ndim != 1 && (buf.ndim != 2 || (buf.shape [0 ] != 1 && buf.shape [1 ] != 1 ))) {
825825 throw std::invalid_argument (" search query must be 1D vector (or 1xDim / Dimx1 2D array)" );
@@ -835,17 +835,18 @@ class SearcherPy {
835835 }
836836
837837 uint32_t count = 0 ;
838- if (ef > 0 ) {
838+ if (eps_or_ef >= 1 .0f ) {
839+ const uint32_t ef = static_cast <uint32_t >(std::lround (eps_or_ef));
839840 if (buf.itemsize == 2 && (buf.format == " H" || buf.format == " h" || buf.format == " e" )) {
840841 count = searcher_->search_ef_f16 (static_cast <const uint16_t *>(buf.ptr ), k, ef, rerank_factor, out_ptr, dist_ptr, unsorted);
841842 } else {
842843 count = searcher_->search_ef_f32 (static_cast <const float *>(buf.ptr ), k, ef, rerank_factor, out_ptr, dist_ptr, unsorted);
843844 }
844845 } else {
845846 if (buf.itemsize == 2 && (buf.format == " H" || buf.format == " h" || buf.format == " e" )) {
846- count = searcher_->search_f16 (static_cast <const uint16_t *>(buf.ptr ), k, eps , rerank_factor, out_ptr, dist_ptr, unsorted);
847+ count = searcher_->search_f16 (static_cast <const uint16_t *>(buf.ptr ), k, eps_or_ef , rerank_factor, out_ptr, dist_ptr, unsorted);
847848 } else {
848- count = searcher_->search_f32 (static_cast <const float *>(buf.ptr ), k, eps , rerank_factor, out_ptr, dist_ptr, unsorted);
849+ count = searcher_->search_f32 (static_cast <const float *>(buf.ptr ), k, eps_or_ef , rerank_factor, out_ptr, dist_ptr, unsorted);
849850 }
850851 }
851852
@@ -855,7 +856,7 @@ class SearcherPy {
855856 return result;
856857 }
857858
858- py::object search_batch (py::array queries, uint32_t k, float eps = 0 .1f , float rerank_factor = 1 .0f , size_t num_threads = 1 , bool return_distances = false , bool unsorted = false , uint32_t ef = 0 ) {
859+ py::object search_batch (py::array queries, uint32_t k, float eps_or_ef = 0 .1f , float rerank_factor = 1 .0f , size_t num_threads = 1 , bool return_distances = false , bool unsorted = false ) {
859860 auto buf = queries.request ();
860861 if (buf.ndim != 2 ) {
861862 throw std::invalid_argument (" search_batch queries must be 2D array" );
@@ -873,17 +874,18 @@ class SearcherPy {
873874
874875 {
875876 py::gil_scoped_release release;
876- if (ef > 0 ) {
877+ if (eps_or_ef >= 1 .0f ) {
878+ const uint32_t ef = static_cast <uint32_t >(std::lround (eps_or_ef));
877879 if (buf.itemsize == 2 && (buf.format == " H" || buf.format == " h" || buf.format == " e" )) {
878880 searcher_->search_batch_ef_f16 (static_cast <const uint16_t *>(buf.ptr ), n_queries, k, ef, rerank_factor, out_ptr, dist_ptr, num_threads, unsorted);
879881 } else {
880882 searcher_->search_batch_ef_f32 (static_cast <const float *>(buf.ptr ), n_queries, k, ef, rerank_factor, out_ptr, dist_ptr, num_threads, unsorted);
881883 }
882884 } else {
883885 if (buf.itemsize == 2 && (buf.format == " H" || buf.format == " h" || buf.format == " e" )) {
884- searcher_->search_batch_f16 (static_cast <const uint16_t *>(buf.ptr ), n_queries, k, eps , rerank_factor, out_ptr, dist_ptr, num_threads, unsorted);
886+ searcher_->search_batch_f16 (static_cast <const uint16_t *>(buf.ptr ), n_queries, k, eps_or_ef , rerank_factor, out_ptr, dist_ptr, num_threads, unsorted);
885887 } else {
886- searcher_->search_batch_f32 (static_cast <const float *>(buf.ptr ), n_queries, k, eps , rerank_factor, out_ptr, dist_ptr, num_threads, unsorted);
888+ searcher_->search_batch_f32 (static_cast <const float *>(buf.ptr ), n_queries, k, eps_or_ef , rerank_factor, out_ptr, dist_ptr, num_threads, unsorted);
887889 }
888890 }
889891 }
@@ -1586,13 +1588,32 @@ PYBIND11_MODULE(deglib_cpp, m) {
15861588 )
15871589 .def (
15881590 " search" , &SearcherPy::search,
1589- py::arg (" query" ), py::arg (" k" ), py::arg (" eps" ) = 0 .1f , py::arg (" rerank_factor" ) = 1 .0f ,
1590- py::arg (" return_distances" ) = false , py::arg (" unsorted" ) = false , py::arg (" ef" ) = 0
1591+ py::arg (" query" ), py::arg (" k" ), py::arg (" eps_or_ef" ) = 0 .1f , py::arg (" rerank_factor" ) = 1 .0f ,
1592+ py::arg (" return_distances" ) = false , py::arg (" unsorted" ) = false ,
1593+ " Search for nearest neighbors of a single query vector.\n\n "
1594+ " Args:\n "
1595+ " query: 1D query vector (or 1xDim 2D array).\n "
1596+ " k: Number of nearest neighbors to return.\n "
1597+ " eps_or_ef: Exploration parameter. Values >= 1.0 are treated as fixed pool size (ef),\n "
1598+ " values < 1.0 are treated as relative distance margin (eps).\n "
1599+ " rerank_factor: Expansion factor for exact candidate reranking (default: 1.0).\n "
1600+ " return_distances: If True, returns a tuple of (indices, distances).\n "
1601+ " unsorted: If True, returns candidates in heap order instead of ascending distance."
15911602 )
15921603 .def (
15931604 " search_batch" , &SearcherPy::search_batch,
1594- py::arg (" queries" ), py::arg (" k" ), py::arg (" eps" ) = 0 .1f , py::arg (" rerank_factor" ) = 1 .0f ,
1595- py::arg (" num_threads" ) = 1 , py::arg (" return_distances" ) = false , py::arg (" unsorted" ) = false , py::arg (" ef" ) = 0
1605+ py::arg (" queries" ), py::arg (" k" ), py::arg (" eps_or_ef" ) = 0 .1f , py::arg (" rerank_factor" ) = 1 .0f ,
1606+ py::arg (" num_threads" ) = 1 , py::arg (" return_distances" ) = false , py::arg (" unsorted" ) = false ,
1607+ " Search for nearest neighbors for a batch of query vectors.\n\n "
1608+ " Args:\n "
1609+ " queries: 2D query array (shape N x Dim).\n "
1610+ " k: Number of nearest neighbors to return per query.\n "
1611+ " eps_or_ef: Exploration parameter. Values >= 1.0 are treated as fixed pool size (ef),\n "
1612+ " values < 1.0 are treated as relative distance margin (eps).\n "
1613+ " rerank_factor: Expansion factor for exact candidate reranking (default: 1.0).\n "
1614+ " num_threads: Number of worker threads for parallel search.\n "
1615+ " return_distances: If True, returns a tuple of (indices, distances).\n "
1616+ " unsorted: If True, returns candidates in heap order instead of ascending distance."
15961617 )
15971618 .def (
15981619 " optimize" , &SearcherPy::optimize,
0 commit comments