Visualiser des arbres XGBoost individuels
Maintenant que vous avez utilisé XGBoost pour construire et évaluer des modèles de régression ainsi que de classification, prenez le temps d'explorer vos modèles de façon visuelle. Ici, vous allez visualiser des arbres individuels du modèle entièrement renforcé (boosted) qu'XGBoost crée en utilisant l'ensemble complet des données sur les maisons.
XGBoost propose une fonction plot_tree() qui facilite ce type de visualisation. Une fois que vous avez entraîné un modèle à l'aide de l'API d'apprentissage de XGBoost, vous pouvez le passer à la fonction plot_tree() en précisant le nombre d'arbres à tracer avec l'argument num_trees.
Cette activité fait partie du cours
Amorçage de gradient avancé avec XGBoost
Instructions de l’exercice
- Créez un dictionnaire de paramètres avec un
"objective"de"reg:squarederror"et un"max_depth"de2. - Entraînez le modèle avec
10itérations de boosting en utilisant le dictionnaire de paramètres que vous avez créé. Enregistrez le résultat dansxg_reg. - Tracez le premier arbre avec
xgb.plot_tree(). Cette fonction prend deux arguments : le modèle (ici,xg_reg) etnum_trees, qui est indexé à partir de 0. Donc, pour tracer le premier arbre, indiqueznum_trees=0. - Tracez le cinquième arbre.
- Tracez le dernier (le dixième) arbre à l'horizontale. Pour ce faire, indiquez l'argument de mot-clé additionnel
rankdir="LR".
Exercice interactif pratique
Essayez cet exercice en complétant ce code d’exemple.
# 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()