Code
library(tidyverse)
library(gt)Tyler Grimes
July 4, 2024
When predicting binary outcomes, if we want to threshold our estimated probabilities to obtain predicted classes, then we devise a strategy for determining the optimal threshold. In plain terms, a threshold is just the cutoff probability above which we call a prediction “positive,” and prevalence is how common the positive outcome actually is. Choosing a threshold always involves a trade-off between sensitivity and specificity (false positives and false negatives). But if we know the prevalence of the outcome varies across some known strata (subgroups), should we choose a different threshold within each stratum? And how does that choice affect performance, both overall and within each group? This post works through a small simulated example to build some intuition.
Let’s simulate data from three groups where the true prevalence ranges from about 0.6 to 0.9 across the groups. Our dataset will be composed of samples from these three groups with unequal sizes, so that one group (“b”) ends up considerably smaller than the other two.
set.seed(0)
n = 300
x = rnorm(n)
groups = c("a", "b", "c")
z = sample(groups, n, replace = TRUE, prob = c(0.3, 0.2, 0.5))
z = factor(z, levels = groups)
beta.x = 1
beta.z = c("a" = 0, "b" = 1, "c" = 2)
epsilon = rnorm(n, 0, 0.2)
prob = 1 / (1 + exp(-(0.5 + beta.x * x + beta.z[z] + epsilon)))
y = rbinom(n, 1, prob)
df = data.frame(x = x, z = z, y = y, prob = prob)group_summary = df %>%
group_by(z) %>%
summarize(n = n(), prevalence = mean(y), .groups = "drop") %>%
mutate(z = as.character(z))
overall_row = df %>%
summarize(n = n(), prevalence = mean(y)) %>%
mutate(z = "Overall")
bind_rows(overall_row, group_summary) %>%
mutate(
group_label = case_match(z, "a" ~ "Group a", "b" ~ "Group b", "c" ~ "Group c", "Overall" ~ "Overall"),
group_label = factor(group_label, levels = c("Overall", "Group a", "Group b", "Group c"))
) %>%
arrange(group_label) %>%
select(group_label, n, prevalence) %>%
gt() %>%
cols_label(group_label = "Group", n = "N", prevalence = "Prevalence of y = 1") %>%
fmt_percent(columns = prevalence, decimals = 1) %>%
fmt_number(columns = n, decimals = 0)Warning: There was 1 warning in `mutate()`.
ℹ In argument: `group_label = case_match(...)`.
Caused by warning:
! `case_match()` was deprecated in dplyr 1.2.0.
ℹ Please use `recode_values()` instead.
| Group | N | Prevalence of y = 1 |
|---|---|---|
| Overall | 300 | 79.7% |
| Group a | 87 | 63.2% |
| Group b | 61 | 78.7% |
| Group c | 152 | 89.5% |
As intended, the observed prevalence ranges from about 63% to 89% across the three groups (Table 1), and group “b” is indeed the smallest of the three.
We’ll fit a logistic regression model with a single continuous predictor (x) and the group variable (z), then compare two strategies for turning predicted probabilities into predicted classes:
(Intercept) x zb zc
0.6803849 0.6825582 0.6736314 1.5837307
The coefficient on x is positive (0.68), and the coefficients on z increase from group a to group b to group c, matching how the simulation was built to have increasing prevalence across those groups.
df$p = predict(fit, newdata = df, type = "response")
df = df %>%
mutate(thr = mean(y)) %>%
group_by(z) %>%
mutate(thr.z = mean(y) - 0.2 * (mean(y) - thr)) %>%
ungroup() %>%
mutate(
y.hat = 1 * (p > thr),
y.hat.z = 1 * (p > thr.z)
)
thr.group = sapply(groups, function(group) df$thr.z[df$z == group][1])The group-specific threshold, thr.z, is a weighted average of the group’s own prevalence (80%) and the overall prevalence (20%) — a simple form of shrinkage, or partial pooling, toward the overall rate. We used a 0.2 shrinkage weight here mostly for illustration; it’s a dial that trades off two things. Leaning more on the overall prevalence (a weight closer to 1) makes the group thresholds more stable when a group is small and its own prevalence estimate is noisy, at the cost of responding less to genuine differences between groups. Leaning more on each group’s own prevalence (a weight closer to 0) does the opposite. In a real analysis, this weight would typically be chosen by cross-validation, or by fitting a proper hierarchical (multilevel) model that estimates the right amount of pooling from the data itself, rather than being fixed by hand as we’ve done here.
performance_by_group = function(data, pred) {
data %>%
group_by(z) %>%
summarize(
n = n(),
prevalence = mean(y),
pred_pos = mean({{ pred }}),
accuracy = mean(y == {{ pred }}),
sensitivity = mean({{ pred }}[y == 1] == 1),
specificity = mean({{ pred }}[y == 0] == 0),
ppv = mean(y[{{ pred }} == 1] == 1),
npv = mean(y[{{ pred }} == 0] == 0),
.groups = "drop"
) %>%
mutate(z = as.character(z))
}
performance_overall = function(data, pred) {
data %>%
summarize(
n = n(),
prevalence = mean(y),
pred_pos = mean({{ pred }}),
accuracy = mean(y == {{ pred }}),
sensitivity = mean({{ pred }}[y == 1] == 1),
specificity = mean({{ pred }}[y == 0] == 0),
ppv = mean(y[{{ pred }} == 1] == 1),
npv = mean(y[{{ pred }} == 0] == 0)
)
}
perf_all = bind_rows(
performance_overall(df, y.hat) %>% mutate(z = "Overall", strategy = "Single threshold"),
performance_overall(df, y.hat.z) %>% mutate(z = "Overall", strategy = "Varying threshold"),
performance_by_group(df, y.hat) %>% mutate(strategy = "Single threshold"),
performance_by_group(df, y.hat.z) %>% mutate(strategy = "Varying threshold")
) %>%
mutate(
group_label = case_match(z, "a" ~ "Group a", "b" ~ "Group b", "c" ~ "Group c", "Overall" ~ "Overall"),
group_label = factor(group_label, levels = c("Overall", "Group a", "Group b", "Group c")),
strategy = factor(strategy, levels = c("Single threshold", "Varying threshold"))
) %>%
arrange(group_label, strategy) %>%
select(group_label, strategy, n, prevalence, pred_pos, accuracy, sensitivity, specificity, ppv, npv)
acc_single = perf_all$accuracy[perf_all$group_label == "Overall" & perf_all$strategy == "Single threshold"]
acc_varying = perf_all$accuracy[perf_all$group_label == "Overall" & perf_all$strategy == "Varying threshold"]
range_single = perf_all %>%
filter(group_label != "Overall", strategy == "Single threshold") %>%
summarize(min = min(accuracy), max = max(accuracy))
range_varying = perf_all %>%
filter(group_label != "Overall", strategy == "Varying threshold") %>%
summarize(min = min(accuracy), max = max(accuracy))Looking only at overall accuracy, the single-threshold strategy comes out slightly ahead: 69.3% versus 67.7% for the varying-threshold strategy. But that one-number comparison hides what’s happening inside each group, which is really the point of this post.
A few quick definitions, since the table below uses standard classification terms:
perf_all %>%
gt(groupname_col = "group_label") %>%
cols_label(
strategy = "Strategy",
n = "N",
prevalence = "Prevalence",
pred_pos = "Predicted positive rate",
accuracy = "Accuracy",
sensitivity = "Sensitivity",
specificity = "Specificity",
ppv = "PPV",
npv = "NPV"
) %>%
fmt_percent(columns = c(prevalence, pred_pos, accuracy, sensitivity, specificity, ppv, npv), decimals = 1) %>%
fmt_number(columns = n, decimals = 0)| Strategy | N | Prevalence | Predicted positive rate | Accuracy | Sensitivity | Specificity | PPV | NPV |
|---|---|---|---|---|---|---|---|---|
| Overall | ||||||||
| Single threshold | 300 | 79.7% | 61.0% | 69.3% | 69.0% | 70.5% | 90.2% | 36.8% |
| Varying threshold | 300 | 79.7% | 57.3% | 67.7% | 65.7% | 75.4% | 91.3% | 35.9% |
| Group a | ||||||||
| Single threshold | 87 | 63.2% | 13.8% | 46.0% | 18.2% | 93.8% | 83.3% | 40.0% |
| Varying threshold | 87 | 63.2% | 40.2% | 63.2% | 52.7% | 81.2% | 82.9% | 50.0% |
| Group b | ||||||||
| Single threshold | 61 | 78.7% | 45.9% | 60.7% | 54.2% | 84.6% | 92.9% | 33.3% |
| Varying threshold | 61 | 78.7% | 45.9% | 60.7% | 54.2% | 84.6% | 92.9% | 33.3% |
| Group c | ||||||||
| Single threshold | 152 | 89.5% | 94.1% | 86.2% | 94.9% | 12.5% | 90.2% | 22.2% |
| Varying threshold | 152 | 89.5% | 71.7% | 73.0% | 75.0% | 56.2% | 93.6% | 20.9% |
Table 2 tells a more complete story than the overall numbers alone. Within groups, the single-threshold strategy’s accuracy ranges from 46.0% to 86.2% — a spread of 40.2 percentage points. The varying-threshold strategy narrows that spread to 12.4 points (from 60.7% to 73.0%), at a cost of only 1.7 points of overall accuracy. In other words: letting the threshold adapt to each group’s prevalence gives much more consistent performance across groups, for a small overall price.
g = df %>%
ggplot(aes(x = p, fill = z, group = z)) +
geom_histogram(position = position_identity(), alpha = 0.6, bins = 30) +
geom_vline(xintercept = unique(df$thr), col = "black", linewidth = 0.8) +
geom_vline(xintercept = thr.group, col = scales::hue_pal()(length(thr.group)), linewidth = 0.8) +
labs(
title = "Predicted probabilities and classification thresholds",
x = "Predicted probability of y = 1",
y = "Count",
fill = "Group"
) +
theme_bw(base_size = 13)
gFigure 1 shows why: group c’s distribution sits well to the right of group a’s, so a single threshold necessarily fits one of them worse than the other. The group-specific thresholds shift to follow each group’s distribution, while still being pulled somewhat toward the black line.
Imagine that the strata in this simulation were gender, race, or some other demographic variable. Although the model itself accounted for the group variable — and so correctly captured the underlying prevalence differences — the additional step of classifying (turning a probability into a yes/no decision) can reintroduce bias if we don’t also account for those prevalence differences when we set the threshold. A single threshold can look attractive because it optimizes overall accuracy, but “overall” can quietly average over very uneven performance across groups, as we saw above (a 40.2-point accuracy spread with a single threshold, versus 12.4 points with varying thresholds). Whether that consistency is worth the small overall accuracy cost depends on the application — but it’s a trade-off worth making deliberately, rather than by default.
---
title: "Adjusting classification threshold by strata"
author: Tyler Grimes
date: '2024-07-04'
categories:
- Predictive modeling
- Classification
tags: []
format:
html: default
image: adjusting-class-threshold-by-strata2-thumb.png
---
When predicting binary outcomes, if we want to threshold our estimated probabilities to obtain predicted classes, then we devise a strategy for determining the optimal threshold. In plain terms, a **threshold** is just the cutoff probability above which we call a prediction "positive," and **prevalence** is how common the positive outcome actually is. Choosing a threshold always involves a trade-off between sensitivity and specificity (false positives and false negatives). But if we know the prevalence of the outcome varies across some known strata (subgroups), should we choose a different threshold within each stratum? And how does that choice affect performance, both overall and within each group? This post works through a small simulated example to build some intuition.
```{r}
#| message: false
library(tidyverse)
library(gt)
```
## Simulating grouped data
Let's simulate data from three groups where the true prevalence ranges from about 0.6 to 0.9 across the groups. Our dataset will be composed of samples from these three groups with unequal sizes, so that one group ("b") ends up considerably smaller than the other two.
```{r}
set.seed(0)
n = 300
x = rnorm(n)
groups = c("a", "b", "c")
z = sample(groups, n, replace = TRUE, prob = c(0.3, 0.2, 0.5))
z = factor(z, levels = groups)
beta.x = 1
beta.z = c("a" = 0, "b" = 1, "c" = 2)
epsilon = rnorm(n, 0, 0.2)
prob = 1 / (1 + exp(-(0.5 + beta.x * x + beta.z[z] + epsilon)))
y = rbinom(n, 1, prob)
df = data.frame(x = x, z = z, y = y, prob = prob)
```
```{r}
#| label: tbl-group-summary
#| tbl-cap: "Group sizes and observed outcome prevalence in the simulated data"
group_summary = df %>%
group_by(z) %>%
summarize(n = n(), prevalence = mean(y), .groups = "drop") %>%
mutate(z = as.character(z))
overall_row = df %>%
summarize(n = n(), prevalence = mean(y)) %>%
mutate(z = "Overall")
bind_rows(overall_row, group_summary) %>%
mutate(
group_label = case_match(z, "a" ~ "Group a", "b" ~ "Group b", "c" ~ "Group c", "Overall" ~ "Overall"),
group_label = factor(group_label, levels = c("Overall", "Group a", "Group b", "Group c"))
) %>%
arrange(group_label) %>%
select(group_label, n, prevalence) %>%
gt() %>%
cols_label(group_label = "Group", n = "N", prevalence = "Prevalence of y = 1") %>%
fmt_percent(columns = prevalence, decimals = 1) %>%
fmt_number(columns = n, decimals = 0)
```
As intended, the observed prevalence ranges from about `r sprintf("%.0f%%", 100 * min(group_summary$prevalence))` to `r sprintf("%.0f%%", 100 * max(group_summary$prevalence))` across the three groups (@tbl-group-summary), and group "b" is indeed the smallest of the three.
## Fitting the model and defining two thresholds
We'll fit a logistic regression model with a single continuous predictor (`x`) and the group variable (`z`), then compare two strategies for turning predicted probabilities into predicted classes:
1. **Single threshold**: one cutoff, based on the overall prevalence, applied to every observation regardless of group.
2. **Varying threshold**: a group-specific cutoff, based mostly on that group's own prevalence but pulled partway back toward the overall prevalence.
```{r}
fit = glm(y ~ x + z, family = "binomial", data = df)
coef(fit)
```
The coefficient on `x` is positive (`r round(coef(fit)[["x"]], 2)`), and the coefficients on `z` increase from group a to group b to group c, matching how the simulation was built to have increasing prevalence across those groups.
```{r}
df$p = predict(fit, newdata = df, type = "response")
df = df %>%
mutate(thr = mean(y)) %>%
group_by(z) %>%
mutate(thr.z = mean(y) - 0.2 * (mean(y) - thr)) %>%
ungroup() %>%
mutate(
y.hat = 1 * (p > thr),
y.hat.z = 1 * (p > thr.z)
)
thr.group = sapply(groups, function(group) df$thr.z[df$z == group][1])
```
The group-specific threshold, `thr.z`, is a weighted average of the group's own prevalence (80%) and the overall prevalence (20%) — a simple form of shrinkage, or partial pooling, toward the overall rate. We used a 0.2 shrinkage weight here mostly for illustration; it's a dial that trades off two things. Leaning more on the overall prevalence (a weight closer to 1) makes the group thresholds more stable when a group is small and its own prevalence estimate is noisy, at the cost of responding less to genuine differences between groups. Leaning more on each group's own prevalence (a weight closer to 0) does the opposite. In a real analysis, this weight would typically be chosen by cross-validation, or by fitting a proper hierarchical (multilevel) model that estimates the right amount of pooling from the data itself, rather than being fixed by hand as we've done here.
## Comparing performance across strata
```{r}
#| label: performance-calc
performance_by_group = function(data, pred) {
data %>%
group_by(z) %>%
summarize(
n = n(),
prevalence = mean(y),
pred_pos = mean({{ pred }}),
accuracy = mean(y == {{ pred }}),
sensitivity = mean({{ pred }}[y == 1] == 1),
specificity = mean({{ pred }}[y == 0] == 0),
ppv = mean(y[{{ pred }} == 1] == 1),
npv = mean(y[{{ pred }} == 0] == 0),
.groups = "drop"
) %>%
mutate(z = as.character(z))
}
performance_overall = function(data, pred) {
data %>%
summarize(
n = n(),
prevalence = mean(y),
pred_pos = mean({{ pred }}),
accuracy = mean(y == {{ pred }}),
sensitivity = mean({{ pred }}[y == 1] == 1),
specificity = mean({{ pred }}[y == 0] == 0),
ppv = mean(y[{{ pred }} == 1] == 1),
npv = mean(y[{{ pred }} == 0] == 0)
)
}
perf_all = bind_rows(
performance_overall(df, y.hat) %>% mutate(z = "Overall", strategy = "Single threshold"),
performance_overall(df, y.hat.z) %>% mutate(z = "Overall", strategy = "Varying threshold"),
performance_by_group(df, y.hat) %>% mutate(strategy = "Single threshold"),
performance_by_group(df, y.hat.z) %>% mutate(strategy = "Varying threshold")
) %>%
mutate(
group_label = case_match(z, "a" ~ "Group a", "b" ~ "Group b", "c" ~ "Group c", "Overall" ~ "Overall"),
group_label = factor(group_label, levels = c("Overall", "Group a", "Group b", "Group c")),
strategy = factor(strategy, levels = c("Single threshold", "Varying threshold"))
) %>%
arrange(group_label, strategy) %>%
select(group_label, strategy, n, prevalence, pred_pos, accuracy, sensitivity, specificity, ppv, npv)
acc_single = perf_all$accuracy[perf_all$group_label == "Overall" & perf_all$strategy == "Single threshold"]
acc_varying = perf_all$accuracy[perf_all$group_label == "Overall" & perf_all$strategy == "Varying threshold"]
range_single = perf_all %>%
filter(group_label != "Overall", strategy == "Single threshold") %>%
summarize(min = min(accuracy), max = max(accuracy))
range_varying = perf_all %>%
filter(group_label != "Overall", strategy == "Varying threshold") %>%
summarize(min = min(accuracy), max = max(accuracy))
```
Looking only at overall accuracy, the single-threshold strategy comes out slightly ahead: `r sprintf("%.1f%%", 100 * acc_single)` versus `r sprintf("%.1f%%", 100 * acc_varying)` for the varying-threshold strategy. But that one-number comparison hides what's happening inside each group, which is really the point of this post.
### A closer look, by group
A few quick definitions, since the table below uses standard classification terms:
- **Accuracy**: the proportion of all predictions that are correct.
- **Sensitivity** (recall): among people who are truly positive, the proportion the classifier catches.
- **Specificity**: among people who are truly negative, the proportion the classifier correctly clears.
- **PPV** (precision): among people the classifier flags as positive, the proportion who really are positive.
- **NPV**: among people the classifier clears as negative, the proportion who really are negative.
```{r}
#| label: tbl-performance
#| tbl-cap: "Classifier performance by group and thresholding strategy. 'Overall' pools all three groups together. PPV or NPV can show as NaN if a group has no predicted positives or negatives, respectively."
perf_all %>%
gt(groupname_col = "group_label") %>%
cols_label(
strategy = "Strategy",
n = "N",
prevalence = "Prevalence",
pred_pos = "Predicted positive rate",
accuracy = "Accuracy",
sensitivity = "Sensitivity",
specificity = "Specificity",
ppv = "PPV",
npv = "NPV"
) %>%
fmt_percent(columns = c(prevalence, pred_pos, accuracy, sensitivity, specificity, ppv, npv), decimals = 1) %>%
fmt_number(columns = n, decimals = 0)
```
@tbl-performance tells a more complete story than the overall numbers alone. Within groups, the single-threshold strategy's accuracy ranges from `r sprintf("%.1f%%", 100 * range_single$min)` to `r sprintf("%.1f%%", 100 * range_single$max)` — a spread of `r sprintf("%.1f", 100 * (range_single$max - range_single$min))` percentage points. The varying-threshold strategy narrows that spread to `r sprintf("%.1f", 100 * (range_varying$max - range_varying$min))` points (from `r sprintf("%.1f%%", 100 * range_varying$min)` to `r sprintf("%.1f%%", 100 * range_varying$max)`), at a cost of only `r sprintf("%.1f", 100 * (acc_single - acc_varying))` points of overall accuracy. In other words: letting the threshold adapt to each group's prevalence gives much more consistent performance across groups, for a small overall price.
## Visualizing the thresholds
```{r}
#| label: fig-thresholds
#| fig-cap: "Distribution of predicted probabilities by group. The black line is the single overall threshold; the colored lines are each group's own (shrinkage-adjusted) threshold, matching that group's fill color."
g = df %>%
ggplot(aes(x = p, fill = z, group = z)) +
geom_histogram(position = position_identity(), alpha = 0.6, bins = 30) +
geom_vline(xintercept = unique(df$thr), col = "black", linewidth = 0.8) +
geom_vline(xintercept = thr.group, col = scales::hue_pal()(length(thr.group)), linewidth = 0.8) +
labs(
title = "Predicted probabilities and classification thresholds",
x = "Predicted probability of y = 1",
y = "Count",
fill = "Group"
) +
theme_bw(base_size = 13)
g
```
```{r}
#| include: false
ggsave("adjusting-class-threshold-by-strata2-thumb.png", g, width = 4, height = 3, dpi = 150)
```
@fig-thresholds shows why: group c's distribution sits well to the right of group a's, so a single threshold necessarily fits one of them worse than the other. The group-specific thresholds shift to follow each group's distribution, while still being pulled somewhat toward the black line.
## Fairness and bias in classification
Imagine that the strata in this simulation were gender, race, or some other demographic variable. Although the model itself accounted for the group variable — and so correctly captured the underlying prevalence differences — the additional step of **classifying** (turning a probability into a yes/no decision) can reintroduce bias if we don't also account for those prevalence differences when we set the threshold. A single threshold can look attractive because it optimizes overall accuracy, but "overall" can quietly average over very uneven performance across groups, as we saw above (a `r sprintf("%.1f", 100 * (range_single$max - range_single$min))`-point accuracy spread with a single threshold, versus `r sprintf("%.1f", 100 * (range_varying$max - range_varying$min))` points with varying thresholds). Whether that consistency is worth the small overall accuracy cost depends on the application — but it's a trade-off worth making deliberately, rather than by default.