Визуализация отдельных деревьев XGBoost
Теперь, когда вы научились использовать XGBoost для построения и оценки как регрессионных, так и классификационных моделей, самое время разобраться, как визуально исследовать полученные модели. В этом упражнении вы визуализируете отдельные деревья из полностью обученной модели XGBoost, построенной на полном наборе данных о жилье.
XGBoost предоставляет функцию plot_tree(), которая упрощает такую визуализацию. После обучения модели с помощью API обучения XGBoost её можно передать в функцию plot_tree() вместе с количеством деревьев для отображения через аргумент num_trees.
Это упражнение является частью курса
Экстремальный градиентный бустинг с XGBoost
Инструкции к упражнению
- Создайте словарь параметров со значением
"reg:squarederror"для ключа"objective"и значением2для ключа"max_depth". - Обучите модель, используя
10раундов бустинга и созданный словарь параметров. Сохраните результат в переменнуюxg_reg. - Постройте первое дерево с помощью
xgb.plot_tree(). Функция принимает два аргумента: модель (в данном случаеxg_reg) иnum_trees— индекс дерева, начинающийся с нуля. Чтобы отобразить первое дерево, укажитеnum_trees=0. - Постройте пятое дерево.
- Постройте последнее (десятое) дерево в горизонтальной ориентации. Для этого добавьте дополнительный именованный аргумент
rankdir="LR".
Интерактивное практическое упражнение
Попробуйте выполнить это упражнение, дополнив этот пример кода.
# Create the DMatrix: housing_dmatrix
housing_dmatrix = xgb.DMatrix(data=X, label=y)
# Create the parameter dictionary: params
params = {"objective":"reg:squarederror", "max_depth":2}
# Train the model: xg_reg
xg_reg = xgb.train(params=params, dtrain=housing_dmatrix, num_boost_round=10)
# Plot the first tree
____
plt.show()
# Plot the fifth tree
____
plt.show()
# Plot the last tree sideways
____
plt.show()