使用 Scikit Learn 進行 Grid Search

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 超參數調校

Grid Search 步驟

 

Grid Search 的步驟:

  1. 用來調整超參數的演算法(有時稱為「estimator」)
  2. 定義要調整哪些超參數
  3. 為每個超參數設定要測試的數值範圍
  4. 設定交叉驗證方式;以及
  5. 定義評分函式,用來決定網格中哪一格「最佳」
  6. 視需要加入其他資訊或功能
Python 超參數調校

GridSearchCV 物件的輸入參數

重要的輸入參數有:

  • estimator
  • param_grid
  • cv
  • scoring
  • refit
  • n_jobs
  • return_train_score
Python 超參數調校

GridSearchCV 的「estimator」

 

estimator 參數:

  • 基本上就是我們的演算法
  • 你已經用過 KNN、Random Forest、GBM、Logistic Regression

 

提醒:

  • 每個 GridSearchCV 物件只能有一個 estimator
Python 超參數調校

GridSearchCV 的「param_grid」

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 超參數調校

GridSearchCV 的「param_grid」

param_grid 參數:

提醒:param_grid 字典中的鍵必須是有效的超參數名稱。

例如,對 Logistic regression 的 estimator:

# Incorrect
param_grid = {'C': [0.1,0.2,0.5],
              'best_choice': [10,20,50]}
ValueError: Invalid parameter best_choice for estimator LogisticRegression
Python 超參數調校

GridSearchCV 的「cv」

cv 參數:

  • 選擇交叉驗證的方式
  • 若給整數,則進行 k 折交叉驗證;常見為 5 或 10

k-fold wikipedia

Python 超參數調校

GridSearchCV 的「scoring」

 

scoring 參數:

  • 用哪個分數來挑選最佳網格(模型)
  • 可用自訂或 Scikit Learn 的 metrics 模組

你可以這樣查看所有內建評分函式:

from sklearn import metrics
sorted(metrics.SCORERS.keys())
Python 超參數調校

GridSearchCV 的「refit」

 

refit 參數:

  • 以最佳超參數在訓練資料上重新擬合
  • GridSearchCV 物件可直接當成 estimator(用於預測)
  • 非常實用!
Python 超參數調校

GridSearchCV 的「n_jobs」

n_jobs 參數:

  • 協助平行執行
  • 可同時建立多個模型,而非逐一建立

實用程式碼:

import os
print(os.cpu_count())

若還要做其他工作,別把所有核心都拿去訓練模型!

Python 超參數調校

GridSearchCV 的「return_train_score」

 

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...