Skip to content
AI-grafen
DAI developerClassical machine learning· about 45 min· fundamentals that rarely change· verified 2026-09-20· EN

Decision trees

Be able to build a decision tree, interpret it and explain overfitting in deep trees.

Prerequisites

Intuition

A decision tree is a series of yes/no questions.

            hours > 10?
           /            \
         no             yes
          |              |
     attendance > 0.8?  PASS
      /        \
    no         yes
     |          |
   FAIL        PASS

Training means finding the questions. The algorithm tries every column and every cut point, and chooses the split that makes the two groups the most «pure» — that is, the most one-sided in their answers.

Then it is repeated in each branch, until something says stop.

Why trees are popular:

  • The interface is understandable — you can read the tree and explain a decision.
  • No scaling is needed (the tree splits on thresholds).
  • It mixes numeric and categorical variables without preprocessing.
  • It catches non-linear relationships and interactions automatically.

And the big problem: an unrestricted tree keeps splitting until every leaf contains a single data point. Then the accuracy on the training data is 100 % and on new data it is dreadful. The tree has memorised.

Formal

How a split is chosen. Two measures of impurity, both doing the same thing in practice:

Gini=1−∑kpk2,Entropy=−∑kpklog⁡2pk\text{Gini} = 1 - \sum_k p_k^2, \qquad \text{Entropy} = -\sum_k p_k \log_2 p_k

Both are 0 when the node is completely pure and largest when the classes are evenly distributed. For two classes: max Gini = 0.5, max entropy = 1 bit.

The algorithm chooses the split that gives the greatest information gain:

Δ=I(parent)−nlnI(left)−nrnI(right)\Delta = I(\text{parent}) - \frac{n_l}{n}I(\text{left}) - \frac{n_r}{n}I(\text{right})

The parts are weighted by how many examples end up in them — a split that breaks three points out of a thousand counts for little.

A worked example. 100 examples, 50 of each class → Gini = 1 − (0.5² + 0.5²) = 0.5. A split gives [40 A, 10 B] and [10 A, 40 B]:

  • Left: 1 − (0.8² + 0.2²) = 0.32
  • Right: the same, 0.32
  • Weighted: 0.5·0.32 + 0.5·0.32 = 0.32
  • Gain: 0.18

Overfitting and how it is limited:

ParameterDoesA reasonable value
max_depththe maximum number of questions in a row3–10
min_samples_leafthe smallest number of examples in a leaf5–50
min_samples_splitthe smallest number needed to split at all10+
ccp_alphaprunes the tree afterwardssearch with cross-validation

The trees' other weakness is instability: swap a few data points out and the tree can look entirely different. That makes them unreliable as explanations of the phenomenon — even though they explain the model's decision well.

Both problems are solved by the same idea: build many trees and let them vote. A random forest trains each tree on a sample of the data and a selection of the columns; gradient boosting builds trees that correct the previous tree's errors. The result is often the best there is for tabular data — but the explainability of the single tree is gone.

Code

import numpy as np
from sklearn.datasets import load_breast_cancer
from sklearn.model_selection import train_test_split
from sklearn.tree import DecisionTreeClassifier, export_text

X, y = load_breast_cancer(return_X_y=True, as_frame=True)
Xtr, Xte, ytr, yte = train_test_split(X, y, test_size=0.3, stratify=y, random_state=0)

for depth in (1, 3, 5, None):
    t = DecisionTreeClassifier(max_depth=depth, random_state=0).fit(Xtr, ytr)
    print(f"depth={str(depth):>4}  leaves={t.get_n_leaves():>3}  "
          f"training {t.score(Xtr, ytr):.3f}  test {t.score(Xte, yte):.3f}")
# depth=   1  leaves=  2  training 0.925  test 0.883
# depth=   3  leaves=  8  training 0.985  test 0.942
# depth=   5  leaves= 14  training 1.000  test 0.930
# depth=None  leaves= 19  training 1.000  test 0.930    ← memorised the training data

# A small tree can be read
t = DecisionTreeClassifier(max_depth=2, random_state=0).fit(Xtr, ytr)
print(export_text(t, feature_names=list(X.columns), max_depth=2))

# Gini by hand
def gini(*counts):
    n = sum(counts)
    return 1 - sum((a / n) ** 2 for a in counts)

before = gini(50, 50)
after = 0.5 * gini(40, 10) + 0.5 * gini(10, 40)
print(round(before, 3), round(after, 3), "gain", round(before - after, 3))   # 0.5 0.32 gain 0.18

# Instability: remove three rows and see how the tree changes
for start in (0, 3):
    t = DecisionTreeClassifier(max_depth=2, random_state=0).fit(Xtr.iloc[start:], ytr.iloc[start:])
    print("root split:", X.columns[t.tree_.feature[0]], round(float(t.tree_.threshold[0]), 3))

Mastery means

  • Builds and interprets a decision tree
  • Explains how a split is chosen
  • Recognises and limits overfitting

Sign in to do the exercises and build your mastery up.

Sources

All the sources and licences