Bắt đầu ngayBắt đầu miễn phí

Tinh chỉnh siêu tham số cho random forest

Cũng như mọi mô hình khác, chúng ta muốn tối ưu hiệu năng bằng cách tinh chỉnh các siêu tham số. Random forest có khá nhiều siêu tham số, nhưng thường quan trọng nhất là số lượng đặc trưng được lấy mẫu tại mỗi lần tách, hay max_features trong RandomForestRegressor của thư viện sklearn. Với các mô hình như random forest vốn có tính ngẫu nhiên, bạn cũng nên đặt random_state để kết quả có thể tái lập.

Thông thường, ta có thể dùng phương thức GridSearchCV() của sklearn để tìm siêu tham số, nhưng với chuỗi thời gian tài chính, ta không muốn dùng cross-validation vì dễ trộn lẫn dữ liệu. Ta muốn huấn luyện mô hình trên dữ liệu cũ nhất và đánh giá trên dữ liệu mới nhất. Vì vậy, ta sẽ dùng ParameterGrid của sklearn để tạo các kết hợp siêu tham số cần tìm.

Bài tập này là một phần của khóa học

Machine Learning cho Tài chính bằng Python

Xem khóa học

Hướng dẫn bài tập

  • Đặt siêu tham số n_estimators là một danh sách với một giá trị (200) trong dictionary grid.
  • Đặt siêu tham số max_features là một danh sách chứa 4 và 8 trong dictionary grid.
  • Huấn luyện mô hình random forest regressor (rfr, đã được tạo sẵn) với train_featurestrain_targets cho mỗi tổ hợp siêu tham số g trong vòng lặp.
  • Tính R\(^2\) bằng cách dùng rfr.score() trên test_features và thêm kết quả vào danh sách test_scores.

Bài tập tương tác thực hành trực tiếp

Hãy thử làm bài tập này bằng cách hoàn thành đoạn mã mẫu này.

from sklearn.model_selection import ParameterGrid

# Create a dictionary of hyperparameters to search
grid = {____, 'max_depth': [3], 'max_features': ____, 'random_state': [42]}
test_scores = []

# Loop through the parameter grid, set the hyperparameters, and save the scores
for g in ParameterGrid(grid):
    rfr.set_params(**g)  # ** is "unpacking" the dictionary
    rfr.fit(____, ____)
    test_scores.append(rfr.score(____, ____))

# Find best hyperparameters from the test score and print
best_idx = np.argmax(test_scores)
print(test_scores[best_idx], ParameterGrid(grid)[best_idx])
Chỉnh sửa và Chạy Mã