Adjusting classification threshold by strata

Predictive modeling
Classification
Author

Tyler Grimes

Published

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.

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

Code
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)
Code
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.
Table 1: Group sizes and observed outcome prevalence in the simulated data
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.

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.
Code
fit = glm(y ~ x + z, family = "binomial", data = df)
coef(fit)
(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.

Code
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

Code
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 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.
Code
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)
Table 2: 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.
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.

Visualizing the thresholds

Code
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
Figure 1: 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.

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

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

Back to top