Grid Search z Scikit-Learn

Strojenie hiperparametrów w Pythonie

Alex Scriven

Data Scientist

Obiekt GridSearchCV

 

Przedstawiamy obiekt 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')
Strojenie hiperparametrów w Pythonie

Kroki w przeszukiwaniu siatki

 

Kroki w przeszukiwaniu siatki:

  1. Algorytm do strojenia hiperparametrów (czasem zwany 'estymatorem')
  2. Określenie, które hiperparametry będą strojone
  3. Określenie zakresu wartości dla każdego hiperparametru
  4. Ustalenie schematu walidacji krzyżowej
  5. Zdefiniowanie funkcji oceniającej, aby wybrać najlepszy model na siatce
  6. Dodanie dodatkowych przydatnych informacji lub funkcji
Strojenie hiperparametrów w Pythonie

Parametry obiektu GridSearchCV

Ważne parametry wejściowe:

  • estimator
  • param_grid
  • cv
  • scoring
  • refit
  • n_jobs
  • return_train_score
Strojenie hiperparametrów w Pythonie

GridSearchCV – 'estimator'

 

Parametr estimator:

  • To zasadniczo nasz algorytm
  • Pracowałeś już z KNN, Random Forest, GBM i regresją logistyczną

 

Pamiętaj:

  • Tylko jeden estymator na obiekt GridSearchCV
Strojenie hiperparametrów w Pythonie

GridSearchCV – 'param_grid'

Parametr param_grid:

  • Określa, które hiperparametry i wartości mają być testowane

Zamiast listy:

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

Zapiszemy to w postaci:

param_grid = {'max_depth': [2, 4, 6, 8],
              'min_samples_leaf': [1, 2, 4, 6]}
Strojenie hiperparametrów w Pythonie

GridSearchCV – 'param_grid'

Parametr param_grid:

Pamiętaj: Klucze w słowniku param_grid muszą być poprawnymi hiperparametrami.

Na przykład dla estymatora regresji logistycznej:

# Incorrect
param_grid = {'C': [0.1,0.2,0.5],
              'best_choice': [10,20,50]}
ValueError: Invalid parameter best_choice for estimator LogisticRegression
Strojenie hiperparametrów w Pythonie

GridSearchCV – 'cv'

Parametr cv:

  • Określa sposób przeprowadzenia walidacji krzyżowej
  • Liczba całkowita oznacza k-krotną walidację krzyżową; standardowo 5 lub 10

k-fold wikipedia

Strojenie hiperparametrów w Pythonie

GridSearchCV – 'scoring'

 

parametr scoring:

  • Metryka służąca do wyboru najlepszego modelu
  • Można użyć własnej lub modułu metrics w Scikit-Learn

Wszystkie wbudowane funkcje oceniające można sprawdzić w następujący sposób:

from sklearn import metrics
sorted(metrics.SCORERS.keys())
Strojenie hiperparametrów w Pythonie

GridSearchCV – 'refit'

 

Parametr refit:

  • Dopasowuje najlepsze hiperparametry do danych treningowych
  • Umożliwia użycie obiektu GridSearchCV jako estymatora (do predykcji)
  • Bardzo przydatna opcja!
Strojenie hiperparametrów w Pythonie

GridSearchCV – 'n_jobs'

Parametr n_jobs:

  • Wspomaga równoległe wykonywanie obliczeń
  • Pozwala tworzyć wiele modeli jednocześnie, zamiast kolejno po sobie

Przydatny kod:

import os
print(os.cpu_count())

Uważaj, aby nie zużywać wszystkich rdzeni do modelowania, jeśli chcesz wykonywać inne zadania!

Strojenie hiperparametrów w Pythonie

GridSearchCV – 'return_train_score'

 

Parametr return_train_score:

  • Rejestruje statystyki dotyczące przeprowadzonych przebiegów treningowych
  • Przydatny do analizy kompromisu między obciążeniem a wariancją, lecz zwiększa koszt obliczeniowy.
  • Nie pomaga w wyborze najlepszego modelu – służy wyłącznie do analizy
Strojenie hiperparametrów w Pythonie

Budowanie obiektu GridSearchCV

 

Budowanie własnego obiektu 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')
Strojenie hiperparametrów w Pythonie

Budowanie obiektu GridSearchCV

 

Łączenie elementów w całość:

grid_rf_class = GridSearchCV(
    estimator = rf_class,
    param_grid = parameter_grid,
    scoring='accuracy',
    n_jobs=4,
    cv = 10,
    refit=True,
    return_train_score=True)
Strojenie hiperparametrów w Pythonie

Używanie obiektu GridSearchCV

 

Ponieważ parametr refit ustawiono na True, można bezpośrednio używać tego obiektu:

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

# Make predictions
grid_rf_class.predict(X_test)
Strojenie hiperparametrów w Pythonie

Czas na ćwiczenia!

Strojenie hiperparametrów w Pythonie

Preparing Video For Download...