Как работает алгоритм KNN (k-nearest neighbors)?

«Как работает алгоритм KNN (k-nearest neighbors)?» — вопрос из категории Классическое ML, который задают на 26% собеседований Data Scientist / ML Инженер. Ниже — развёрнутый ответ с разбором ключевых моментов.

Ответ

KNN (k-ближайших соседей) — это алгоритм для задач классификации и регрессии, основанный на метриках близости. Он не строит явную модель во время обучения, а «запоминает» всю обучающую выборку.

Алгоритм работы:

  1. Вычисление расстояний: Для нового объекта вычисляются расстояния до всех точек обучающей выборки. Чаще всего используется Евклидово расстояние, но могут применяться Манхэттенское, косинусное и другие метрики.
  2. Выбор соседей: Выбираются k объектов с наименьшими расстояниями.
  3. Принятие решения:
    • Для классификации: Присваивается класс, наиболее часто встречающийся среди k соседей (мажоритарное голосование).
    • Для регрессии: Вычисляется среднее (или медиана) значений целевой переменной соседей.

Пример классификации с использованием scikit-learn:

from sklearn.neighbors import KNeighborsClassifier
from sklearn.preprocessing import StandardScaler
from sklearn.model_selection import train_test_split

# Масштабирование данных критически важно для KNN
scaler = StandardScaler()
X_train_scaled = scaler.fit_transform(X_train)
X_test_scaled = scaler.transform(X_test)

# Создание и обучение модели
model = KNeighborsClassifier(n_neighbors=5, metric='euclidean')
model.fit(X_train_scaled, y_train)

# Предсказание
predictions = model.predict(X_test_scaled)

Ключевые особенности и настройки:

  • Чувствительность к масштабу: Обязательна нормализация или стандартизация признаков.
  • Выбор k: Слишком малое k (например, 1) ведет к переобучению и высокой чувствительности к шуму. Слишком большое k сглаживает модель и может привести к недообучению. Оптимальное k подбирается через кросс-валидацию.
  • Вычислительная сложность: Предсказание медленное на больших данных, так как требует вычисления расстояний до всех обучающих примеров. Для ускорения используются структуры данных вроде KD-деревьев или шаровых деревьев.
  • Вес соседей: Часто используется взвешенное голосование, где голос ближайшего соседа имеет больший вес.