Улучшите значение метрики FPR на данных о пассажирских перевозках.
- Создайте словарь с метриками ROC-AUC и FPR.
- Запустите автоматизированный поиск с мультискорингом. В конце обучите модель заново на метрике FPR.
- Проверьте значение метрики FPR на тестовой выборке, используя лучшую модель из
GridSearchCV.
# импортируем библиотеки и объявляем константы
import pandas as pd
from sklearn.metrics import confusion_matrix, make_scorer
from sklearn.model_selection import train_test_split, GridSearchCV
from sklearn.preprocessing import OneHotEncoder
from sklearn.tree import DecisionTreeClassifier
RANDOM_STATE = 42
# подготавливаем данные заранее созданной функцией
X_train, X_test, y_train, y_test = prepare_data('train_satisfaction.csv')
# инициализируем модель дерева решений
model = DecisionTreeClassifier(random_state=RANDOM_STATE)
# создаём словарь с гиперпараметрами
parameters = {
'min_samples_split': range(2, 6),
'min_samples_leaf': range(1, 6),
'max_depth': range(2, 6)
}
# создаём функцию для подсчёта метрики FPR
def get_false_positive_rate(y_true, y_pred):
conf_matrix = confusion_matrix(y_true, y_pred)
return conf_matrix[0][1] / (conf_matrix[0][1] + conf_matrix[0][0])
# считаем метрику FPR вместе с make_scorer()
fpr_score = make_scorer(
get_false_positive_rate,
greater_is_better=False
)
# Создайте словарь с двумя метриками:
# - roc-auc с именем roc_auc_score;
# - FPR с именем false_positive_rate.
scoring = scoring = {
'roc_auc_score': 'roc_auc',
'false_positive_rate': fpr_score
}
# Инициализируйте класс для автоматизированного поиска:
# значение кросс-валидации 5, метрики из словаря,
# в конце обучить на метрике FPR.
gs_multiscoring = GridSearchCV(
model,
parameters,
n_jobs=-1,
cv=5,
scoring=scoring,
refit='false_positive_rate'
)
# запускаем поиск гиперпараметров с мультискорингом
gs_multiscoring.fit(X_train, y_train)
# получите предсказания на тестовых данных с помощью лучшей модели
y_pred = gs_multiscoring.best_estimator_.predict(X_test)
# считаем метрику FPR на тестовых данных
print(get_false_positive_rate(y_test, y_pred))
Результат
0.6553191489361702