CommencezCommencez gratuitement

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

Voir le cours

Instructions de l’exercice

  • Créez un dictionnaire de paramètres avec un "objective" de "reg:squarederror" et un "max_depth" de 2.
  • Entraînez le modèle avec 10 itérations de boosting en utilisant le dictionnaire de paramètres que vous avez créé. Enregistrez le résultat dans xg_reg.
  • Tracez le premier arbre avec xgb.plot_tree(). Cette fonction prend deux arguments : le modèle (ici, xg_reg) et num_trees, qui est indexé à partir de 0. Donc, pour tracer le premier arbre, indiquez num_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()
Modifier et exécuter le code