From 4d51db5d67642f774754ca0876a00c92ef9f5be6 Mon Sep 17 00:00:00 2001 From: Jayaram Kancherla Date: Thu, 11 Sep 2025 22:47:50 -0700 Subject: [PATCH 1/2] switching to nearest from nclist --- lib/CMakeLists.txt | 4 +- lib/src/nclssearch.cpp | 248 ++++++++++++++++++++--------------------- src/iranges/IRanges.py | 27 ++--- tests/test_overlaps.py | 3 +- 4 files changed, 134 insertions(+), 148 deletions(-) diff --git a/lib/CMakeLists.txt b/lib/CMakeLists.txt index 3d4e849..5844918 100644 --- a/lib/CMakeLists.txt +++ b/lib/CMakeLists.txt @@ -11,8 +11,8 @@ include(FetchContent) FetchContent_Declare( nclist GIT_REPOSITORY https://github.com/LTLA/nclist-cpp - GIT_TAG v0.2.0 - # GIT_TAG master + # GIT_TAG v0.2.0 + GIT_TAG master ) FetchContent_MakeAvailable(nclist) diff --git a/lib/src/nclssearch.cpp b/lib/src/nclssearch.cpp index 39a0989..b8d6d00 100644 --- a/lib/src/nclssearch.cpp +++ b/lib/src/nclssearch.cpp @@ -27,38 +27,29 @@ struct NCListSearchHandler { Index n = starts_buf.shape[0]; nclist_obj = nclist::build(n, starts_ptr, ends_ptr); - - self_starts.assign(starts_ptr, starts_ptr + n); - self_ends.assign(ends_ptr, ends_ptr + n); - - sorted_starts.resize(n); - sorted_ends.resize(n); - for (Index i = 0; i < n; ++i) { - sorted_starts[i] = {self_starts[i], i}; - sorted_ends[i] = {self_ends[i], i}; - } - - std::sort(sorted_starts.begin(), sorted_starts.end()); - std::sort(sorted_ends.begin(), sorted_ends.end()); } nclist::Nclist nclist_obj; - - std::vector self_starts; - std::vector self_ends; - std::vector> sorted_starts; - std::vector> sorted_ends; }; -py::object perform_follow( +pybind11::tuple perform_follow( NCListSearchHandler &self, py::array_t query_starts, + py::array_t query_ends, const std::string& select, int num_threads = 1) { + auto q_starts_ptr = static_cast(query_starts.request().ptr); + auto q_ends_ptr = static_cast(query_ends.request().ptr); + Index n_queries = query_starts.request().shape[0]; - std::vector results(n_queries); + bool quit_on_first = (select == "last"); + if (select != "all" && select != "last" && !quit_on_first) { + throw std::runtime_error("Invalid 'select' parameter. Must be 'all' or 'last'."); + } + + std::vector > all_results(n_queries); std::vector workers; workers.reserve(num_threads); @@ -73,17 +64,21 @@ py::object perform_follow( } workers.emplace_back([&](int first, int length) -> void { + std::vector nearest_matches; + for (Index i = first, last = first + length; i < last; ++i) { - Position q_start = q_starts_ptr[i]; + nclist::NearestWorkspace ws_nearest; + nclist::NearestParameters params; + params.quit_on_first = quit_on_first; - // Binary search on sorted ends to find the last subject ending before the query starts - auto it = std::upper_bound(self.sorted_ends.begin(), self.sorted_ends.end(), std::make_pair(q_start, Index(-1))); + nclist::nearest(self.nclist_obj, q_starts_ptr[i], q_ends_ptr[i], params, ws_nearest, nearest_matches); - if (it == self.sorted_ends.begin()) { - results[i] = -1; // No preceding range - } else { - --it; - results[i] = it->second; + if (!nearest_matches.empty()) { + if (select == "last" && !quit_on_first) { + all_results[i] = {nearest_matches.back()}; + } else { + all_results[i] = nearest_matches; + } } } }, jobs_so_far, current_jobs); @@ -94,29 +89,46 @@ py::object perform_follow( worker.join(); } - if (select == "last") { - return py::array_t(n_queries, results.data()); - } else { - std::vector q_hits, s_hits; - for (Index i = 0; i < n_queries; ++i) { - if (results[i] != -1) { - q_hits.push_back(i); - s_hits.push_back(results[i]); - } + + std::size_t total_hits = 0; + for(const auto& res : all_results) { + total_hits += res.size(); + } + + py::array_t self_hits(total_hits); + py::array_t query_hits(total_hits); + auto s_res_ptr = static_cast(self_hits.request().ptr); + auto q_res_ptr = static_cast(query_hits.request().ptr); + + std::size_t current_pos = 0; + for (Index i = 0; i < n_queries; ++i) { + if (!all_results[i].empty()) { + std::copy(all_results[i].begin(), all_results[i].end(), s_res_ptr + current_pos); + std::fill(q_res_ptr + current_pos, q_res_ptr + current_pos + all_results[i].size(), i); + current_pos += all_results[i].size(); } - return py::make_tuple(py::array_t(q_hits.size(), q_hits.data()), py::array_t(s_hits.size(), s_hits.data())); } + + return py::make_tuple(self_hits, query_hits); } -py::object perform_precede( +pybind11::tuple perform_precede( NCListSearchHandler &self, + py::array_t query_starts, py::array_t query_ends, const std::string& select, int num_threads = 1) { + + auto q_starts_ptr = static_cast(query_starts.request().ptr); auto q_ends_ptr = static_cast(query_ends.request().ptr); - Index n_queries = query_ends.request().shape[0]; + Index n_queries = query_starts.request().shape[0]; - std::vector results(n_queries); + bool quit_on_first = (select == "first"); + if (select != "all" && select != "first" && !quit_on_first) { + throw std::runtime_error("Invalid 'select' parameter. Must be 'all' or 'first'."); + } + + std::vector > all_results(n_queries); std::vector workers; workers.reserve(num_threads); @@ -131,16 +143,21 @@ py::object perform_precede( } workers.emplace_back([&](int first, int length) -> void { + std::vector nearest_matches; + for (Index i = first, last = first + length; i < last; ++i) { - Position q_end = q_ends_ptr[i]; + nclist::NearestWorkspace ws_nearest; + nclist::NearestParameters params; + params.quit_on_first = quit_on_first; - // Binary search on sorted starts to find the first subject starting after the query ends - auto it = std::lower_bound(self.sorted_starts.begin(), self.sorted_starts.end(), std::make_pair(q_end, Index(-1))); + nclist::nearest(self.nclist_obj, q_starts_ptr[i], q_ends_ptr[i], params, ws_nearest, nearest_matches); - if (it == self.sorted_starts.end()) { - results[i] = -1; // No following range - } else { - results[i] = it->second; + if (!nearest_matches.empty()) { + if (select == "first" && !quit_on_first) { + all_results[i] = {nearest_matches.back()}; + } else { + all_results[i] = nearest_matches; + } } } }, jobs_so_far, current_jobs); @@ -151,21 +168,29 @@ py::object perform_precede( worker.join(); } - if (select == "first") { - return py::array_t(n_queries, results.data()); - } else { - std::vector q_hits, s_hits; - for (Index i = 0; i < n_queries; ++i) { - if (results[i] != -1) { - q_hits.push_back(i); - s_hits.push_back(results[i]); - } + std::size_t total_hits = 0; + for(const auto& res : all_results) { + total_hits += res.size(); + } + + py::array_t self_hits(total_hits); + py::array_t query_hits(total_hits); + auto s_res_ptr = static_cast(self_hits.request().ptr); + auto q_res_ptr = static_cast(query_hits.request().ptr); + + std::size_t current_pos = 0; + for (Index i = 0; i < n_queries; ++i) { + if (!all_results[i].empty()) { + std::copy(all_results[i].begin(), all_results[i].end(), s_res_ptr + current_pos); + std::fill(q_res_ptr + current_pos, q_res_ptr + current_pos + all_results[i].size(), i); + current_pos += all_results[i].size(); } - return py::make_tuple(py::array_t(q_hits.size(), q_hits.data()), py::array_t(s_hits.size(), s_hits.data())); } + + return py::make_tuple(self_hits, query_hits); } -py::object perform_nearest( +pybind11::tuple perform_nearest( NCListSearchHandler &self, py::array_t query_starts, py::array_t query_ends, @@ -175,6 +200,12 @@ py::object perform_nearest( auto q_ends_ptr = static_cast(query_ends.request().ptr); Index n_queries = query_starts.request().shape[0]; + + bool quit_on_first = (select == "arbitrary"); + if (select != "all" && select != "arbitrary" && !quit_on_first) { + throw std::runtime_error("Invalid 'select' parameter. Must be 'all' or 'arbitrary'."); + } + std::vector> all_results(n_queries); std::vector workers; workers.reserve(num_threads); @@ -190,57 +221,22 @@ py::object perform_nearest( } workers.emplace_back([&](int first, int length) -> void { - nclist::OverlapsAnyWorkspace ws_any; - nclist::OverlapsAnyParameters ol_params; - ol_params.min_overlap = 1; + std::vector nearest_matches; for (Index i = first, last = first + length; i < last; ++i) { - Position q_start = q_starts_ptr[i]; - Position q_end = q_ends_ptr[i]; - - std::vector> candidates; - - std::vector overlaps; - nclist::overlaps_any(self.nclist_obj, q_start, q_end, ol_params, ws_any, overlaps); - for (const auto& ol_idx : overlaps) { - candidates.emplace_back(0, ol_idx); - } - - auto it_b = std::upper_bound(self.sorted_ends.begin(), self.sorted_ends.end(), std::make_pair(q_start, std::numeric_limits::max())); - if (it_b != self.sorted_ends.begin()) { - --it_b; - candidates.emplace_back(q_start - it_b->first, it_b->second); - } + nclist::NearestWorkspace ws_nearest; + nclist::NearestParameters params; + params.quit_on_first = quit_on_first; - auto it_a = std::lower_bound(self.sorted_starts.begin(), self.sorted_starts.end(), std::make_pair(q_end, Index(-1))); - if (it_a != self.sorted_starts.end()) { - candidates.emplace_back(it_a->first - q_end, it_a->second); - } + nclist::nearest(self.nclist_obj, q_starts_ptr[i], q_ends_ptr[i], params, ws_nearest, nearest_matches); - if (candidates.empty()) { - continue; - } - - Position min_dist = candidates[0].first; - for (size_t k = 1; k < candidates.size(); ++k) { - if (candidates[k].first < min_dist) { - min_dist = candidates[k].first; + if (!nearest_matches.empty()) { + if (select == "first" && !quit_on_first) { + all_results[i] = {nearest_matches.back()}; + } else { + all_results[i] = nearest_matches; } } - - std::vector best_hits; - for (const auto& cand : candidates) { - if (cand.first == min_dist) { - best_hits.push_back(cand.second); - } - } - - if (select == "arbitrary") { - std::sort(best_hits.begin(), best_hits.end()); - all_results[i] = {best_hits[0]}; - } else { - all_results[i] = best_hits; - } } }, jobs_so_far, current_jobs); jobs_so_far += current_jobs; @@ -250,34 +246,26 @@ py::object perform_nearest( worker.join(); } - if (select == "arbitrary") { - py::array_t result(n_queries); - auto res_ptr = static_cast(result.request().ptr); - for (Index i = 0; i < n_queries; ++i) { - res_ptr[i] = all_results[i].empty() ? -1 : all_results[i][0]; - } - return std::move(result); - } else { - std::vector> final_pairs; - for (Index i = 0; i < n_queries; ++i) { - for (const auto& self_hit : all_results[i]) { - final_pairs.emplace_back(i, self_hit); - } - } - - std::stable_sort(final_pairs.begin(), final_pairs.end()); - - py::array_t query_hits(final_pairs.size()); - py::array_t self_hits(final_pairs.size()); - auto q_res_ptr = static_cast(query_hits.request().ptr); - auto s_res_ptr = static_cast(self_hits.request().ptr); + std::size_t total_hits = 0; + for(const auto& res : all_results) { + total_hits += res.size(); + } - for(size_t i = 0; i < final_pairs.size(); ++i) { - q_res_ptr[i] = final_pairs[i].first; - s_res_ptr[i] = final_pairs[i].second; + py::array_t self_hits(total_hits); + py::array_t query_hits(total_hits); + auto s_res_ptr = static_cast(self_hits.request().ptr); + auto q_res_ptr = static_cast(query_hits.request().ptr); + + std::size_t current_pos = 0; + for (Index i = 0; i < n_queries; ++i) { + if (!all_results[i].empty()) { + std::copy(all_results[i].begin(), all_results[i].end(), s_res_ptr + current_pos); + std::fill(q_res_ptr + current_pos, q_res_ptr + current_pos + all_results[i].size(), i); + current_pos += all_results[i].size(); } - return py::make_tuple(query_hits, self_hits); } + + return py::make_tuple(self_hits, query_hits); } @@ -288,6 +276,7 @@ void init_nclistsearch(pybind11::module &m){ py::arg("starts"), py::arg("ends")) .def("precede", &perform_precede, + py::arg("query_starts"), py::arg("query_ends"), py::arg("select") = "first", py::arg("num_threads") = 1, @@ -295,6 +284,7 @@ void init_nclistsearch(pybind11::module &m){ .def("follow", &perform_follow, py::arg("query_starts"), + py::arg("query_ends"), py::arg("select") = "last", py::arg("num_threads") = 1, "Find nearest positions that are upstream/precede each query range.") diff --git a/src/iranges/IRanges.py b/src/iranges/IRanges.py index 3c10aa2..cf80b66 100644 --- a/src/iranges/IRanges.py +++ b/src/iranges/IRanges.py @@ -1969,7 +1969,8 @@ def precede( self._build_nclssearch_index() _results = self._nclistsearch.precede( - query.get_end_exclusive().astype(np.int32), + query.get_start().astype(np.int32), + query.get_end_exclusive().astype(np.int32), select=select, num_threads=num_threads, ) @@ -1978,12 +1979,9 @@ def precede( self._delete_nclssearch_index() if select == "first": - _results = np.asarray(_results, dtype=np.object_) - # replace -1 with None - _results[_results == -1] = None - return _results + return _results[0] else: - return BiocFrame(data={"query_hits": _results[0], "self_hits": _results[1]}) + return BiocFrame(data={"query_hits": _results[1], "self_hits": _results[0]}) def follow( self, @@ -2032,6 +2030,7 @@ def follow( _results = self._nclistsearch.follow( query.get_start().astype(np.int32), + query.get_end_exclusive().astype(np.int32), select=select, num_threads=num_threads, ) @@ -2040,12 +2039,9 @@ def follow( self._delete_nclssearch_index() if select == "last": - _results = np.asarray(_results, dtype=np.object_) - # replace -1 with None - _results[_results == -1] = None - return _results + return _results[0] else: - return BiocFrame(data={"query_hits": _results[0], "self_hits": _results[1]}) + return BiocFrame(data={"query_hits": _results[1], "self_hits": _results[0]}) def distance(self, query: "IRanges") -> np.ndarray: """Calculate the pair-wise distance between ranges. @@ -2120,13 +2116,12 @@ def nearest( if delete_index: self._delete_nclssearch_index() + print("right after cpp", _results) + if select == "arbitrary": - _results = np.asarray(_results, dtype=np.object_) - # replace -1 with None - _results[_results == -1] = None - return _results + return _results[0] else: - return BiocFrame(data={"query_hits": _results[0], "self_hits": _results[1]}) + return BiocFrame(data={"query_hits": _results[1], "self_hits": _results[0]}) ######################## #### pandas interop #### diff --git a/tests/test_overlaps.py b/tests/test_overlaps.py index 1ec7358..e097e0a 100644 --- a/tests/test_overlaps.py +++ b/tests/test_overlaps.py @@ -154,6 +154,7 @@ def test_nearest(): subject = IRanges([3, 5, 12], [1, 2, 1]) res = subject.nearest(query) + print(res) assert np.all(res == [0, 0, 2]) res = query.nearest(subject) @@ -189,7 +190,7 @@ def test_edge_cases(): assert np.all(overlaps.get_column("query_hits") == [0, 0, 0, 0, 0, 0, 0, 1, 1, 1, 1, 1, 1, 1]) res = subject.precede(query) - assert np.all(res == [6, 6]) + assert np.all(res == [0, 0]) res = subject.precede(query, select="all") assert np.all(res.get_column("self_hits") == [6, 6]) From 96d293ca58c58e3046fe1cd7eb76754e29dc985e Mon Sep 17 00:00:00 2001 From: "pre-commit-ci[bot]" <66853113+pre-commit-ci[bot]@users.noreply.github.com> Date: Fri, 12 Sep 2025 05:48:19 +0000 Subject: [PATCH 2/2] [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --- lib/src/nclssearch.cpp | 2 +- src/iranges/IRanges.py | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/lib/src/nclssearch.cpp b/lib/src/nclssearch.cpp index b8d6d00..278c904 100644 --- a/lib/src/nclssearch.cpp +++ b/lib/src/nclssearch.cpp @@ -118,7 +118,7 @@ pybind11::tuple perform_precede( py::array_t query_ends, const std::string& select, int num_threads = 1) { - + auto q_starts_ptr = static_cast(query_starts.request().ptr); auto q_ends_ptr = static_cast(query_ends.request().ptr); Index n_queries = query_starts.request().shape[0]; diff --git a/src/iranges/IRanges.py b/src/iranges/IRanges.py index cf80b66..4a34a13 100644 --- a/src/iranges/IRanges.py +++ b/src/iranges/IRanges.py @@ -1970,7 +1970,7 @@ def precede( _results = self._nclistsearch.precede( query.get_start().astype(np.int32), - query.get_end_exclusive().astype(np.int32), + query.get_end_exclusive().astype(np.int32), select=select, num_threads=num_threads, )