Wizualizacja pojedynczych drzew XGBoost
Skoro udało ci się już użyć XGBoost do budowania i oceny modeli regresji oraz klasyfikacji, czas przyjrzeć się modelom od środka. W tym ćwiczeniu zwizualizujesz poszczególne drzewa z w pełni wytrenowanego modelu XGBoost, zbudowanego na całym zbiorze danych dotyczących cen domów.
XGBoost udostępnia funkcję plot_tree(), która znacznie ułatwia tego typu wizualizację. Po wytrenowaniu modelu za pomocą API uczenia XGBoost możesz przekazać go do funkcji plot_tree() wraz z liczbą drzew do wykreślenia – za pomocą argumentu num_trees.
To ćwiczenie jest częścią kursu
Extreme Gradient Boosting with XGBoost
Instrukcje do ćwiczenia
- Utwórz słownik parametrów z kluczem
"objective"ustawionym na"reg:squarederror"oraz"max_depth"równym2. - Wytrenuj model, używając
10rund boostingu i utworzonego słownika parametrów. Zapisz wynik w zmiennejxg_reg. - Wykreśl pierwsze drzewo za pomocą
xgb.plot_tree(). Funkcja przyjmuje dwa argumenty: model (w tym przypadkuxg_reg) oraznum_trees– indeksowany od 0. Aby wykreślić pierwsze drzewo, podajnum_trees=0. - Wykreśl piąte drzewo.
- Wykreśl ostatnie (dziesiąte) drzewo w układzie poziomym. W tym celu dodaj argument
rankdir="LR".
Interaktywne ćwiczenie praktyczne
Spróbuj tego ćwiczenia, uzupełniając ten przykładowy kod.
# 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()