Поиск по сетке с помощью Scikit-Learn

Подбор гиперпараметров в Python

Alex Scriven

Data Scientist

Объект GridSearchCV

 

Объект GridSearchCV:

sklearn.model_selection.GridSearchCV(
    estimator,
    param_grid, scoring=None, fit_params=None,
    n_jobs=None, refit=True, cv='warn',
    verbose=0, pre_dispatch='2*n_jobs',
    error_score='raise-deprecating',
    return_train_score='warn')
Подбор гиперпараметров в Python

Шаги поиска по сетке

 

Шаги поиска по сетке:

  1. Алгоритм для настройки гиперпараметров (иногда называемый «оценщиком»).
  2. Выбор настраиваемых гиперпараметров.
  3. Определение диапазона значений для каждого гиперпараметра.
  4. Выбор схемы кросс-валидации.
  5. Определение метрики для выбора наилучшей ячейки сетки.
  6. Добавление дополнительной полезной информации и функций.
Подбор гиперпараметров в Python

Параметры объекта GridSearchCV

Основные параметры:

  • estimator
  • param_grid
  • cv
  • scoring
  • refit
  • n_jobs
  • return_train_score
Подбор гиперпараметров в Python

Параметр «estimator» в GridSearchCV

 

Параметр estimator:

  • Определяет используемый алгоритм.
  • Вы уже работали с KNN, случайным лесом, GBM и логистической регрессией.

 

Помните:

  • Один объект GridSearchCV — один оценщик.
Подбор гиперпараметров в Python

Параметр «param_grid» в GridSearchCV

Параметр param_grid:

  • Задаёт гиперпараметры и значения для перебора.

Вместо отдельных списков:

max_depth_list = [2, 4, 6, 8]
min_samples_leaf_list = [1, 2, 4, 6]

Используется словарь:

param_grid = {'max_depth': [2, 4, 6, 8],
              'min_samples_leaf': [1, 2, 4, 6]}
Подбор гиперпараметров в Python

Параметр «param_grid» в GridSearchCV

Параметр param_grid:

Помните: ключи словаря param_grid должны быть допустимыми гиперпараметрами.

Например, для логистической регрессии:

# Incorrect
param_grid = {'C': [0.1,0.2,0.5],
              'best_choice': [10,20,50]}
ValueError: Invalid parameter best_choice for estimator LogisticRegression
Подбор гиперпараметров в Python

Параметр «cv» в GridSearchCV

Параметр cv:

  • Определяет способ проведения кросс-валидации.
  • Целое число задаёт k-fold кросс-валидацию; стандартные значения — 5 или 10.

k-fold wikipedia

Подбор гиперпараметров в Python

Параметр «scoring» в GridSearchCV

 

Параметр scoring:

  • Метрика для выбора наилучшей ячейки сетки (модели).
  • Используйте собственную метрику или модуль metrics из Scikit-Learn.

Просмотреть все встроенные метрики можно так:

from sklearn import metrics
sorted(metrics.SCORERS.keys())
Подбор гиперпараметров в Python

Параметр «refit» в GridSearchCV

 

Параметр refit:

  • Обучает модель с лучшими гиперпараметрами на обучающих данных.
  • Позволяет использовать объект GridSearchCV как оценщик (для предсказаний).
  • Очень удобная опция!
Подбор гиперпараметров в Python

Параметр «n_jobs» в GridSearchCV

Параметр n_jobs:

  • Обеспечивает параллельное выполнение.
  • Позволяет обучать несколько моделей одновременно, а не последовательно.

Полезный код:

import os
print(os.cpu_count())

Будьте осторожны: не используйте все ядра для обучения, если нужно выполнять другие задачи!

Подбор гиперпараметров в Python

Параметр «return_train_score» в GridSearchCV

 

Параметр return_train_score:

  • Сохраняет статистику по выполненным обучающим запускам.
  • Полезен для анализа компромисса между смещением и дисперсией, но увеличивает вычислительные затраты.
  • Не влияет на выбор лучшей модели — служит только для анализа.
Подбор гиперпараметров в Python

Создание объекта GridSearchCV

 

Создаём объект GridSearchCV:

# Create the grid
param_grid = {'max_depth': [2, 4, 6, 8], 'min_samples_leaf': [1, 2, 4, 6]}

#Get a base classifier with some set parameters. rf_class = RandomForestClassifier(criterion='entropy', max_features='auto')
Подбор гиперпараметров в Python

Создание объекта GridSearchCV

 

Собираем всё вместе:

grid_rf_class = GridSearchCV(
    estimator = rf_class,
    param_grid = parameter_grid,
    scoring='accuracy',
    n_jobs=4,
    cv = 10,
    refit=True,
    return_train_score=True)
Подбор гиперпараметров в Python

Использование объекта GridSearchCV

 

Так как refit установлен в True, объект можно использовать напрямую:

#Fit the object to our data
grid_rf_class.fit(X_train, y_train)

# Make predictions
grid_rf_class.predict(X_test)
Подбор гиперпараметров в Python

Давайте потренируемся!

Подбор гиперпараметров в Python

Preparing Video For Download...