Tìm kiếm lưới với Scikit Learn

Tinh chỉnh siêu tham số trong Python

Alex Scriven

Data Scientist

Đối tượng GridSearchCV

 

Giới thiệu đối tượng 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')
Tinh chỉnh siêu tham số trong Python

Các bước trong Grid Search

 

Các bước của Grid Search:

  1. Chọn thuật toán để tinh chỉnh siêu tham số (estimator)
  2. Xác định siêu tham số sẽ tinh chỉnh
  3. Đặt phạm vi giá trị cho mỗi siêu tham số
  4. Chọn sơ đồ cross-validation
  5. Chọn hàm điểm để quyết định ô lưới “tốt nhất”
  6. Thêm thông tin hoặc chức năng bổ trợ
Tinh chỉnh siêu tham số trong Python

Các đầu vào của đối tượng GridSearchCV

Các đầu vào quan trọng:

  • estimator
  • param_grid
  • cv
  • scoring
  • refit
  • n_jobs
  • return_train_score
Tinh chỉnh siêu tham số trong Python

'estimator' của GridSearchCV

 

Đầu vào estimator:

  • Về cơ bản là thuật toán của chúng ta
  • Bạn đã dùng KNN, Random Forest, GBM, Logistic Regression

 

Lưu ý:

  • Mỗi đối tượng GridSearchCV chỉ có một estimator
Tinh chỉnh siêu tham số trong Python

'param_grid' của GridSearchCV

Đầu vào param_grid:

  • Xác định siêu tham số và các giá trị cần thử

Thay vì danh sách:

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

Sẽ là:

param_grid = {'max_depth': [2, 4, 6, 8],
              'min_samples_leaf': [1, 2, 4, 6]}
Tinh chỉnh siêu tham số trong Python

'param_grid' của GridSearchCV

Đầu vào param_grid:

Lưu ý: Các khóa trong từ điển param_grid phải là siêu tham số hợp lệ.

Ví dụ với bộ ước lượng Logistic Regression:

# Incorrect
param_grid = {'C': [0.1,0.2,0.5],
              'best_choice': [10,20,50]}
ValueError: Invalid parameter best_choice for estimator LogisticRegression
Tinh chỉnh siêu tham số trong Python

'cv' của GridSearchCV

Đầu vào cv:

  • Cách thực hiện cross-validation
  • Dùng số nguyên sẽ là k-fold cross-validation; thường dùng 5 hoặc 10

k-fold wikipedia

Tinh chỉnh siêu tham số trong Python

'scoring' của GridSearchCV

 

Đầu vào scoring:

  • Chọn điểm số để chọn ô lưới (mô hình) tốt nhất
  • Dùng hàm riêng hoặc module metrics của Scikit Learn

Bạn có thể xem tất cả hàm chấm điểm dựng sẵn như sau:

from sklearn import metrics
sorted(metrics.SCORERS.keys())
Tinh chỉnh siêu tham số trong Python

'refit' của GridSearchCV

 

Đầu vào refit:

  • Fit bộ siêu tham số tốt nhất trên dữ liệu huấn luyện
  • Cho phép dùng đối tượng GridSearchCV như một estimator (để dự đoán)
  • Rất hữu ích!
Tinh chỉnh siêu tham số trong Python

'n_jobs' của GridSearchCV

Đầu vào n_jobs:

  • Hỗ trợ chạy song song
  • Tạo nhiều mô hình cùng lúc thay vì tuần tự

Đoạn mã hữu ích:

import os
print(os.cpu_count())

Cẩn thận dùng hết lõi CPU nếu còn làm việc khác!

Tinh chỉnh siêu tham số trong Python

'return_train_score' của GridSearchCV

 

Đầu vào return_train_score:

  • Ghi lại thống kê về các lần huấn luyện đã chạy
  • Hữu ích để phân tích đánh đổi bias-variance nhưng tốn tài nguyên
  • Không giúp chọn mô hình tốt nhất, chỉ để phân tích
Tinh chỉnh siêu tham số trong Python

Xây dựng đối tượng GridSearchCV

 

Tự xây dựng đối tượng 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')
Tinh chỉnh siêu tham số trong Python

Xây dựng đối tượng GridSearchCv

 

Ghép các phần lại:

grid_rf_class = GridSearchCV(
    estimator = rf_class,
    param_grid = parameter_grid,
    scoring='accuracy',
    n_jobs=4,
    cv = 10,
    refit=True,
    return_train_score=True)
Tinh chỉnh siêu tham số trong Python

Dùng đối tượng GridSearchCV

 

Vì đặt refitTrue, ta có thể dùng trực tiếp đối tượng:

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

# Make predictions
grid_rf_class.predict(X_test)
Tinh chỉnh siêu tham số trong Python

Ayo berlatih!

Tinh chỉnh siêu tham số trong Python

Preparing Video For Download...