Kom igångKom igång gratis

Visualisera enskilda XGBoost-träd

Nu när du har använt XGBoost för att både bygga och utvärdera regressions- och klassificeringsmodeller är det dags att utforska hur du kan visualisera dina modeller. Här visualiserar du enskilda träd från den färdigboostade modell som XGBoost skapar med hela bostadsdatamängden.

XGBoost har en plot_tree()-funktion som gör den här typen av visualisering enkel. När du har tränat en modell med XGBoosts inlärnings-API kan du skicka den till plot_tree() tillsammans med antalet träd du vill rita, via argumentet num_trees.

Den här övningen är en del av kursen

Extreme Gradient Boosting med XGBoost

Visa kurs

Övningsinstruktioner

  • Skapa en parameterordbok med "objective" satt till "reg:squarederror" och "max_depth" satt till 2.
  • Träna modellen med 10 boostringsrundor och parameterordboken du skapade. Spara resultatet i xg_reg.
  • Rita det första trädet med xgb.plot_tree(). Funktionen tar två argument – modellen (i det här fallet xg_reg) och num_trees, som är 0-indexerat. Ange alltså num_trees=0 för att rita det första trädet.
  • Rita det femte trädet.
  • Rita det sista (tionde) trädet liggande. Gör det genom att ange det extra nyckelordsargumentet rankdir="LR".

Interaktiv övning med praktiskt arbete

Testa den här övningen genom att slutföra den här exempelkoden.

# 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()
Redigera och kör kod