Ajustarea hiperparametrilor în Python
Alex Scriven
Data Scientist
Unii hiperparametri sunt mai importanți decât alții pentru ajustare inițială.
Dar ce valori să încercăm?
Hai să vedem principalele sfaturi!
Atenție la combinații conflictuale de hiperparametri.
LogisticRegression() are opțiuni conflictuale pentru solver și penalty.The 'newton-cg', 'sag' and 'lbfgs' solvers support only l2 penalties.
Unii nu sunt expliciți, ci sunt pur și simplu ignorați (ex. ElasticNet cu hiperparametrul normalize):
This parameter is ignored when fit_intercept is set to False
Consultați întotdeauna documentația Scikit-Learn!
Evitați valorile „absurde" pentru diferiți algoritmi:
Documentarea valorilor rezonabile pentru hiperparametri este o activitate valoroasă.
În exercițiul anterior, am construit modele astfel:
knn_5 = KNeighborsClassifier(n_neighbors=5)
knn_10 = KNeighborsClassifier(n_neighbors=10)
knn_20 = KNeighborsClassifier(n_neighbors=20)
Acesta este ineficient. Putem face mai bine?
Utilizați o buclă for pentru a itera prin opțiuni:
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)
Putem stoca rezultatele într-un DataFrame:
results_df = pd.DataFrame({'neighbors':neighbors_list, 'accuracy':accuracy_list})
print(results_df)

Vom crea un grafic al curbei de învățare
De data aceasta testăm mai multe valori
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})
Putem reprezenta grafic DataFrame-ul extins:
plt.plot(results_df['neighbors'], results_df['accuracy'])# Add the labels and title plt.gca().set(xlabel='n_neighbors', ylabel='Accuracy', title='Accuracy for different n_neighbors') plt.show()
Graficul obținut:

Funcția range din Python nu acceptă pași zecimali.
O metodă utilă este np.linspace(start, end, num) din NumPy
num valori distribuite uniform în intervalul (start, end) specificat.print(np.linspace(1,2,5))
[1. 1.25 1.5 1.75 2. ]
Ajustarea hiperparametrilor în Python