Tinh chỉnh siêu tham số trong Python
Alex Scriven
Data Scientist
Một số siêu tham số đáng ưu tiên tinh chỉnh hơn.
Nhưng nên thử các giá trị nào cho siêu tham số?
Cùng xem vài mẹo hàng đầu!
Lưu ý các lựa chọn siêu tham số mâu thuẫn.
LogisticRegression() có các tùy chọn solver và penalty có thể xung đột.The 'newton-cg', 'sag' and 'lbfgs' solvers support only l2 penalties.
Một số không nêu rõ mà sẽ “bỏ qua” (ví dụ ElasticNet với siêu tham số normalize):
This parameter is ignored when fit_intercept is set to False
Hãy tham khảo tài liệu Scikit-learn!
Tránh đặt giá trị “ngớ ngẩn” cho các thuật toán:
Ghi chép các giá trị hợp lý cho siêu tham số là rất hữu ích.
Ở bài trước, ta xây dựng các mô hình:
knn_5 = KNeighborsClassifier(n_neighbors=5)
knn_10 = KNeighborsClassifier(n_neighbors=10)
knn_20 = KNeighborsClassifier(n_neighbors=20)
Cách này kém hiệu quả. Có thể làm tốt hơn không?
Dùng vòng lặp for để duyệt các lựa chọn:
neighbors_list = [3,5,10,20,50,75]accuracy_list = []for test_number in neighbors_list: model = KNeighborsClassifier(n_neighbors=test_number) predictions = model.fit(X_train, y_train).predict(X_test)accuracy = accuracy_score(y_test, predictions) accuracy_list.append(accuracy)
Ta có thể lưu kết quả vào DataFrame để xem:
results_df = pd.DataFrame({'neighbors':neighbors_list, 'accuracy':accuracy_list})
print(results_df)

Hãy tạo đồ thị đường học tập
Lần này sẽ thử nhiều giá trị hơn
neighbors_list = list(range(5,500, 5))accuracy_list = [] for test_number in neighbors_list: model = KNeighborsClassifier(n_neighbors=test_number) predictions = model.fit(X_train, y_train).predict(X_test) accuracy = accuracy_score(y_test, predictions) accuracy_list.append(accuracy) results_df = pd.DataFrame({'neighbors':neighbors_list, 'accuracy':accuracy_list})
Ta có thể vẽ DataFrame lớn hơn:
plt.plot(results_df['neighbors'], results_df['accuracy'])# Thêm nhãn và tiêu đề plt.gca().set(xlabel='n_neighbors', ylabel='Accuracy', title='Accuracy for different n_neighbors') plt.show()
Đồ thị của chúng ta:

Hàm range của Python không hỗ trợ bước thập phân.
Một mẹo tiện dùng NumPy np.linspace(start, end, num)
num giá trị cách đều trong khoảng (start, end) bạn chỉ định.print(np.linspace(1,2,5))
[1. 1.25 1.5 1.75 2. ]
Tinh chỉnh siêu tham số trong Python