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:
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:
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:
| Parameter | Does | A reasonable value |
|---|---|---|
max_depth | the maximum number of questions in a row | 3–10 |
min_samples_leaf | the smallest number of examples in a leaf | 5–50 |
min_samples_split | the smallest number needed to split at all | 10+ |
ccp_alpha | prunes the tree afterwards | search 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
- scikit-learn User Guide (BSD-3) — BSD-3-Clause
- Dive into Deep Learning (CC BY-SA 4.0) — CC BY-SA 4.0
- Wikipedia — Decision tree learning (CC BY-SA 4.0) — CC BY-SA 4.0