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
Övningsinstruktioner
- Skapa en parameterordbok med
"objective"satt till"reg:squarederror"och"max_depth"satt till2. - Träna modellen med
10boostringsrundor och parameterordboken du skapade. Spara resultatet ixg_reg. - Rita det första trädet med
xgb.plot_tree(). Funktionen tar två argument – modellen (i det här falletxg_reg) ochnum_trees, som är 0-indexerat. Ange alltsånum_trees=0fö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()