Machine Learningscikit-learn 1.9.1 · xgboost 3.4.1 · Python 3.12+
Dashboard
0%
1
Curious builder0 XP earned · 300 to level 2
0 daysFinish a lesson to begin
Badge collection0 of 6 unlocked
52 small wins to finish your pathNext lesson →

Decision tree in scikit-learn

DecisionTreeClassifier is the scikit-learn class that builds a classification tree, and plot_tree draws the fitted tree with each node's question, Gini impurity, sample count and class counts.

Last updated: 04 Oct, 2026 · scikit-learn 1.9.1

Decision tree regression and pruning ended the hand-worked part of trees. The video's notebook fits a tree on the Iris flowers and draws it, so the questions, Gini values and class counts from the earlier tree lessons all show up on one picture.

DecisionTreeClassifier and its hyperparameters · from the Complete Machine Learning in 6 Hours video · 244:35 to 245:45

Reading the DecisionTreeClassifier parameters

The video plans to let the tree overfit first and prune it afterwards, so it keeps the defaults and walks through the parameters:

  • criterion: the impurity measure, "gini" by default; "entropy" and "log_loss" use entropy.
  • splitter: "best" tries every candidate split; "random" picks among random ones. The video recommends best.
  • max_depth, min_samples_leaf and max_features: hyperparameters that limit how far the tree grows and how many features each split looks at.

The dataset is Iris: 150 flowers, 50 of each species (setosa, versicolor, virginica), with four measurements each: sepal length, sepal width, petal length and petal width.

Fitting and drawing the Iris tree

Loading Iris

python
import matplotlib.pyplot as plt
from sklearn import tree
from sklearn.datasets import load_iris
from sklearn.tree import DecisionTreeClassifier

iris = load_iris()           # 150 flowers, 4 measurements, 3 species

Fitting the classifier

python
# Default settings: criterion="gini", splitter="best", no depth limit
classifier = DecisionTreeClassifier(random_state=0)
classifier.fit(iris.data, iris.target)

The video's first attempt calls tree.plot, which does not exist, then plots before fitting and gets a NotFittedError. The working order is fit first, then tree.plot_tree. The video's call passes only filled=True, so its nodes read X[0] to X[3]; feature_names and class_names are added here, and random_state=0 makes the tie-breaks repeat.

ExampleFrom the video, run on scikit-learn 1.9.1
plt.figure(figsize=(15, 10))
tree.plot_tree(classifier, filled=True, feature_names=iris.feature_names,
               class_names=list(iris.target_names))
plt.title("Decision tree on the Iris dataset")
plt.show()
print("depth:", classifier.get_depth(), " leaves:", classifier.get_n_leaves())
print("root gini:", round(classifier.tree_.impurity[0], 3))
The fitted Iris decision tree: the root asks petal width <= 0.8 with gini 0.667 on 150 samples, the left leaf holds the 50 setosa flowers, and the right side splits the 100 versicolor and virginica flowers further.
Reading the Iris tree · from the Complete Machine Learning in 6 Hours video · 247:05 to 248:21

Reading the Iris tree

  • The root asks petal width ≤ 0.8 on all 150 samples, value [50, 50, 50], gini 0.667.
  • The left branch is a leaf: 50 samples, value [50, 0, 0], gini 0.0. One question separates all the setosa flowers.
  • The right branch holds the other 100 flowers, [0, 50, 50], gini 0.5, and splits on petal width ≤ 1.75 into [0, 49, 5] and [0, 1, 45].
  • Deep in the tree a node holds [0, 47, 1]. The video asks whether such a split is worth making at all: one flower out of 48 is a lot of tree for little gain. That is the post-pruning question.

The clip names the left leaf versicolor; value [50, 0, 0] is class 0, setosa. The root's gini of 0.667 is above the 0.5 ceiling from Entropy and Gini impurity because Iris has three classes: 1 − 3 × (1/3)² = 0.667. The 0.5 ceiling holds for two classes, which is why the node with only versicolor and virginica shows exactly 0.5.

Testing the tree with a confusion matrix

The video fits the tree on all 150 flowers and never tests it. Holding out 30 flowers (20%) as a test set shows how the full tree does on flowers it never saw. A Confusion matrix and the classification report then show which species it mixes up.

Holding out a test set

python
from sklearn.model_selection import train_test_split
from sklearn.metrics import confusion_matrix, classification_report, accuracy_score

# 120 flowers to train on, 30 to test on
X_train, X_test, y_train, y_test = train_test_split(iris.data, iris.target, test_size=0.2, random_state=10)
tree_classifier = DecisionTreeClassifier(random_state=0).fit(X_train, y_train)
ExampleRun on scikit-learn 1.9.1
y_pred = tree_classifier.predict(X_test)
print("depth:", tree_classifier.get_depth(), " leaves:", tree_classifier.get_n_leaves(),
      " train accuracy:", tree_classifier.score(X_train, y_train))
print(confusion_matrix(y_test, y_pred))
print(classification_report(y_test, y_pred, target_names=iris.target_names))

Pre-pruning the Iris tree with GridSearchCV

Pre-pruning limits the tree while it grows, as in Decision tree regression and pruning. GridSearchCV tries every mix of criterion, splitter, depth and max_features with 5-fold cross-validation on the 120 training flowers, and keeps the best. Older code for this grid lists max_features="auto", which scikit-learn removed in 1.3; None (all features) stands in for it.

python
from sklearn.model_selection import GridSearchCV

param = {
    "criterion": ["gini", "entropy", "log_loss"],
    "splitter": ["best", "random"],
    "max_depth": [1, 2, 3, 4, 5],
    "max_features": [None, "sqrt", "log2"],
}
grid = GridSearchCV(DecisionTreeClassifier(random_state=0), param_grid=param, cv=5, scoring="accuracy")
grid.fit(X_train, y_train)
ExampleRun on scikit-learn 1.9.1
print("best params:", grid.best_params_)
print("best cv accuracy:", round(grid.best_score_, 3))
best = grid.best_estimator_
print("depth:", best.get_depth(), " leaves:", best.get_n_leaves())
y_pred = grid.predict(X_test)
print(confusion_matrix(y_test, y_pred))
print("test accuracy:", round(accuracy_score(y_test, y_pred), 3))

What the test and the grid search show

  • The full tree has depth 6 and 10 leaves and a training accuracy of 1.0: every training flower is classified right, a sign of overfitting.
  • On the 30 test flowers it makes one mistake: the confusion matrix shows one versicolor predicted as virginica, so versicolor recall is 12/13 = 0.92 and virginica precision 7/8 = 0.88, for an accuracy of 0.97.
  • The grid search picks gini, max_depth 3, the random splitter and all features, with a cross-validation accuracy of 0.967 on the training flowers. The tree it keeps has 5 leaves, half as many.
  • The grid-searched tree gets all 30 test flowers right. On 30 flowers one mistake is 0.033 of accuracy, so the cross-validation score, which averages five folds, is the steadier guide; the test set checks the choice once.

Full tree vs grid-searched tree

Full tree (defaults)Grid-searched tree
Depth and leaves6 and 103 and 5
Chosen byThe defaults: grow until every leaf is pure5-fold cross-validation over 90 settings
Training accuracy1.0Below 1.0 is allowed
Test accuracy (30 flowers)0.97, one versicolor called virginica1.0
Readable with plot_treeYes, but crowdedYes, at a glance
RiskOverfittingUnderfitting if the depth is too small

Where you use plot_tree

  • Explaining a model: the drawing shows every question the model asks and how many training rows reach each node.
  • Spotting overfitting: leaves with 1 or 2 samples, such as [0, 0, 1], are splits made for single rows.
  • Checking a pruning choice: draw the tree before and after setting max_depth or ccp_alpha.
Watch out. Without random_state, two runs can choose different splits when features tie, so the drawing and the scores can change between runs. A deep tree also turns into an unreadable plot: pass max_depth=2 to plot_tree to draw only the top, or print the rules with export_text.
Try it yourself
  • Fit the Iris tree with criterion="entropy" and plot it. What does the root's impurity read now (log₂ 3 for three equal classes)?
  • Draw only the top of the tree with tree.plot_tree(classifier, max_depth=2, filled=True).
  • Print the rules as text with print(export_text(classifier, feature_names=iris.feature_names)) (import export_text from sklearn.tree).

Every expert started right here.