Class imbalance in clinical data: what actually worked
In short: oversampling and class weights mostly move you along the same precision-recall curve. What actually worked was picking the decision threshold on purpose, on a validation set, against a…
- published
- read time
- 5 min
- words
- 1,011
- lang
- en
- filed under
- Engineering
In short: oversampling and class weights mostly move you along the same precision-recall curve. What actually worked was picking the decision threshold on purpose, on a validation set, against a precision the clinic could live with.
Almost every clinical dataset I've touched has been lopsided. The condition you care about shows up in a few percent of patients, or a few percent of recordings, and everything else is the healthy majority. Train a model on that and it learns, correctly, that saying "healthy" is a safe bet. Accuracy looks great. Recall on the cases that matter is poor.
The internet's answer is a list of tricks: oversample the minority, undersample the majority, generate synthetic cases, weight the loss. I've tried all of them on health data, at a hospital network and on voice recordings. Most of the time they were solving the wrong problem.
The default threshold is the real problem
A classifier gives you a score. Something turns that score into a yes or no, and in most libraries that something is score >= 0.5. Nobody chose 0.5. It's just what predict() does.
With five percent positives, a well calibrated model rarely gives any single patient a score above 0.5, so it rarely says yes. That's the low recall people blame on imbalance. Resampling and class weights "fix" it by pushing all the scores up, so more of them cross 0.5. You get more yeses. You also get a lot more false alarms, because you moved the whole model instead of moving the line.
Moving the line is cheaper, easier to explain, and you can change it later without retraining.
A toy you can run
I can't share clinical data, so here's the same effect on synthetic data with about five percent positives. Three logistic regressions: one plain, one trained on oversampled data, one with class_weight="balanced". Each is scored twice: at the default 0.5, and at a threshold picked on a separate validation set as the one with the best recall while precision stays at or above 0.5.
| Model | Threshold | Recall | Precision |
|---|---|---|---|
| plain | 0.5 | 0.33 | 0.68 |
| oversampled | 0.5 | 0.87 | 0.26 |
| class weights | 0.5 | 0.87 | 0.26 |
| plain | tuned | 0.59 | 0.50 |
| oversampled | tuned | 0.51 | 0.51 |
| class weights | tuned | 0.52 | 0.50 |
These numbers are the output of the script below on synthetic data, not results from any patient set. Look at the top three rows and resampling looks like magic: recall jumps from a third to most of the positives. Then look at precision. Three out of four alarms are false. In a clinic that means a clinician chasing three healthy people for every real case, and they will stop trusting the tool within a week.
The bottom three rows hold precision at the same level and compare recall fairly. The gap is gone. The plain model with a chosen threshold is as good as the others, a little better here, and it's the simplest thing in the table.
import numpy as np
from sklearn.datasets import make_classification
from sklearn.linear_model import LogisticRegression
from sklearn.model_selection import train_test_split
from sklearn.metrics import precision_recall_curve
X, y = make_classification(n_samples=20000, n_features=20, n_informative=6,
weights=[0.95], flip_y=0.01, random_state=0)
X_tr, X_te, y_tr, y_te = train_test_split(X, y, test_size=0.5, stratify=y, random_state=0)
X_tr, X_va, y_tr, y_va = train_test_split(X_tr, y_tr, test_size=0.5, stratify=y_tr, random_state=0)
def pick_threshold(y_true, scores, min_precision):
p, r, t = precision_recall_curve(y_true, scores)
ok = p[:-1] >= min_precision
return t[ok][np.argmax(r[:-1][ok])] # best recall that keeps precision
model = LogisticRegression(max_iter=1000).fit(X_tr, y_tr)
threshold = pick_threshold(y_va, model.predict_proba(X_va)[:, 1], min_precision=0.5)
pred = model.predict_proba(X_te)[:, 1] >= threshold
The oversampled and weighted variants in the table use the same split; oversampling is plain duplication of positives in the training set, nothing fancier.
Reading the curve
The picture explains the table better than the table does.
Plot precision against recall for the plain model and the weighted one and the curves nearly overlap. The weighted model isn't a better model. It's the same ranking of patients with a different default cut. If your ranking is good, you can stand anywhere on that curve by moving the threshold. If your ranking is bad, no resampling will save it.
The flat line near the bottom is the prevalence. A model that guesses at random sits there. It's worth drawing every time, because it reminds you how low the floor is and how much of the area above it you've actually earned.
What I do now, in order
- Split by patient first. Before any of this, make sure no person is in both train and test. Imbalance tricks on a leaky split just produce confident nonsense.
- Report the PR curve, not accuracy. Average precision and the curve itself. Accuracy at five percent prevalence tells you almost nothing.
- Ask the clinic what a false alarm costs. That gives you a minimum precision, or a maximum alarms per week. It's a clinical decision, not a modelling one.
- Pick the threshold on validation data to hit that number, then check it once on the test set. Write the threshold down next to the model version.
- Only then try weights or resampling, and judge them at the same precision, never at 0.5.
When the tricks do earn their place
I don't throw them away. Class weights help when the positives are so rare that the model barely sees them during training, and the ranking itself gets better, which you'll see as a curve that actually sits higher. They also help with neural networks, where a mini-batch might hold zero positives. Undersampling the majority is a fine way to cut training time on a large dataset. Synthetic oversampling has rarely helped me on clinical features, which tend to be messy, correlated and full of missing values.
The test is always the same. Does the curve move up, or did the point just slide along it?
predict() with predict_proba() and a threshold picked on validation data at the precision your users accept. Compare recall to your resampled model at that same precision. If the gap is small, delete the resampling step.related