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.

Lesson 8 stepsPractice 5 problemsExercise Python, 8 functionsQuiz 5 questions

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 asks yes-or-no questions, one after another:Studied less than 4.8 hours?yes: slept less than 6.8 hours?yes: fail (0 of 7)no: pass (2 of 3)no: studied less than 6.65 hours?yes: pass (3 of 4)no: pass (6 of 6)The new student (5.2 hours, 6 hours): pass.Each question is a line on the plot. The leaves are the boxes.0246810468hours studiedhours sleptpassedfailed
01/08

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.

0246810456789hours studiedhours slept

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
StudentStudySleepResultstudy < 2.50?
11.44.2failedyes
22.65.7failedno
31.68.7failedyes
43.66.4failedno
56.94.1passedno
60.85.9failedyes
768.5passedno
87.18.1passedno
96.15.5passedno
101.94.1failedyes
116.27.6passedno
123.17.2passedno
132.84failedno
148.28.3passedno
158.74.6passedno
168.65.7passedno
177.68.3passedno
186.46.1failedno
191.97.9passedyes
201.44.7failedyes

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

  1. Problem 1

    1 point

    A 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

  2. Problem 2

    2 points

    A 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

  3. Problem 3

    2 points

    And 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

  4. Problem 4

    2 points

    What 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

  5. Problem 5

    3 points

    A 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

  1. 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.

  2. 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.

  3. 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.

  4. 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.

  5. 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%.