Zacznij terazZacznij za darmo

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

Zobacz kurs

Instrukcje do ćwiczenia

  • Utwórz słownik parametrów z kluczem "objective" ustawionym na "reg:squarederror" oraz "max_depth" równym 2.
  • Wytrenuj model, używając 10 rund boostingu i utworzonego słownika parametrów. Zapisz wynik w zmiennej xg_reg.
  • Wykreśl pierwsze drzewo za pomocą xgb.plot_tree(). Funkcja przyjmuje dwa argumenty: model (w tym przypadku xg_reg) oraz num_trees – indeksowany od 0. Aby wykreślić pierwsze drzewo, podaj num_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()
Edytuj i uruchom kod