Кросс-валидация и подбор гиперпараметров в scikit-learn

Один train-validation split может случайно оказаться удобным или трудным. Кросс-валидация повторяет обучение на нескольких фолдах и показывает не только среднее качество, но и…

Один train-validation split может случайно оказаться удобным или трудным. Кросс-валидация повторяет обучение на нескольких фолдах и показывает не только среднее качество, но и его разброс. По этим оценкам выбирают гиперпараметры, сохраняя test для финальной проверки.

Параметры и гиперпараметры

Коэффициенты линейной модели или разбиения дерева находятся алгоритмом во время fit — это параметры. Глубина дерева, число соседей и сила регуляризации задают способ обучения — это гиперпараметры. Их варианты сравнивает внешний цикл model selection.

Каждая проверка гиперпараметра должна заново обучать весь процесс, включая imputer, scaler, selector и модель. Иначе статистика validation просочится в преобразования.

K-fold cross-validation

Данные делятся на k частей. Модель k раз обучается на всех фолдах, кроме одного, и оценивается на оставшемся. Затем значения агрегируются. Для классификации StratifiedKFold сохраняет примерные доли классов.

Обычный KFold не подходит для любых данных. GroupKFold удерживает связанные объекты в одной стороне, TimeSeriesSplit сохраняет временной порядок. Стратегию выбирают по тому, что будет новым после запуска.

GridSearchCV с Pipeline

from sklearn.datasets import load_breast_cancer
from sklearn.linear_model import LogisticRegression
from sklearn.model_selection import GridSearchCV, StratifiedKFold, train_test_split
from sklearn.pipeline import make_pipeline
from sklearn.preprocessing import StandardScaler

X, y = load_breast_cancer(return_X_y=True)
X_train, X_test, y_train, y_test = train_test_split(
    X, y, test_size=0.2, random_state=42, stratify=y
)
pipeline = make_pipeline(
    StandardScaler(),
    LogisticRegression(max_iter=2_000),
)

search = GridSearchCV(
    pipeline,
    param_grid={
        'logisticregression__C': [0.01, 0.1, 1.0, 10.0],
        'logisticregression__class_weight': [None, 'balanced'],
    },
    scoring='f1_macro',
    cv=StratifiedKFold(n_splits=5, shuffle=True, random_state=42),
    n_jobs=-1,
    refit=True,
)
search.fit(X_train, y_train)

print(search.best_params_)
print(search.best_score_)
print(search.score(X_test, y_test))

После выбора refit=True обучает лучший Pipeline на всём X_train. Вызов score для теста должен быть редким финальным действием, а не частью дальнейшего перебора.

Среднее без разброса недостаточно

Модели со средним F1 0,82 могут иметь стандартное отклонение 0,01 и 0,12. Второй вариант сильно зависит от состава данных. Посмотрите результаты всех фолдов и исследуйте, какие группы вызывают провалы.

Разница в тысячные доли часто меньше естественного разброса. Выберите более простую, быструю или устойчивую модель, если статистических оснований для сложной нет.

Цена большого поиска

Сетка из 20×10×5 комбинаций содержит 1000 кандидатов и при пяти фолдах означает 5000 CV-обучений. При refit=True после них выполняется ещё одно обучение победителя на всём train. RandomizedSearchCV исследует ограниченное число случайных комбинаций, а successive halving постепенно отбрасывает слабые варианты. Но любой масштабный перебор повышает риск подгонки к CV.

Фиксируйте пространство поиска заранее, сохраняйте таблицу экспериментов и не расширяйте его бесконечно после каждого взгляда на результат. При особенно важной оценке применяют nested cross-validation.

Практика: спроектируйте честный CV

Для истории покупок одного пользователя во многих строках сравните StratifiedKFold и GroupKFold. Объясните, почему первый может завысить качество. Затем создайте маленькую сетку из двух гиперпараметров, посчитайте число обучений и назовите критерий остановки поиска.

Что важно запомнить

Источники