Ajustement des hyperparamètres en Python
Alex Scriven
Data Scientist
Certains hyperparamètres sont prioritaires à régler d'abord.
Mais quelles valeurs essayer pour les hyperparamètres ?
Voyons quelques astuces clés !
Attention aux choix d'hyperparamètres incompatibles.
LogisticRegression() a des options solver et penalty qui peuvent entrer en conflit.The 'newton-cg', 'sag' and 'lbfgs' solvers support only l2 penalties.
Certains conflits ne sont pas explicites : ils seront simplement « ignorés » (dans ElasticNet avec l'hyperparamètre normalize) :
This parameter is ignored when fit_intercept is set to False
Consultez la documentation de Scikit-Learn !
Évitez de fixer des valeurs « insensées » selon l'algorithme :
Prendre le temps de documenter des valeurs raisonnables d'hyperparamètres est très utile.
Dans l'exercice précédent, nous avons bâti des modèles ainsi :
knn_5 = KNeighborsClassifier(n_neighbors=5)
knn_10 = KNeighborsClassifier(n_neighbors=10)
knn_20 = KNeighborsClassifier(n_neighbors=20)
C'est peu efficace. Peut-on faire mieux ?
Essayez une boucle for pour parcourir les options :
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)
On peut stocker les résultats dans un DataFrame pour les consulter :
results_df = pd.DataFrame({'neighbors':neighbors_list, 'accuracy':accuracy_list})
print(results_df)

Créons une courbe d'apprentissage
Cette fois, nous testerons bien plus de valeurs
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})
Traçons le DataFrame plus volumineux :
plt.plot(results_df['neighbors'], results_df['accuracy'])# Ajouter les étiquettes et le titre plt.gca().set(xlabel='n_neighbors', ylabel='Accuracy', title='Accuracy for different n_neighbors') plt.show()
Notre graphique :

La fonction range de Python ne gère pas les pas décimaux.
Astuces : utilisez np.linspace(start, end, num) de NumPy
num valeurs également espacées dans l'intervalle (start, end) que vous précisez.print(np.linspace(1,2,5))
[1. 1.25 1.5 1.75 2. ]
Ajustement des hyperparamètres en Python