Spec-Zone.ru › scikit-learn

Примечание

Перейти к концу для загрузки полного примера кода. или для запуска этого примера в вашем браузере через JupyterLite или Binder

Кэширование ближайших соседей

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

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

Classification accuracy, Fit time (with caching)
# 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 секунды)

Launch binder
Launch JupyterLite

Download Jupyter notebook: plot_caching_nearest_neighbors.ipynb

Download Python source code: plot_caching_nearest_neighbors.py

Download zipped: plot_caching_nearest_neighbors.zip

Связанные примеры

Сравнение ближайших соседей с и без анализа компонентов близости

Классификация ближайших соседей

Приближенные ближайшие соседи в TSNE

Основные моменты выпуска scikit-learn 0.22

© 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

Spec-Zone.ru

Настройки Оффлайн Что нового Помощь О нас
Spec-Zone .ru
спецификации, руководства, описания, API