Создайте пайплайн для подготовки бинарных признаков. Он должен состоять из двух шагов:
- Обработка пропусков: пропущенные значения сначала заполняются nan, а затем самым частотным значением ('most_frequent').
- OHE-кодирование. При этом нужно учесть, что данные могут содержать пропуски.
import numpy as np
import pandas as pd
from sklearn.model_selection import train_test_split
# загружаем нужные классы
from sklearn.pipeline import Pipeline
from sklearn.compose import ColumnTransformer
from sklearn.preprocessing import OneHotEncoder, OrdinalEncoder, StandardScaler
# импортируйте нужные классы для работы с пропусками
from sklearn.impute import SimpleImputer
# загружаем нужные метрики
from sklearn.metrics import roc_auc_score
# импортируем модель
from sklearn.tree import DecisionTreeClassifier
RANDOM_STATE = 42
TEST_SIZE = 0.25
# загружаем данные
df_full = pd.read_csv('railway_full.csv')
X_train, X_test, y_train, y_test = train_test_split(
df_full.drop(
[
'Удовлетворён предоставленной услугой',
'Общая оценка качества предоставленной услуги'
],
axis=1
),
df_full['Удовлетворён предоставленной услугой'],
test_size = TEST_SIZE,
random_state = RANDOM_STATE,
stratify = df_full['Удовлетворён предоставленной услугой']
)
# создаём списки с названиями признаков
ohe_columns = [
'Пол', 'Путешествует с детьми', 'Путешествует по работе',
'Тип', 'Оценка качества питания'
]
ord_columns = [
'Оценка комфортности покупки билета онлайн', 'Оценка качества wifi',
'Оценка комфортности времени отправления/прибытия'
]
num_columns = ['Возраст', 'Расстояние']
# Создайте пайплайн для подготовки признаков из списка ohe_columns:
# 1) заполните пропуски,
# 2) проведите OHE-кодирование.
ohe_pipe = Pipeline(
[
('simpleImputer_ohe', SimpleImputer(missing_values=np.nan, strategy='most_frequent')),
('ohe', OneHotEncoder(drop='first', handle_unknown='ignore', sparse=False))
]
)
print(ohe_pipe)
Результат
Pipeline(steps=[('simpleImputer_ohe', SimpleImputer(strategy='most_frequent')),
('ohe',
OneHotEncoder(drop='first', handle_unknown='ignore',
sparse=False))])