Примечание
Перейти к концу для загрузки полного примера кода. или для запуска этого примера в вашем браузере через JupyterLite или Binder
Кэширование ближайших соседей
Этот пример демонстрирует, как предварительно вычислить k ближайших соседей перед использованием их в KNeighborsClassifier. KNeighborsClassifier может вычислить ближайших соседей внутри, но предварительное вычисление может иметь несколько преимуществ, таких как более точный контроль параметров, кэширование для многократного использования или настраиваемые реализации.
Здесь мы используем свойство кэширования конвейеров для кэширования графа ближайших соседей между несколькими вызовами KNeighborsClassifier. Первый вызов медленный, поскольку он вычисляет граф соседей, а последующие вызовы быстрее, поскольку им не нужно перевычислять граф. Здесь длительности небольшие, так как набор данных небольшой, но выгода может быть более существенной, когда набор данных становится больше или когда сетка параметров для поиска велика.

# Authors: The scikit-learn developers
# SPDX-License-Identifier: BSD-3-Clause
from tempfile import TemporaryDirectory
import matplotlib.pyplot as plt
from sklearn.datasets import load_digits
from sklearn.model_selection import GridSearchCV
from sklearn.neighbors import KNeighborsClassifier, KNeighborsTransformer
from sklearn.pipeline import Pipeline
X, y = load_digits(return_X_y=True)
n_neighbors_list = [1, 2, 3, 4, 5, 6, 7, 8, 9]
# The transformer computes the nearest neighbors graph using the maximum number
# of neighbors necessary in the grid search. The classifier model filters the
# nearest neighbors graph as required by its own n_neighbors parameter.
graph_model = KNeighborsTransformer(n_neighbors=max(n_neighbors_list), mode="distance")
classifier_model = KNeighborsClassifier(metric="precomputed")
# Note that we give `memory` a directory to cache the graph computation
# that will be used several times when tuning the hyperparameters of the
# classifier.
with TemporaryDirectory(prefix="sklearn_graph_cache_") as tmpdir:
full_model = Pipeline(
steps=[("graph", graph_model), ("classifier", classifier_model)], memory=tmpdir
)
param_grid = {"classifier__n_neighbors": n_neighbors_list}
grid_model = GridSearchCV(full_model, param_grid)
grid_model.fit(X, y)
# Plot the results of the grid search.
fig, axes = plt.subplots(1, 2, figsize=(8, 4))
axes[0].errorbar(
x=n_neighbors_list,
y=grid_model.cv_results_["mean_test_score"],
yerr=grid_model.cv_results_["std_test_score"],
)
axes[0].set(xlabel="n_neighbors", title="Classification accuracy")
axes[1].errorbar(
x=n_neighbors_list,
y=grid_model.cv_results_["mean_fit_time"],
yerr=grid_model.cv_results_["std_fit_time"],
color="r",
)
axes[1].set(xlabel="n_neighbors", title="Fit time (with caching)")
fig.tight_layout()
plt.show()
Общее время выполнения скрипта: (0 минут 1,384 секунды)
Связанные примеры
© 2007–2025 The scikit-learn developers
Licensed under the 3-clause BSD License.
https://scikit-learn.org/1.6/auto_examples/neighbors/plot_caching_nearest_neighbors.html