diff --git a/cpp/src/neighbors/mg/snmg.cuh b/cpp/src/neighbors/mg/snmg.cuh index 057a3e8272..a6b057693b 100644 --- a/cpp/src/neighbors/mg/snmg.cuh +++ b/cpp/src/neighbors/mg/snmg.cuh @@ -355,11 +355,22 @@ void sharded_search_with_direct_merge( d_trans.view(), raft::make_host_vector_view(h_trans.data(), index.num_ranks_)); + // Results from each rank are packed using the current batch size. Use matching logical views + // so the final partial batch is merged with the same per-rank stride used above. + auto in_distances_batch = raft::make_device_matrix_view( + in_distances.data_handle(), index.num_ranks_ * n_rows_of_current_batch, n_neighbors); + auto in_neighbors_batch = raft::make_device_matrix_view( + in_neighbors.data_handle(), index.num_ranks_ * n_rows_of_current_batch, n_neighbors); + auto out_distances_batch = raft::make_device_matrix_view( + out_distances.data_handle(), n_rows_of_current_batch, n_neighbors); + auto out_neighbors_batch = raft::make_device_matrix_view( + out_neighbors.data_handle(), n_rows_of_current_batch, n_neighbors); + knn_merge_parts(root_handle_, - in_distances.view(), - in_neighbors.view(), - out_distances.view(), - out_neighbors.view(), + in_distances_batch, + in_neighbors_batch, + out_distances_batch, + out_neighbors_batch, d_trans.view()); raft::copy(