Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions lib/CMakeLists.txt
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
248 changes: 119 additions & 129 deletions lib/src/nclssearch.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -27,38 +27,29 @@ struct NCListSearchHandler {
Index n = starts_buf.shape[0];

nclist_obj = nclist::build<Index, Position>(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<Index, Position> nclist_obj;

std::vector<Position> self_starts;
std::vector<Position> self_ends;
std::vector<std::pair<Position, Index>> sorted_starts;
std::vector<std::pair<Position, Index>> sorted_ends;
};

py::object perform_follow(
pybind11::tuple perform_follow(
NCListSearchHandler &self,
py::array_t<Position> query_starts,
py::array_t<Position> query_ends,
const std::string& select,
int num_threads = 1) {

auto q_starts_ptr = static_cast<const Position*>(query_starts.request().ptr);
auto q_ends_ptr = static_cast<const Position*>(query_ends.request().ptr);

Index n_queries = query_starts.request().shape[0];

std::vector<Index> 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<std::vector<Index> > all_results(n_queries);
std::vector<std::thread> workers;
workers.reserve(num_threads);

Expand All @@ -73,17 +64,21 @@ py::object perform_follow(
}

workers.emplace_back([&](int first, int length) -> void {
std::vector<Index> nearest_matches;

for (Index i = first, last = first + length; i < last; ++i) {
Position q_start = q_starts_ptr[i];
nclist::NearestWorkspace<Index> ws_nearest;
nclist::NearestParameters<Position> 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);
Expand All @@ -94,29 +89,46 @@ py::object perform_follow(
worker.join();
}

if (select == "last") {
return py::array_t<Index>(n_queries, results.data());
} else {
std::vector<Index> 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<Index> self_hits(total_hits);
py::array_t<Index> query_hits(total_hits);
auto s_res_ptr = static_cast<Index*>(self_hits.request().ptr);
auto q_res_ptr = static_cast<Index*>(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<Index>(q_hits.size(), q_hits.data()), py::array_t<Index>(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<Position> query_starts,
py::array_t<Position> query_ends,
const std::string& select,
int num_threads = 1) {

auto q_starts_ptr = static_cast<const Position*>(query_starts.request().ptr);
auto q_ends_ptr = static_cast<const Position*>(query_ends.request().ptr);
Index n_queries = query_ends.request().shape[0];
Index n_queries = query_starts.request().shape[0];

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<Index> results(n_queries);
std::vector<std::vector<Index> > all_results(n_queries);
std::vector<std::thread> workers;
workers.reserve(num_threads);

Expand All @@ -131,16 +143,21 @@ py::object perform_precede(
}

workers.emplace_back([&](int first, int length) -> void {
std::vector<Index> nearest_matches;

for (Index i = first, last = first + length; i < last; ++i) {
Position q_end = q_ends_ptr[i];
nclist::NearestWorkspace<Index> ws_nearest;
nclist::NearestParameters<Position> 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);
Expand All @@ -151,21 +168,29 @@ py::object perform_precede(
worker.join();
}

if (select == "first") {
return py::array_t<Index>(n_queries, results.data());
} else {
std::vector<Index> 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<Index> self_hits(total_hits);
py::array_t<Index> query_hits(total_hits);
auto s_res_ptr = static_cast<Index*>(self_hits.request().ptr);
auto q_res_ptr = static_cast<Index*>(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<Index>(q_hits.size(), q_hits.data()), py::array_t<Index>(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<Position> query_starts,
py::array_t<Position> query_ends,
Expand All @@ -175,6 +200,12 @@ py::object perform_nearest(
auto q_ends_ptr = static_cast<const Position*>(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<std::vector<Index>> all_results(n_queries);
std::vector<std::thread> workers;
workers.reserve(num_threads);
Expand All @@ -190,57 +221,22 @@ py::object perform_nearest(
}

workers.emplace_back([&](int first, int length) -> void {
nclist::OverlapsAnyWorkspace<Index> ws_any;
nclist::OverlapsAnyParameters<Position> ol_params;
ol_params.min_overlap = 1;
std::vector<Index> 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<std::pair<Position, Index>> candidates;

std::vector<Index> 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<Index>::max()));
if (it_b != self.sorted_ends.begin()) {
--it_b;
candidates.emplace_back(q_start - it_b->first, it_b->second);
}
nclist::NearestWorkspace<Index> ws_nearest;
nclist::NearestParameters<Position> 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<Index> 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;
Expand All @@ -250,34 +246,26 @@ py::object perform_nearest(
worker.join();
}

if (select == "arbitrary") {
py::array_t<Index> result(n_queries);
auto res_ptr = static_cast<Index*>(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<std::pair<Index, Index>> 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<Index> query_hits(final_pairs.size());
py::array_t<Index> self_hits(final_pairs.size());
auto q_res_ptr = static_cast<Index*>(query_hits.request().ptr);
auto s_res_ptr = static_cast<Index*>(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<Index> self_hits(total_hits);
py::array_t<Index> query_hits(total_hits);
auto s_res_ptr = static_cast<Index*>(self_hits.request().ptr);
auto q_res_ptr = static_cast<Index*>(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);
}


Expand All @@ -288,13 +276,15 @@ 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,
"Find nearest positions that are downstream/follow each query range.")

.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.")
Expand Down
Loading
Loading