Precision, recall and other metrics
A model can be 99% accurate and useless. Learn the confusion matrix, precision and recall, the trade-off a threshold makes between them, the F1 score, the ROC curve and its AUC, and what a rare class does to precision.
Warm-up
One question before the lesson. Choose an answer and check it.
A model that checks payments for fraud is right about 99 times in every 100. Is it a good model?
Show the answer
B: It depends on how many payments are fraud, and how many of those it catches. If 1 payment in 100 is fraud, a model that always answers “not fraud” is right 99 times in 100 and catches no fraud at all. Accuracy on its own cannot tell a useful model from a useless one. This lesson measures the two kinds of mistake separately.
Step 1 Accuracy can mislead
Lesson 4’s model gave each student a chance and a prediction, pass or fail. How good are a model’s predictions? The first measure most people reach for is accuracy, the share of predictions that are right:
accuracy = right predictions ÷ all predictions
Accuracy treats every mistake the same, and that can hide a model that is no use at all. Take 1,000 payments, 10 of them fraud. A model that answers “not fraud” for every payment is right about the 990 that are not fraud:
accuracy = 990 ÷ 1,000 = 0.99
That is 99% accurate, and it catches none of the 10 frauds. The rarer the thing a model looks for, the easier it is to score well by never finding it.
This lesson uses a screening test for an illness. A model trained as in Lesson 4, on 200 earlier patients, gives each of 20 new patients a chance p of having the illness, from the level of a marker in their blood, a number from 0 to 10. 6 of the 20 are ill:
| Patient | Marker | p | Truth |
|---|---|---|---|
| 1 | 0.2 | 0.01 | healthy |
| 2 | 0.6 | 0.01 | healthy |
| 3 | 1.2 | 0.02 | healthy |
| 4 | 1.8 | 0.03 | healthy |
| 5 | 2.3 | 0.04 | healthy |
| 6 | 2.6 | 0.04 | healthy |
| 7 | 3.3 | 0.07 | healthy |
| 8 | 3.8 | 0.1 | healthy |
| 9 | 4.2 | 0.13 | healthy |
| 10 | 4.8 | 0.19 | healthy |
| 11 | 5.2 | 0.25 | healthy |
| 12 | 5.9 | 0.35 | healthy |
| 13 | 6.2 | 0.41 | ill |
| 14 | 6.8 | 0.52 | ill |
| 15 | 7.2 | 0.59 | ill |
| 16 | 7.6 | 0.66 | healthy |
| 17 | 8.1 | 0.74 | healthy |
| 18 | 8.7 | 0.82 | ill |
| 19 | 9.3 | 0.87 | ill |
| 20 | 9.7 | 0.9 | ill |
With a threshold of 0.5, as in Lesson 4, the model flags (predicts ill for) every patient whose p is at least 0.5. It is right about 17 of the 20: an accuracy of 17 ÷ 20 = 0.85. A model that calls everyone healthy is right about the 14 who are healthy: 14 ÷ 20 = 0.7.
At a threshold of 0.8 the model is also right about 17. But at 0.5 it sends 1 ill patient home and flags 2 healthy ones, while at 0.8 it sends 3 ill patients home and flags no healthy ones. Accuracy cannot tell these apart. The next steps count each kind of mistake on its own.
Try it yourself
Move the threshold and watch the confusion matrix, precision, recall and F1 change, and the point move along the ROC curve. Find the threshold with the best F1, then the lowest threshold with no false alarms.
Filled dots are ill, hollow dots healthy. The shaded side of the dashed line is flagged, and a red ring marks a mistake.
The ROC curve. The dot is this threshold. Area under the curve: 0.93
Accuracy = (TP + TN) ÷ 20 = (5 + 12) ÷ 20 = 0.85
Precision = TP ÷ (TP + FP) = 5 ÷ 7 = 0.71
Recall = TP ÷ (TP + FN) = 5 ÷ 6 = 0.83
Specificity = TN ÷ (TN + FP) = 12 ÷ 14 = 0.86
F1 = 2TP ÷ (2TP + FP + FN) = 10 ÷ 13 = 0.77
False positive rate = FP ÷ (FP + TN) = 2 ÷ 14 = 0.14
Best F1: 0.86, at a threshold of 0.41
Everyone to the right of the threshold is flagged. Red rings mark the mistakes: false alarms on the healthy row, misses on the ill row.
Show the calculation
| Patient | Marker | p | Truth | p ≥ 0.50? | Outcome |
|---|---|---|---|---|---|
| 1 | 0.2 | 0.01 | healthy | no, cleared | TN |
| 2 | 0.6 | 0.01 | healthy | no, cleared | TN |
| 3 | 1.2 | 0.02 | healthy | no, cleared | TN |
| 4 | 1.8 | 0.03 | healthy | no, cleared | TN |
| 5 | 2.3 | 0.04 | healthy | no, cleared | TN |
| 6 | 2.6 | 0.04 | healthy | no, cleared | TN |
| 7 | 3.3 | 0.07 | healthy | no, cleared | TN |
| 8 | 3.8 | 0.1 | healthy | no, cleared | TN |
| 9 | 4.2 | 0.13 | healthy | no, cleared | TN |
| 10 | 4.8 | 0.19 | healthy | no, cleared | TN |
| 11 | 5.2 | 0.25 | healthy | no, cleared | TN |
| 12 | 5.9 | 0.35 | healthy | no, cleared | TN |
| 13 | 6.2 | 0.41 | ill | no, cleared | FN |
| 14 | 6.8 | 0.52 | ill | yes, flagged | TP |
| 15 | 7.2 | 0.59 | ill | yes, flagged | TP |
| 16 | 7.6 | 0.66 | healthy | yes, flagged | FP |
| 17 | 8.1 | 0.74 | healthy | yes, flagged | FP |
| 18 | 8.7 | 0.82 | ill | yes, flagged | TP |
| 19 | 9.3 | 0.87 | ill | yes, flagged | TP |
| 20 | 9.7 | 0.9 | ill | yes, flagged | TP |
TP = ill and flagged = 5; FP = healthy and flagged = 2
FN = ill and cleared = 1; TN = healthy and cleared = 12
5 + 2 + 1 + 12 = 20 patients
p is the model’s chance that a patient is ill, to 2 decimal places. A patient is flagged when p is at least the threshold.
Drag the threshold, or press a preset. Find the threshold with the best F1, then the lowest threshold with no false alarms.
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 spam filter checks 1,000 emails. It flags 40 as spam, and 30 of those really are spam. It lets 20 spam emails through without flagging them. What is its accuracy?
Hint 1
Accuracy is the share of all the emails it gets right: (TP + TN) ÷ total.
Hint 2
TN is every email left once the other three groups are counted: 1,000 − 30 − 10 − 20.
Solution
TP = 30, FP = 40 − 30 = 10, FN = 20
TN = 1,000 − 30 − 10 − 20 = 940
accuracy = (TP + TN) ÷ total
= (30 + 940) ÷ 1,000 = 0.97
Problem 2
2 pointsFor the same filter, what is its precision?
Hint 1
Precision asks: of the emails it flagged, what share really are spam?
Hint 2
precision = TP ÷ (TP + FP).
Solution
precision = TP ÷ (TP + FP)
= 30 ÷ (30 + 10)
= 30 ÷ 40 = 0.75
Problem 3
2 pointsAnd its recall?
Hint 1
Recall asks: of all the spam, what share did it flag?
Hint 2
recall = TP ÷ (TP + FN). The spam it let through is FN.
Solution
recall = TP ÷ (TP + FN)
= 30 ÷ (30 + 20)
= 30 ÷ 50 = 0.6
Problem 4
2 pointsAnd its F1 score? Give it to 3 decimal places.
Hint 1
F1 = 2 × precision × recall ÷ (precision + recall).
Hint 2
Or straight from the counts: F1 = 2TP ÷ (2TP + FP + FN).
Solution
F1 = 2 × precision × recall ÷ (precision + recall)
= 2 × 0.75 × 0.6 ÷ (0.75 + 0.6)
= 0.9 ÷ 1.35 = 0.6667
From the counts: 2 × 30 ÷ (2 × 30 + 10 + 20) = 60 ÷ 90 = 0.6667
Problem 5
3 pointsA test for an illness catches 90% of the people who have it, and wrongly flags 5% of the people who do not. In a town of 10,000, 1% of the people have the illness. Of the people the test flags, what share really have it? Give it to 3 decimal places.
Hint 1
Turn the percentages into counts first: how many people have the illness, and how many do not?
Hint 2
The test flags 90% of the 100 who have it (TP) and 5% of the 9,900 who do not (FP). Then precision = TP ÷ (TP + FP).
Solution
ill = 1% of 10,000 = 100
healthy = 10,000 − 100 = 9,900
TP = 90% of 100 = 90
FP = 5% of 9,900 = 495
precision = TP ÷ (TP + FP) = 90 ÷ (90 + 495)
= 90 ÷ 585 = 0.1538
Programming exercise
Write the confusion matrix, accuracy, precision, recall, F1, the ROC curve and the area under it, in plain Python. Save metrics.py and test_metrics.py in the same folder, fill in each function in metrics.py, and run the tests:
python test_metrics.py
"""Precision, recall and other metrics: programming exercise. Write the measures of a classifier in plain Python: the four counts of theconfusion matrix, accuracy, precision, recall, F1, and the ROC curve with thearea under it. Run the tests from this folder: python test_metrics.py actual is a list of the true classes: 1 for positive (such as ill) and 0 fornegative (healthy). scores is a list of the model's chances that each exampleis positive, in the same order, and predicted is a list of 1s and 0s.""" def predict(scores, threshold): """Return a list with 1 for each score that is at least threshold, and 0 for every other score.""" raise NotImplementedError def confusion_counts(actual, predicted): """Return (tp, fp, fn, tn). tp: positive and predicted positive. fp: negative but predicted positive. fn: positive but predicted negative. tn: negative and predicted negative. """ raise NotImplementedError def accuracy(tp, fp, fn, tn): """Return the share of all predictions that are right: (tp + tn) / (tp + fp + fn + tn).""" raise NotImplementedError def precision(tp, fp): """Return tp / (tp + fp), the share of positive predictions that are right. If nothing is predicted positive, tp + fp is 0: return 0.0, as scikit-learn does. """ raise NotImplementedError def recall(tp, fn): """Return tp / (tp + fn), the share of the positives that are found. Return 0.0 if there are no positives.""" raise NotImplementedError def f1(tp, fp, fn): """Return the F1 score, 2 * tp / (2 * tp + fp + fn). Return 0.0 if the bottom is 0.""" raise NotImplementedError def roc_points(actual, scores): """Return the ROC curve as a list of (false positive rate, true positive rate) pairs. Start at (0.0, 0.0), where nothing is predicted positive. Then use each different score as the threshold, from the highest to the lowest, and add the point it gives. Equal scores are predicted positive together, so they make one point. The last point is (1.0, 1.0). The false positive rate is fp / (fp + tn), and the true positive rate is tp / (tp + fn). """ raise NotImplementedError def auc(points): """Return the area under a curve given as (x, y) points in order of x. Add up one trapezoid for each pair of neighbouring points: (x2 - x1) * (y1 + y2) / 2. """ raise NotImplementedError Stuck? Count TP, FP, FN and TN first: every other measure is worked out from those four. For the ROC curve, go through the scores from highest to lowest and flag one more group of equal scores at each step. The Solution tab has one way to write each function.
In practice: measuring a classifier
Choose the threshold on the validation set. If you pick the threshold with the best F1 on the test set and then report that F1, the number is too good, because the threshold was fitted to those examples. Choose it on the validation set of Lesson 2, then measure the model once on the test set.
The share of positives changes after launch. A fraud model tested on a sample where 1 payment in 10 is fraud shows a far higher precision than it gets in use, where fraud may be 1 in 1,000. Recall and the false positive rate stay the same, so work out the precision for the real share, as in step 8, and keep measuring it once the model is live.
More than two classes. scikit-learn’s classification_report gives precision, recall and F1 for each class, then two averages. The macro average weights every class equally. The weighted average weights each class by how many examples it has, so a large class doing well can hide a small one doing badly. Read the rows for each class before the averages.
Test your knowledge
01A model is 99% accurate on data where 1% of the examples are positive. What would you check first?Show answer
The confusion matrix, and the recall on the positive class. A model that always answers no is also 99% accurate on that data, and it finds nothing.
02When would you choose high recall over high precision?Show answer
When a miss costs more than a false alarm. A screening test for a serious illness should catch nearly every case, and a second test can clear the false alarms. A spam filter is the other way round: hiding a real email costs more than letting a spam through.
03Why is F1 a harmonic mean rather than an ordinary average?Show answer
The harmonic mean is pulled towards the smaller number, so a model cannot score well by being excellent at one and poor at the other. With precision 1 and recall 0.1, the ordinary average is 0.55 but F1 is 2 × 1 × 0.1 ÷ (1 + 0.1) = 0.18.
04What does an AUC of 0.5 mean, and an AUC of 0.93?Show answer
AUC is the chance that the model gives a random positive a higher score than a random negative. 0.5 is what guessing gets. 0.93 means 93 of every 100 such pairs are in the right order. It does not depend on the threshold, so it compares models before a threshold is chosen.
05Why can a model with an AUC of 0.95 still be poor at finding a rare disease?Show answer
Neither of the ROC curve’s rates depends on how rare the disease is. If 1 person in 1,000 has it, a test that catches 90% of cases and flags just 1% of healthy people still flags about 10 healthy people for every person who is ill, so fewer than 1 flag in 10 is right. Judge it by precision and recall on data as rare as the real thing.
Exit ticket
One last question on the main idea of the lesson.
You lower a spam filter’s threshold from 0.5 to 0.3. What happens?
Show the answer
A: It flags more emails. Recall rises or stays the same, and precision usually falls. A lower threshold flags every email the old one did, and some more. Every spam caught before is still caught, so recall cannot fall. The extra emails it flags are more often real ones, so precision usually falls.