Skip to content

Commit 77643f0

Browse files
Reduce NEIGHBORS_ALL_NEIGHBORS_TEST combinatorial test space (#2573)
Moves NEIGHBORS_ALL_NEIGHBORS_TEST runtime from 34 minutes to 4 minutes by making the combinatorial space a little sparse. Primarily this is done by doing: 1. Data locality/movement testing doesn't need to happen across all n_row * dim * graph_degree inputs 2. Make n_row * dim * graph_degree sparser for batched tests. This is done by testing min, max and ~diagonal for this test space. Authors: - Robert Maynard (https://github.com/robertmaynard) - Corey J. Nolet (https://github.com/cjnolet) Approvers: - Tamas Bela Feher (https://github.com/tfeher) URL: #2573
1 parent c740818 commit 77643f0

2 files changed

Lines changed: 144 additions & 30 deletions

File tree

‎cpp/tests/neighbors/all_neighbors.cuh‎

Lines changed: 120 additions & 26 deletions
Original file line numberDiff line numberDiff line change
@@ -279,12 +279,66 @@ const std::vector<AllNeighborsInputs> inputsSingle =
279279
{5000, 7151}, // n_rows
280280
{64, 137}, // dim
281281
{16, 23}, // graph_degree
282-
{false, true}, // data on host
282+
{false}, // data on host
283+
{false}, // mutual_reach
284+
{false} // output on host
285+
);
286+
287+
const std::vector<AllNeighborsInputs> inputsSingleDataTransfer =
288+
raft::util::itertools::product<AllNeighborsInputs>(
289+
{std::make_tuple(BRUTE_FORCE, cuvs::distance::DistanceType::L2Expanded, 0.9),
290+
std::make_tuple(IVF_PQ, cuvs::distance::DistanceType::L2Expanded, 0.9),
291+
std::make_tuple(NN_DESCENT, cuvs::distance::DistanceType::InnerProduct, 0.8)},
292+
{std::make_tuple(1lu, 2lu)}, // min_recall, n_clusters, overlap_factor
293+
{5000}, // n_rows
294+
{137}, // dim
295+
{23}, // graph_degree
296+
{true}, // data on host
283297
{false}, // mutual_reach
284298
{false, true} // output on host
285299
);
286300

287-
const std::vector<AllNeighborsInputs> inputsBatch =
301+
const std::vector<AllNeighborsInputs> inputsBatchLow =
302+
raft::util::itertools::product<AllNeighborsInputs>(
303+
{std::make_tuple(BRUTE_FORCE, cuvs::distance::DistanceType::L2Expanded, 0.9),
304+
std::make_tuple(BRUTE_FORCE, cuvs::distance::DistanceType::L2SqrtExpanded, 0.9),
305+
std::make_tuple(BRUTE_FORCE, cuvs::distance::DistanceType::CosineExpanded, 0.9),
306+
std::make_tuple(BRUTE_FORCE, cuvs::distance::DistanceType::InnerProduct, 0.9),
307+
std::make_tuple(IVF_PQ, cuvs::distance::DistanceType::L2Expanded, 0.9),
308+
std::make_tuple(NN_DESCENT, cuvs::distance::DistanceType::L2Expanded, 0.9),
309+
std::make_tuple(NN_DESCENT, cuvs::distance::DistanceType::L2SqrtExpanded, 0.9),
310+
std::make_tuple(NN_DESCENT, cuvs::distance::DistanceType::CosineExpanded, 0.9),
311+
std::make_tuple(NN_DESCENT, cuvs::distance::DistanceType::InnerProduct, 0.9)},
312+
{std::make_tuple(4lu, 2lu)}, // min_recall, n_clusters, overlap_factor
313+
{5000}, // n_rows
314+
{64, 137}, // dim
315+
{16, 23}, // graph_degree
316+
{true}, // data on host
317+
{false}, // mutual_reach
318+
{true} // output on host
319+
);
320+
321+
const std::vector<AllNeighborsInputs> inputsBatchMed =
322+
raft::util::itertools::product<AllNeighborsInputs>(
323+
{std::make_tuple(BRUTE_FORCE, cuvs::distance::DistanceType::L2Expanded, 0.9),
324+
std::make_tuple(BRUTE_FORCE, cuvs::distance::DistanceType::L2SqrtExpanded, 0.9),
325+
std::make_tuple(BRUTE_FORCE, cuvs::distance::DistanceType::CosineExpanded, 0.9),
326+
std::make_tuple(BRUTE_FORCE, cuvs::distance::DistanceType::InnerProduct, 0.9),
327+
std::make_tuple(IVF_PQ, cuvs::distance::DistanceType::L2Expanded, 0.9),
328+
std::make_tuple(NN_DESCENT, cuvs::distance::DistanceType::L2Expanded, 0.9),
329+
std::make_tuple(NN_DESCENT, cuvs::distance::DistanceType::L2SqrtExpanded, 0.9),
330+
std::make_tuple(NN_DESCENT, cuvs::distance::DistanceType::CosineExpanded, 0.9),
331+
std::make_tuple(NN_DESCENT, cuvs::distance::DistanceType::InnerProduct, 0.9)},
332+
{std::make_tuple(7lu, 2lu)}, // min_recall, n_clusters, overlap_factor
333+
{7151}, // n_rows
334+
{137}, // dim
335+
{23}, // graph_degree
336+
{true}, // data on host
337+
{false}, // mutual_reach
338+
{false} // output on host
339+
);
340+
341+
const std::vector<AllNeighborsInputs> inputsBatchHigh =
288342
raft::util::itertools::product<AllNeighborsInputs>(
289343
{std::make_tuple(BRUTE_FORCE, cuvs::distance::DistanceType::L2Expanded, 0.9),
290344
std::make_tuple(BRUTE_FORCE, cuvs::distance::DistanceType::L2SqrtExpanded, 0.9),
@@ -295,17 +349,13 @@ const std::vector<AllNeighborsInputs> inputsBatch =
295349
std::make_tuple(NN_DESCENT, cuvs::distance::DistanceType::L2SqrtExpanded, 0.9),
296350
std::make_tuple(NN_DESCENT, cuvs::distance::DistanceType::CosineExpanded, 0.9),
297351
std::make_tuple(NN_DESCENT, cuvs::distance::DistanceType::InnerProduct, 0.9)},
298-
{
299-
std::make_tuple(4lu, 2lu),
300-
std::make_tuple(7lu, 2lu),
301-
std::make_tuple(10lu, 2lu),
302-
}, // min_recall, n_clusters, overlap_factor
303-
{5000, 7151}, // n_rows
304-
{64, 137}, // dim
305-
{16, 23}, // graph_degree
306-
{true}, // data on host
307-
{false}, // mutual_reach
308-
{false, true} // output on host
352+
{std::make_tuple(10lu, 2lu)}, // min_recall, n_clusters, overlap_factor
353+
{5000}, // n_rows
354+
{64}, // dim
355+
{16}, // graph_degree
356+
{true}, // data on host
357+
{false}, // mutual_reach
358+
{false} // output on host
309359
);
310360

311361
const std::vector<AllNeighborsInputs> mutualReachSingle =
@@ -320,30 +370,74 @@ const std::vector<AllNeighborsInputs> mutualReachSingle =
320370
{5000, 7151}, // n_rows
321371
{64, 137}, // dim
322372
{16, 23}, // graph_degree
323-
{false, true}, // data on host
373+
{false}, // data on host
374+
{true}, // mutual_reach
375+
{false} // output on host
376+
);
377+
const std::vector<AllNeighborsInputs> mutualReachSingleDataTransfer =
378+
raft::util::itertools::product<AllNeighborsInputs>(
379+
{std::make_tuple(BRUTE_FORCE, cuvs::distance::DistanceType::L2Expanded, 0.9),
380+
std::make_tuple(BRUTE_FORCE, cuvs::distance::DistanceType::L2SqrtExpanded, 0.9),
381+
std::make_tuple(BRUTE_FORCE, cuvs::distance::DistanceType::CosineExpanded, 0.9),
382+
std::make_tuple(NN_DESCENT, cuvs::distance::DistanceType::L2Expanded, 0.9),
383+
std::make_tuple(NN_DESCENT, cuvs::distance::DistanceType::L2SqrtExpanded, 0.9),
384+
std::make_tuple(NN_DESCENT, cuvs::distance::DistanceType::CosineExpanded, 0.9)},
385+
{std::make_tuple(1lu, 2lu)}, // n_clusters, overlap_factor
386+
{5000}, // n_rows
387+
{137}, // dim
388+
{23}, // graph_degree
389+
{true}, // data on host
324390
{true}, // mutual_reach
325391
{false, true} // output on host
326392
);
327393

328-
const std::vector<AllNeighborsInputs> mutualReachBatch =
394+
const std::vector<AllNeighborsInputs> mutualReachBatchLow =
395+
raft::util::itertools::product<AllNeighborsInputs>(
396+
{std::make_tuple(BRUTE_FORCE, cuvs::distance::DistanceType::L2Expanded, 0.9),
397+
std::make_tuple(BRUTE_FORCE, cuvs::distance::DistanceType::L2SqrtExpanded, 0.9),
398+
std::make_tuple(BRUTE_FORCE, cuvs::distance::DistanceType::CosineExpanded, 0.9),
399+
std::make_tuple(NN_DESCENT, cuvs::distance::DistanceType::L2Expanded, 0.9),
400+
std::make_tuple(NN_DESCENT, cuvs::distance::DistanceType::L2SqrtExpanded, 0.9),
401+
std::make_tuple(NN_DESCENT, cuvs::distance::DistanceType::CosineExpanded, 0.9)},
402+
{std::make_tuple(4lu, 2lu)}, // n_clusters, overlap_factor
403+
{5000}, // n_rows
404+
{64, 137}, // dim
405+
{16, 23}, // graph_degree
406+
{true}, // data on host
407+
{true}, // mutual_reach
408+
{true} // output on host
409+
);
410+
const std::vector<AllNeighborsInputs> mutualReachBatchMed =
411+
raft::util::itertools::product<AllNeighborsInputs>(
412+
{std::make_tuple(BRUTE_FORCE, cuvs::distance::DistanceType::L2Expanded, 0.9),
413+
std::make_tuple(BRUTE_FORCE, cuvs::distance::DistanceType::L2SqrtExpanded, 0.9),
414+
std::make_tuple(BRUTE_FORCE, cuvs::distance::DistanceType::CosineExpanded, 0.9),
415+
std::make_tuple(NN_DESCENT, cuvs::distance::DistanceType::L2Expanded, 0.9),
416+
std::make_tuple(NN_DESCENT, cuvs::distance::DistanceType::L2SqrtExpanded, 0.9),
417+
std::make_tuple(NN_DESCENT, cuvs::distance::DistanceType::CosineExpanded, 0.9)},
418+
{std::make_tuple(7lu, 2lu)}, // n_clusters, overlap_factor
419+
{5000}, // n_rows
420+
{137}, // dim
421+
{16}, // graph_degree
422+
{true}, // data on host
423+
{true}, // mutual_reach
424+
{false} // output on host
425+
);
426+
const std::vector<AllNeighborsInputs> mutualReachBatchHigh =
329427
raft::util::itertools::product<AllNeighborsInputs>(
330428
{std::make_tuple(BRUTE_FORCE, cuvs::distance::DistanceType::L2Expanded, 0.9),
331429
std::make_tuple(BRUTE_FORCE, cuvs::distance::DistanceType::L2SqrtExpanded, 0.9),
332430
std::make_tuple(BRUTE_FORCE, cuvs::distance::DistanceType::CosineExpanded, 0.9),
333431
std::make_tuple(NN_DESCENT, cuvs::distance::DistanceType::L2Expanded, 0.9),
334432
std::make_tuple(NN_DESCENT, cuvs::distance::DistanceType::L2SqrtExpanded, 0.9),
335433
std::make_tuple(NN_DESCENT, cuvs::distance::DistanceType::CosineExpanded, 0.9)},
336-
{
337-
std::make_tuple(4lu, 2lu),
338-
std::make_tuple(7lu, 2lu),
339-
std::make_tuple(10lu, 2lu),
340-
}, // n_clusters, overlap_factor
341-
{5000, 7151}, // n_rows
342-
{64, 137}, // dim
343-
{16, 23}, // graph_degree
344-
{true}, // data on host
345-
{true}, // mutual_reach
346-
{false, true} // output on host
434+
{std::make_tuple(10lu, 2lu)}, // n_clusters, overlap_factor
435+
{7151}, // n_rows
436+
{64}, // dim
437+
{23}, // graph_degree
438+
{true}, // data on host
439+
{true}, // mutual_reach
440+
{false} // output on host
347441
);
348442

349443
} // namespace cuvs::neighbors::all_neighbors

‎cpp/tests/neighbors/all_neighbors/test_float.cu‎

Lines changed: 24 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
11
/*
2-
* SPDX-FileCopyrightText: Copyright (c) 2025, NVIDIA CORPORATION.
2+
* SPDX-FileCopyrightText: Copyright (c) 2025-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
33
* SPDX-License-Identifier: Apache-2.0
44
*/
55

@@ -15,14 +15,34 @@ TEST_P(AllNeighborsTestF, AllNeighbors) { this->run(); }
1515
INSTANTIATE_TEST_CASE_P(AllNeighborsSingleTest,
1616
AllNeighborsTestF,
1717
::testing::ValuesIn(inputsSingle));
18+
INSTANTIATE_TEST_CASE_P(AllNeighborsSingleTestDataTransfer,
19+
AllNeighborsTestF,
20+
::testing::ValuesIn(inputsSingleDataTransfer));
1821

19-
INSTANTIATE_TEST_CASE_P(AllNeighborsBatchTest, AllNeighborsTestF, ::testing::ValuesIn(inputsBatch));
22+
INSTANTIATE_TEST_CASE_P(AllNeighborsBatchTestLow,
23+
AllNeighborsTestF,
24+
::testing::ValuesIn(inputsBatchLow));
25+
INSTANTIATE_TEST_CASE_P(AllNeighborsBatchTestMed,
26+
AllNeighborsTestF,
27+
::testing::ValuesIn(inputsBatchMed));
28+
INSTANTIATE_TEST_CASE_P(AllNeighborsBatchTestHigh,
29+
AllNeighborsTestF,
30+
::testing::ValuesIn(inputsBatchHigh));
2031

2132
INSTANTIATE_TEST_CASE_P(AllNeighborsSingleMutualTest,
2233
AllNeighborsTestF,
2334
::testing::ValuesIn(mutualReachSingle));
35+
INSTANTIATE_TEST_CASE_P(AllNeighborsSingleMutualTestDataTransfer,
36+
AllNeighborsTestF,
37+
::testing::ValuesIn(mutualReachSingleDataTransfer));
2438

25-
INSTANTIATE_TEST_CASE_P(AllNeighborsBatchMutualTest,
39+
INSTANTIATE_TEST_CASE_P(AllNeighborsBatchMutualTestLow,
40+
AllNeighborsTestF,
41+
::testing::ValuesIn(mutualReachBatchLow));
42+
INSTANTIATE_TEST_CASE_P(AllNeighborsBatchMutualTestMed,
43+
AllNeighborsTestF,
44+
::testing::ValuesIn(mutualReachBatchMed));
45+
INSTANTIATE_TEST_CASE_P(AllNeighborsBatchMutualTestHigh,
2646
AllNeighborsTestF,
27-
::testing::ValuesIn(mutualReachBatch));
47+
::testing::ValuesIn(mutualReachBatchHigh));
2848
} // namespace cuvs::neighbors::all_neighbors

0 commit comments

Comments
 (0)