Decision trees
A decision tree classifies by asking yes-or-no questions about the inputs. Learn how each question is chosen with Gini impurity or entropy, how a tree is grown, why a deep tree overfits, how trees predict numbers, and how to read which inputs a tree relies on.
Warm-up
One question before the lesson. Choose an answer and check it.
One group of 10 students has 5 passes and 5 fails. Another has 10 passes and no fails. Which group is easier to predict?
Show the answer
A: The 10 passes and no fails: predict pass and be right every time. A group where everyone agrees can be predicted perfectly, and a 50–50 group is as mixed as a group can be. A decision tree asks the questions that split students into groups that agree, and measures how mixed a group is with Gini impurity.
Step 1 A tree of questions
A decision tree classifies by asking yes-or-no questions about the inputs, one after another. Each answer leads to the next question, until an answer leads to a prediction. Here is a small tree grown on Lesson 6’s 20 students:
Did they study less than 4.8 hours?
yes: did they sleep less than 6.8 hours?
yes: fail (0 of 7 passed)
no: pass (2 of 3 passed)
no: did they study less than 6.65 hours?
yes: pass (3 of 4 passed)
no: pass (6 of 6 passed)
Each question splits the students in two by one input: less than a threshold, or not. On the board each question is a straight line across or up the plot, so the tree carves the plot into boxes, called leaves. Each leaf predicts what most of its training students did.
The new student from Lesson 6 studied 5.2 hours and slept 6:
5.2 is not less than 4.8, so no
5.2 is less than 6.65, so yes
prediction: pass
A tree is easy to read: anyone can follow its questions. Training it means choosing the questions, and the next steps show how.
Try it yourself
Pick an input and drag the threshold to split the students, and watch the Gini impurity of each side and the gain. Find the best split, then grow the tree a level at a time and compare its accuracy on the training and validation students.
Filled dots passed, hollow dots failed. Students below the threshold answer yes.
Gain 0.126
yes (6): 1 of 6 passed, Gini = 1 − (1/6)² − (5/6)² = 10/36 = 0.2778
no (14): 10 of 14 passed, Gini = 1 − (10/14)² − (4/14)² = 80/196 = 0.4082
impurity = 6/20 × 10/36 + 14/20 × 80/196 = 0.369
gain = 0.495 − 0.369 = 0.126
The best split there is gains 0.245: study < 4.8.
Show the calculation
| Student | Study | Sleep | Result | study < 2.50? |
|---|---|---|---|---|
| 1 | 1.4 | 4.2 | failed | yes |
| 2 | 2.6 | 5.7 | failed | no |
| 3 | 1.6 | 8.7 | failed | yes |
| 4 | 3.6 | 6.4 | failed | no |
| 5 | 6.9 | 4.1 | passed | no |
| 6 | 0.8 | 5.9 | failed | yes |
| 7 | 6 | 8.5 | passed | no |
| 8 | 7.1 | 8.1 | passed | no |
| 9 | 6.1 | 5.5 | passed | no |
| 10 | 1.9 | 4.1 | failed | yes |
| 11 | 6.2 | 7.6 | passed | no |
| 12 | 3.1 | 7.2 | passed | no |
| 13 | 2.8 | 4 | failed | no |
| 14 | 8.2 | 8.3 | passed | no |
| 15 | 8.7 | 4.6 | passed | no |
| 16 | 8.6 | 5.7 | passed | no |
| 17 | 7.6 | 8.3 | passed | no |
| 18 | 6.4 | 6.1 | failed | no |
| 19 | 1.9 | 7.9 | passed | yes |
| 20 | 1.4 | 4.7 | failed | yes |
Move the threshold and find the largest gain, then press Best split.
Practice problems
Work each problem out on paper, then type your answer and press Check. Every problem has hints and a full solution.
Score: 0 of 10 points
Problem 1
1 pointA group has 3 passes and 1 fail. What is its Gini impurity?
Hint 1
p is the share who passed. Gini = 1 − p² − (1 − p)².
Solution
p = 3 ÷ 4 = 0.75
Gini = 1 − (3/4)² − (1/4)² = 1 − 9/16 − 1/16 = 6/16 = 0.375
Problem 2
2 pointsA question splits 10 students into a left side with 4 passes and 0 fails, and a right side with 2 passes and 4 fails. What is the split’s weighted Gini impurity? Give it to 2 decimal places.
Hint 1
Work out each side’s Gini.
Hint 2
Weight each by its share of the 10 students, then add.
Solution
left: 4 students, Gini = 1 − (4/4)² − (0/4)² = 1 − 16/16 − 0/16 = 0/16 = 0
right: 6 students, Gini = 1 − (2/6)² − (4/6)² = 1 − 4/36 − 16/36 = 16/36 = 0.4444
impurity = 4/10 × 0/16 + 6/10 × 16/36 = 0.2667
Problem 3
2 pointsAnd the split’s gain, the parent group’s Gini minus the split’s weighted impurity? Give it to 2 decimal places.
Hint 1
The parent group is all 10 students: 6 passes and 4 fails.
Solution
parent: Gini = 1 − (6/10)² − (4/10)² = 1 − 36/100 − 16/100 = 48/100 = 0.48
gain = 0.48 − 0.2667 = 0.2133
Problem 4
2 pointsWhat is the entropy of a group with 3 passes and 1 fail, in bits? Give it to 2 decimal places.
Hint 1
H = −p log₂ p − (1 − p) log₂(1 − p).
Hint 2
log₂ x is ln x ÷ ln 2.
Solution
p = 0.75
H = −0.75 × log₂ 0.75 − 0.25 × log₂ 0.25 = 0.3113 + 0.5 = 0.8113
Problem 5
3 pointsA regression tree’s leaf holds three training students, who scored 50, 54 and 58. The leaf predicts their mean. What is the leaf’s sum of squared errors (SSE)?
Hint 1
The leaf predicts the mean of the three scores.
Hint 2
Add up each score’s squared distance from that mean.
Solution
mean = (50 + 54 + 58) ÷ 3 = 54
SSE = (−4)² + 0² + 4² = 16 + 0 + 16 = 32
Programming exercise
Write Gini impurity, entropy, the best split and a whole decision tree in plain Python. Save tree.py and test_tree.py in the same folder, fill in each function in tree.py, and run the tests:
python test_tree.py
"""Decision trees: programming exercise. Write a decision tree classifier in plain Python: Gini impurity and entropy,splitting a group, finding the best question, and growing the whole tree.Run the tests from this folder: python test_tree.py A point is a list of numbers, such as [hours studied, hours slept], andfeature is a position in it: 0 for the first number, 1 for the second. Alabel is 1 or 0. A group's labels are a list of 1s and 0s.""" def gini(labels): """Return the Gini impurity of the labels: 1 - p**2 - (1 - p)**2, where p is the share of 1s. An empty list has Gini 0.""" raise NotImplementedError def entropy(labels): """Return the entropy of the labels in bits: -p * log2(p) - (1 - p) * log2(1 - p). Leave out a term whose share is 0.""" raise NotImplementedError def split(points, labels, feature, threshold): """Return (left, right): the labels of the points whose feature is less than threshold, and the labels of the rest.""" raise NotImplementedError def weighted_impurity(left, right): """Return the Gini of each side, weighted by its share of all the labels: len(left)/n * gini(left) + len(right)/n * gini(right).""" raise NotImplementedError def best_split(points, labels): """Return (feature, threshold, gain) for the question with the largest gain. Try every feature, in order, and every threshold halfway between two neighbouring different values of that feature. The gain is the group's Gini minus the split's weighted impurity. If two questions have the same gain, keep the first one found. Return None if no question has a gain above 0. """ raise NotImplementedError def build(points, labels, depth): """Grow a tree, at most depth questions deep. A leaf is {"leaf": 1} or {"leaf": 0}: what most of its labels are (a tie gives 1). A question is {"feature": f, "threshold": t, "below": tree, "above": tree}, where below is the tree for the points whose feature is less than t. Make a leaf when depth is 0 or no question has a gain. """ raise NotImplementedError def predict(tree, point): """Follow the tree's questions for the point and return the label of the leaf it reaches.""" raise NotImplementedError def accuracy(tree, points, labels): """Return the share of the points the tree labels right.""" raise NotImplementedError Stuck? Write gini first, then split, then best_split, which tries every threshold halfway between neighbouring values. build calls itself on each side. The Solution tab has one way to write each function.
In practice: trees in real projects
Controlling the size. scikit-learn’s DecisionTreeClassifier grows until every leaf is pure unless you stop it. The usual controls are max_depth, min_samples_leaf (the fewest training examples a leaf may hold) and min_samples_split. Choose them on a validation set, or with cross-validation.
Categories and missing values. Trees split numbers by thresholds, so a category such as a city is usually turned into numbers first, one column per city. Some libraries, such as LightGBM, split categories and handle missing values directly.
Importance can mislead. The impurity-based importance of this lesson favours inputs with many different values, which offer more thresholds to try. Permutation importance, which shuffles one input on the validation set and measures how much worse the model gets, is a fairer check.
Test your knowledge
01A group has 6 passes and 2 fails. What is its Gini impurity?Show answer
p = 6 ÷ 8 = 0.75, so Gini = 1 − 0.75² − 0.25² = 1 − 0.5625 − 0.0625 = 0.375.
02How does a tree choose its first question?Show answer
It tries every input and every threshold, splits the training examples, and works out the impurity of the two sides, each weighted by its share. It picks the question that lowers the impurity most, the largest gain.
03Why does a fully grown tree overfit?Show answer
It keeps splitting until each leaf is pure, so its last questions separate a few training examples, often by their noise. It is perfect on the training data and worse on new data. Limit the depth or the leaf size.
04Does a decision tree need its inputs scaled?Show answer
No. Each question compares one input with a threshold. Changing the input’s unit moves the threshold the same way, and the tree makes the same predictions.
05What does a leaf of a regression tree predict?Show answer
The mean of the training values that reach it. A split is chosen to lower the squared error most, and the tree’s predictions form a staircase.
Exit ticket
One last question on the main idea of the lesson.
A decision tree grown until every leaf is pure gets every training example right. What should you expect on new examples?
Show the answer
B: Worse: its last questions fit the noise in a few training examples. Limit its depth, chosen on a validation set. In this lesson the fully grown tree gets all 20 training students right but only 72.5% of the validation students, while a tree with one question gets 82.5%.