-
Notifications
You must be signed in to change notification settings - Fork 50
make all neighbors more robust #769
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from 2 commits
2f518f0
0f6798f
2e45b35
ba3e7f5
a6b615d
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -143,13 +143,15 @@ def neighbors( | |
|
|
||
| * 'algo': The algorithm to use. Valid options are: 'ivf_pq' and 'nn_descent'. Default is 'nn_descent'. | ||
|
|
||
| * 'n_clusters': Number of clusters/batches to partition the dataset into (> overlap_factor). Default is number of GPUs. | ||
| * 'n_clusters': Number of clusters/batches to partition the dataset into (> overlap_factor). Default is 1 on a single GPU and the smallest multiple of the device count greater than `overlap_factor` otherwise. | ||
|
|
||
| * 'overlap_factor': Number of clusters each point is assigned to (must be < n_clusters). Default is 1. | ||
| * 'overlap_factor': Number of clusters each point is assigned to (must be < n_clusters). Default is `max(2, ceil(log2(n_clusters)))`. Lower values are faster but lose neighbors at cluster boundaries. | ||
|
|
||
| * 'n_lists': Number of inverted lists for IVF indexing. Default is 2 * next_power_of_2(sqrt(n_samples)). Only available for `ivf_pq` algorithm. | ||
|
|
||
| * 'intermediate_graph_degree': The degree of the intermediate graph. Default is None. It is recommended to set it to `>= 1.5 * n_neighbors`. Only available for `nn_descent` algorithm. | ||
| * 'graph_degree': The degree of the graph nn-descent builds before selecting the final `n_neighbors`. Default is 64, raised to `n_neighbors` if larger. Only available for `nn_descent` algorithm. | ||
|
|
||
| * 'intermediate_graph_degree': The degree of the intermediate graph. Default is 128, raised to `graph_degree` if larger. It is recommended to set it to `>= 1.5 * graph_degree`. Only available for `nn_descent` algorithm. | ||
|
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 🎯 Functional Correctness | 🟠 Major | ⚡ Quick win Document the actual intermediate graph-degree default. The runtime uses As per coding guidelines, public functions must have accurate docstrings with documented parameters and notes about GPU-specific behavior differences where relevant. 🤖 Prompt for AI AgentsSource: Coding guidelines |
||
|
|
||
| For `mg_ivfflat` and `mg_ivfpq` algorithms, the following parameters can be specified: | ||
|
|
||
|
|
@@ -208,7 +210,7 @@ def neighbors( | |
| ) | ||
|
|
||
| X = _choose_representation(adata, use_rep=use_rep, n_pcs=n_pcs) | ||
| X_contiguous = _check_neighbors_X(X, algorithm) | ||
| X_contiguous = _check_neighbors_X(X, algorithm, algorithm_kwds) | ||
| _check_metrics(algorithm, metric) | ||
|
|
||
| knn_indices, knn_dist = KNN_ALGORITHMS[algorithm]( | ||
|
|
@@ -402,7 +404,7 @@ def bbknn( | |
| adata._init_as_actual(adata.copy()) | ||
|
|
||
| X = _choose_representation(adata, use_rep=use_rep, n_pcs=n_pcs) | ||
| X_contiguous = _check_neighbors_X(X, algorithm) | ||
| X_contiguous = _check_neighbors_X(X, algorithm, algorithm_kwds) | ||
| _check_metrics(algorithm, metric) | ||
|
|
||
| n_obs = adata.shape[0] | ||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
| @@ -1,5 +1,6 @@ | ||
| from __future__ import annotations | ||
|
|
||
| import math | ||
| from typing import TYPE_CHECKING | ||
|
|
||
| import cupy as cp | ||
|
|
@@ -17,6 +18,36 @@ | |
| from rapids_singlecell.preprocessing._neighbors import _Metrics | ||
|
|
||
|
|
||
| def _default_overlap_factor(n_clusters: int) -> int: | ||
| """Overlap needed to hold recall as the dataset is split into more clusters.""" | ||
| if n_clusters <= 1: | ||
| return 1 | ||
| return max(2, math.ceil(math.log2(n_clusters))) | ||
|
|
||
|
|
||
| def _all_neighbors_batching(algorithm_kwds: Mapping) -> tuple[int, int]: | ||
| """Resolve ``(n_clusters, overlap_factor)`` for the cuVS all-neighbors build.""" | ||
| n_devices = cp.cuda.runtime.getDeviceCount() | ||
| n_clusters = algorithm_kwds.get("n_clusters") | ||
| overlap_factor = algorithm_kwds.get("overlap_factor") | ||
| if n_clusters is None: | ||
| n_clusters = 1 if n_devices == 1 else n_devices | ||
| while n_clusters > 1 and n_clusters <= ( | ||
| _default_overlap_factor(n_clusters) | ||
| if overlap_factor is None | ||
| else overlap_factor | ||
| ): | ||
| n_clusters += n_devices | ||
| if overlap_factor is None: | ||
| overlap_factor = _default_overlap_factor(n_clusters) | ||
| if n_clusters > 1 and overlap_factor >= n_clusters: | ||
| raise ValueError( | ||
| f"'n_clusters' ({n_clusters}) must be greater than 'overlap_factor' " | ||
| f"({overlap_factor}) when batching the all_neighbors build." | ||
| ) | ||
|
Comment on lines
+31
to
+49
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 🎯 Functional Correctness | 🟠 Major | ⚡ Quick win Validate batching values before returning them.
Require positive integer values for both settings before applying the relative bound. Preserve the single-cluster exception for 🤖 Prompt for AI Agents |
||
| return n_clusters, overlap_factor | ||
|
|
||
|
|
||
| def _all_neighbors_knn( | ||
| X: np.ndarray, | ||
| Y: np.ndarray, | ||
|
|
@@ -41,23 +72,32 @@ def _all_neighbors_knn( | |
| from cuvs.common import MultiGpuResources | ||
|
|
||
| res = MultiGpuResources() | ||
| n_clusters = algorithm_kwds.get("n_clusters", n_devices) | ||
| overlap_factor = algorithm_kwds.get("overlap_factor", 1) | ||
| n_clusters, overlap_factor = _all_neighbors_batching(algorithm_kwds) | ||
| cuvs_metric = "sqeuclidean" if metric == "euclidean" else metric | ||
| if algo == "ivf_pq" or algo == "ivfpq": | ||
| from cuvs.neighbors import ivf_pq | ||
|
|
||
| algo = "ivf_pq" | ||
| if cuvs_metric != "sqeuclidean": | ||
| raise ValueError( | ||
| f"all_neighbors with algo='ivf_pq' only supports 'euclidean' and " | ||
| f"'sqeuclidean' metrics, got {metric!r}. Use algo='nn_descent' instead." | ||
| ) | ||
| n_lists = algorithm_kwds.get("n_lists", _compute_nlist(X.shape[0])) | ||
| ivf_pq_params = ivf_pq.IndexParams(n_lists=n_lists) | ||
| ivf_pq_params = ivf_pq.IndexParams(n_lists=n_lists, metric=cuvs_metric) | ||
| nn_descent_params = None | ||
| elif algo == "nn_descent": | ||
| from cuvs.neighbors import nn_descent | ||
|
|
||
| graph_degree = max(algorithm_kwds.get("graph_degree", 64), k) | ||
| intermediate_graph_degree = algorithm_kwds.get( | ||
| "intermediate_graph_degree", None | ||
| "intermediate_graph_degree", max(128, int(1.5 * graph_degree)) | ||
| ) | ||
| intermediate_graph_degree = max(intermediate_graph_degree, graph_degree) | ||
|
Comment on lines
+93
to
+97
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 🎯 Functional Correctness | 🟠 Major | ⚡ Quick win Treat explicit When callers pass Use an explicit As per coding guidelines, use 🤖 Prompt for AI AgentsSource: Coding guidelines |
||
| nn_descent_params = nn_descent.IndexParams( | ||
| graph_degree=k, intermediate_graph_degree=intermediate_graph_degree | ||
| graph_degree=graph_degree, | ||
| intermediate_graph_degree=intermediate_graph_degree, | ||
| metric=cuvs_metric, | ||
| ) | ||
| ivf_pq_params = None | ||
| else: | ||
|
|
@@ -66,7 +106,7 @@ def _all_neighbors_knn( | |
| algo=algo, | ||
| overlap_factor=overlap_factor, | ||
| n_clusters=n_clusters, | ||
| metric="sqeuclidean", | ||
| metric=cuvs_metric, | ||
| ivf_pq_params=ivf_pq_params, | ||
| nn_descent_params=nn_descent_params, | ||
| ) | ||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -83,6 +83,71 @@ def test_all_neighbors(algo): | |
| _calc_recall(distances, adata.obsp["distances"], tolerance=tolerance) | ||
|
|
||
|
|
||
| @pytest.mark.parametrize( | ||
| ("n_devices", "expected"), | ||
| [(1, (1, 1)), (2, (4, 2)), (3, (3, 2)), (4, (4, 2)), (8, (8, 3)), (16, (16, 4))], | ||
| ) | ||
| def test_all_neighbors_batching_defaults(monkeypatch, n_devices, expected): | ||
| import cupy as cp | ||
|
|
||
| from rapids_singlecell.preprocessing._neighbors._algorithms._all_neighbors import ( | ||
| _all_neighbors_batching, | ||
| ) | ||
|
|
||
| monkeypatch.setattr(cp.cuda.runtime, "getDeviceCount", lambda: n_devices) | ||
| n_clusters, overlap_factor = _all_neighbors_batching({}) | ||
| assert (n_clusters, overlap_factor) == expected | ||
| assert n_clusters == 1 or overlap_factor < n_clusters | ||
|
|
||
|
|
||
| def test_all_neighbors_batching_overrides(): | ||
| from rapids_singlecell.preprocessing._neighbors._algorithms._all_neighbors import ( | ||
| _all_neighbors_batching, | ||
| ) | ||
|
|
||
| assert _all_neighbors_batching({"n_clusters": 16}) == (16, 4) | ||
| assert _all_neighbors_batching({"n_clusters": 8, "overlap_factor": 2}) == (8, 2) | ||
| with pytest.raises(ValueError, match="must be greater than"): | ||
| _all_neighbors_batching({"n_clusters": 3, "overlap_factor": 3}) | ||
|
|
||
|
|
||
| @pytest.mark.parametrize("n_clusters", [4, 8]) | ||
| def test_all_neighbors_batched(n_clusters): | ||
| """These recall 0.91 and 0.47 with the previous ``overlap_factor=1``.""" | ||
| if parse_version(cuvs.__version__) <= parse_version("25.08"): | ||
| pytest.skip("Skipping All-Neighbors") | ||
| adata = pbmc68k_reduced() | ||
| rsc.pp.neighbors( | ||
| adata, | ||
| n_pcs=50, | ||
| n_neighbors=15, | ||
| algorithm="all_neighbors", | ||
| algorithm_kwds={"n_clusters": n_clusters}, | ||
| ) | ||
| distances = adata.obsp["distances"].copy() | ||
| rsc.pp.neighbors(adata, n_pcs=50, n_neighbors=15, algorithm="brute") | ||
| _calc_recall(distances, adata.obsp["distances"], tolerance=0.95) | ||
|
|
||
|
|
||
| @pytest.mark.parametrize("n_clusters", [1, 4]) | ||
| @pytest.mark.parametrize("metric", ["cosine", "sqeuclidean"]) | ||
| def test_all_neighbors_metrics(metric, n_clusters): | ||
| if parse_version(cuvs.__version__) <= parse_version("25.08"): | ||
| pytest.skip("Skipping All-Neighbors") | ||
| adata = pbmc68k_reduced() | ||
| rsc.pp.neighbors( | ||
| adata, | ||
| n_pcs=50, | ||
| n_neighbors=15, | ||
| algorithm="all_neighbors", | ||
| metric=metric, | ||
| algorithm_kwds={"n_clusters": n_clusters}, | ||
| ) | ||
| distances = adata.obsp["distances"].copy() | ||
| rsc.pp.neighbors(adata, n_pcs=50, n_neighbors=15, algorithm="brute", metric=metric) | ||
| _calc_recall(distances, adata.obsp["distances"], tolerance=0.95) | ||
|
Comment on lines
+136
to
+152
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. 🎯 Functional Correctness | 🟠 Major | 🏗️ Heavy lift Add independent coverage for The new Add an As per coding guidelines, tests must validate numerical correctness against scanpy, squidpy, pertpy, or SciPy references rather than only checking that code runs. 🤖 Prompt for AI AgentsSources: Coding guidelines, Path instructions |
||
|
|
||
|
|
||
| @pytest.mark.parametrize("algo", ["mg_ivfflat", "mg_ivfpq"]) | ||
| def test_mg_bbknn(algo): | ||
| if parse_version(cuvs.__version__) <= parse_version("25.08"): | ||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
🎯 Functional Correctness | 🟠 Major | ⚡ Quick win
Qualify metric support by all-neighbors subalgorithm.
This entry implies that every
all_neighborsconfiguration supportscosineandinner_product. The IVF-PQ branch rejects both metrics and only accepts squared Euclidean. State that these metrics are available withalgo="nn_descent", or state the IVF-PQ limitation.As per path instructions, check accuracy of code examples and consistency with current code.
🤖 Prompt for AI Agents
Source: Path instructions