Візуалізація окремих дерев XGBoost
Тепер, коли ви використали XGBoost для побудови й оцінювання моделей регресії та класифікації, варто навчитися візуально досліджувати свої моделі. Тут ви візуалізуєте окремі дерева з повністю підсиленої моделі, яку XGBoost створює, використовуючи весь набір даних про житло.
У XGBoost є функція plot_tree(), яка полегшує таку візуалізацію. Після того як ви натренуєте модель за допомогою API навчання XGBoost, ви можете передати її у функцію plot_tree() разом із кількістю дерев, які хочете відобразити, через аргумент num_trees.
Ця вправа є частиною курсу
Екстремальний градієнтний бустинг з XGBoost
Інструкції до вправи
- Створіть словник параметрів із
"objective"="reg:squarederror"та"max_depth"=2. - Натренуйте модель, використавши
10раундів підсилення та створений вами словник параметрів. Збережіть результат уxg_reg. - Побудуйте графік першого дерева за допомогою
xgb.plot_tree(). Функція приймає два аргументи: модель (у цьому випадкуxg_reg) іnum_trees, індексація якого починається з 0. Тож для першого дерева вкажіть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()