Ir para o conteúdo

Class Imbalance

Desbalanceamento de Classes

Fraude é 0,1% das transações. Doença é 2% das triagens. Falha de equipamento é um punhado de horas num ano de logs. Desbalanceamento é a condição normal de qualquer problema que valha a pena resolver, porque a classe interessante é a rara — em geral é por isso que ela é interessante.

O conselho padrão é consertar isso: reamostrar, repesar, gerar pontos sintéticos da minoria. A maior parte desse conselho mira a coisa errada e o laboratório desta página mede o quão pouco ele ajuda.

O que está errado de verdade

Desbalanceamento, por si só, não é um defeito nos dados. Duas outras coisas são e levam a culpa por ele:

  1. A métrica está errada. Acurácia com 2% de positivos é um relatório sobre a classe majoritária. Ela tem de sair e o que a substitui depende do que você está fazendo.
  2. O limiar está errado. Um modelo devolve uma pontuação; transformar essa pontuação em decisão exige um corte e 0,5 é um padrão, não uma resposta. O corte certo vem do que cada tipo de erro custa.

Conserte esses dois e em geral sobra muito pouco para a reamostragem fazer — e reamostrar tem um custo próprio: destrói o sentido das probabilidades que o modelo devolve.


1. A métrica

A primeira tabela do laboratório defende o caso melhor que qualquer argumento. Um modelo, um conjunto de pontuações, quatro misturas de classe diferentes:

prevalence majority_class_acc roc_auc avg_precision
0.300 0.700 0.968 0.934
0.100 0.900 0.969 0.836
0.020 0.980 0.972 0.652
0.005 0.995 0.967 0.457
"""One model, one set of scores, four different class mixes.

The model is trained once and never touched again. All that changes is how many
negatives are kept in the evaluation set — so the model's ability to rank is
identical in every row, by construction.

Watch the three metrics disagree about what happened:
  majority_acc    what you get by predicting "no" every time — it goes UP
  roc_auc         flat, because it is a property of the ranking alone
  avg_precision   falls by half, because finding the positives really does get
                  harder as they get rarer

Printed as a markdown table, in identifiers only, so one artifact serves both
the English and the Portuguese page.
"""

import numpy as np
from sklearn.linear_model import LogisticRegression
from sklearn.metrics import average_precision_score, roc_auc_score
from sklearn.pipeline import make_pipeline
from sklearn.preprocessing import StandardScaler

N, DIM, TRAIN = 60_000, 10, 20_000

rng = np.random.default_rng(11)
direction = rng.normal(size=DIM)
X = rng.normal(size=(N, DIM))
risk = X @ direction + 1.0 * rng.normal(size=N)
y = (risk > np.quantile(risk, 0.70)).astype(int)

fitted = make_pipeline(StandardScaler(), LogisticRegression(max_iter=2000)).fit(X[:TRAIN], y[:TRAIN])
scores = fitted.predict_proba(X[TRAIN:])[:, 1]
held = y[TRAIN:]

positives, negatives = np.flatnonzero(held == 1), np.flatnonzero(held == 0)
draw = np.random.default_rng(0)

print("| `prevalence` | `majority_class_acc` | `roc_auc` | `avg_precision` |")
print("|---:|---:|---:|---:|")
for target in (0.30, 0.10, 0.02, 0.005):
    keep = min(len(positives), int(target * len(negatives) / (1 - target)))
    rows = np.concatenate([draw.choice(positives, keep, replace=False), negatives])
    subset, p = held[rows], scores[rows]
    print(f"| {subset.mean():.3f} | {max(subset.mean(), 1 - subset.mean()):.3f} "
          f"| {roc_auc_score(subset, p):.3f} | **{average_precision_score(subset, p):.3f}** |")

O modelo é treinado uma vez e nunca mais tocado. Conforme os positivos ficam mais raros:

  • A acurácia do preditor trivial sobe, de 0,700 para 0,995. A métrica diz que o problema ficou mais fácil. Não ficou; a métrica ficou sem sentido.
  • A ROC-AUC fica parada, 0,968 a 0,967. É uma propriedade só do ordenamento e portanto cega para o quão raros são os positivos. Útil, mas não vai avisar que a sua tarefa ficou difícil.
  • A precisão média cai pela metade, 0,934 para 0,457. Foi a que percebeu. Encontrar os positivos de fato ficou mais difícil, porque, para o mesmo recall, você agora atravessa muito mais negativos.
Métrica O que responde Use quando
Acurácia com que frequência acerto? as classes são balanceadas e os erros custam o mesmo
Precisão do que sinalizei, quanto era real? alarmes falsos são caros
Recall / sensibilidade do que era real, quanto peguei? casos perdidos são caros
F1 um meio-termo entre as duas você precisa de um número e os custos são mais ou menos simétricos
Precisão média (PR-AUC) sobre todos os limiares, quão bom é o ordenamento dos positivos? o resumo padrão sob desbalanceamento
ROC-AUC quão bem os positivos são ordenados acima dos negativos? comparar modelos; ciente de que ela é insensível à prevalência
MCC um resumo balanceado da matriz de confusão inteira binário, muito desbalanceado, sem forte assimetria de custo

A ROC-AUC lisonjeia um problema desbalanceado

Com 99% de negativos, um número absoluto grande de falsos positivos é uma taxa de falsos positivos pequena e a curva ROC é desenhada em taxas. É por isso que uma ROC-AUC de 0,99 pode conviver com uma precisão de 0,07, como acontece no painel abaixo. Reporte também a precisão média e sempre reporte a prevalência ao lado — uma AP de 0,45 é ruim com 30% de prevalência e excelente com 0,5%.


2. O limiar

Um classificador não devolve uma decisão. Devolve uma pontuação e alguém precisa escolher onde cortar. O predict() corta em 0,5 porque precisa cortar em algum lugar.

Ponha o painel em 2% de positivos e um caso perdido que custa dez alarmes falsos: o limiar mais barato é 0,85, não 0,50 — levá-lo até lá corta a conta mais ou menos pela metade. Torne o caso perdido cinquenta vezes mais caro e o melhor corte cai para 0,65. Torne-o barato e ele sobe para 0,97. Aumente a prevalência para 30% e o melhor corte cai para abaixo de 0,5.

O limiar não é propriedade do modelo. É onde a saída do modelo encontra a economia do seu problema e ele é livre para mudar — sem retreinar, sem reamostrar, sem dados novos.

from sklearn.metrics import precision_recall_curve

p = model.predict_proba(X_val)[:, 1]                 # na validação, nunca no teste
precision, recall, thresholds = precision_recall_curve(y_val, p)

# se você conhece os custos, minimize o custo diretamente
custo = lambda t: CUSTO_PERDA * ((p < t) & (y_val == 1)).sum() + ((p >= t) & (y_val == 0)).sum()
melhor = min(thresholds, key=custo)

# se não conhece, ao menos escolha pela restrição que você de fato tem
suficiente = recall >= 0.90                          # "precisamos pegar 90% das fraudes"
melhor = thresholds[np.argmax(precision[:-1][suficiente[:-1]])]

A pergunta que substitui 'como lido com desbalanceamento?'

Quanto custa um caso perdido, em relação a um alarme falso? Um tumor não detectado e um exame de acompanhamento desnecessário não são erros comparáveis e a razão entre eles — mesmo aproximada, mesmo em ordem de grandeza — determina o limiar e portanto o comportamento inteiro do sistema. Se ninguém sabe responder, esse é o achado e precisa ser resolvido antes de qualquer modelagem.


3. Reamostragem e pesos de classe

Agora as intervenções às quais todo mundo recorre primeiro.

  • Pesos de classe multiplicam a perda nos exemplos da minoria. No sklearn, class_weight='balanced'; no PyTorch, pos_weight no BCEWithLogitsLoss.
  • Sobreamostragem aleatória duplica linhas da minoria até as classes ficarem iguais.
  • SMOTE sintetiza pontos novos da minoria interpolando entre vizinhos, para evitar duplicatas literais.
  • Subamostragem joga fora linhas da maioria, que é a única que também torna o treino mais rápido.
strategy roc_auc avg_precision f1@0.5 f1@best brier mean_pred
baseline 0.993 0.761 0.683 0.697 0.0088 0.020
threshold_tuned 0.993 0.761 0.683 0.697 0.0088 0.020
class_weight_balanced 0.992 0.758 0.443 0.693 0.0360 0.075
random_oversample 0.992 0.758 0.441 0.688 0.0355 0.074
oversample_then_corrected 0.992 0.758 0.681 0.688 0.0090 0.022
"""Four ways to "handle imbalance", and what each one actually changes.

A 2% positive rate, one model family, four treatments. Read the table by
column rather than by row:

  roc_auc         barely moves — which is the first thing to notice about it
  avg_precision   barely moves either: none of these improves the ranking
  f1@0.5          moves a lot, and downward for the resampled variants
  f1@best         back to level once the threshold is tuned for each
  brier, mean_pred  wrecked by weighting and by oversampling, and restored by
                  the prior correction on the last row

Logistic regression rather than a network, because `class_weight` is available
and the run is deterministic. For a network the equivalent of class weights is
`pos_weight` in `BCEWithLogitsLoss`, with exactly the same consequence for the
probabilities it outputs.

Printed as a markdown table, in identifiers only, so one artifact serves both
the English and the Portuguese page.
"""

import numpy as np
from sklearn.linear_model import LogisticRegression
from sklearn.metrics import (average_precision_score, brier_score_loss,
                             f1_score, roc_auc_score)
from sklearn.model_selection import train_test_split
from sklearn.pipeline import make_pipeline
from sklearn.preprocessing import StandardScaler

N, DIM, POSITIVE = 20_000, 10, 0.02

rng = np.random.default_rng(11)
direction = rng.normal(size=DIM)
X = rng.normal(size=(N, DIM))
risk = X @ direction + 1.0 * rng.normal(size=N)
y = (risk > np.quantile(risk, 1 - POSITIVE)).astype(int)
X_train, X_test, y_train, y_test = train_test_split(
    X, y, test_size=0.3, stratify=y, random_state=0
)

model = lambda **kw: make_pipeline(StandardScaler(), LogisticRegression(max_iter=2000, **kw))


def row(name, p):
    grid = np.quantile(p, np.linspace(0.90, 0.9999, 300))
    best = max(f1_score(y_test, p >= t) for t in grid)
    print(f"| `{name}` | {roc_auc_score(y_test, p):.3f} | **{average_precision_score(y_test, p):.3f}** "
          f"| {f1_score(y_test, p >= 0.5):.3f} | **{best:.3f}** "
          f"| {brier_score_loss(y_test, p):.4f} | {p.mean():.3f} |")


print("| `strategy` | `roc_auc` | `avg_precision` | `f1@0.5` | `f1@best` | `brier` | `mean_pred` |")
print("|---|---:|---:|---:|---:|---:|---:|")

plain = model().fit(X_train, y_train).predict_proba(X_test)[:, 1]
row("baseline", plain)
row("threshold_tuned", plain)                 # the same model — only the cut moves

weighted = model(class_weight="balanced").fit(X_train, y_train)
row("class_weight_balanced", weighted.predict_proba(X_test)[:, 1])

minority, majority = np.flatnonzero(y_train == 1), np.flatnonzero(y_train == 0)
copies = np.random.default_rng(0).choice(minority, len(majority) - len(minority), replace=True)
index = np.concatenate([np.arange(len(y_train)), copies])
resampled = model().fit(X_train[index], y_train[index])
p_over = resampled.predict_proba(X_test)[:, 1]
row("random_oversample", p_over)

# Undo the base rate the resampling invented (Elkan / King & Zeng): shift the odds
# by the ratio of the true prior to the training prior.
prior_train, prior_true = y_train[index].mean(), y_train.mean()
odds = p_over / (1 - p_over) * (prior_true / (1 - prior_true)) * ((1 - prior_train) / prior_train)
row("oversample_then_corrected", odds / (1 + odds))

Leia essa tabela por coluna e ela diz algo que o conselho de sempre não diz.

Nada melhorou o ordenamento. A precisão média fica em 0,758–0,761 nas cinco linhas, a ROC-AUC em 0,992–0,993. Seja lá o que esses tratamentos fazem, eles não tornam o modelo melhor em distinguir positivos de negativos.

As diferenças de F1 são diferenças de limiar. No padrão 0,5, pesar e sobreamostrar parecem piores que a linha de base — 0,443 e 0,441 contra 0,683. No melhor limiar de cada modelo, empatam de novo: 0,697, 0,693, 0,688. Repesar não melhorou o modelo; moveu o ponto de operação, o que o limiar faz de graça.

E quebraram as probabilidades. O escore de Brier vai de 0,0088 para 0,0360, quatro vezes pior. A probabilidade média prevista vai de 0,020 — que é exatamente a prevalência verdadeira — para 0,075. O modelo agora acredita que positivos são quase quatro vezes mais comuns do que são, porque você lhe disse isso ao mostrar um conjunto de treino em que eram.

Reamostrar muda a taxa-base que o modelo aprende

Se qualquer coisa a jusante consome as suas probabilidades — valor esperado, um escore de risco, um cálculo de custo, um gráfico de calibração — um modelo reamostrado está mentindo para ela. A distorção é corrigível em forma fechada: desloque o log das chances pela razão entre a prior verdadeira e a prior de treino.

# correção de Elkan, desfazendo uma prevalência de treino pi_train
odds = p / (1 - p) * (pi_true / (1 - pi_true)) * ((1 - pi_train) / pi_train)
p_corrigido = odds / (1 + odds)

A última linha do laboratório aplica exatamente isso: Brier de volta a 0,0090, previsão média de volta a 0,022.

Então quando reamostrar vale a pena?

Há casos reais e eles são mais estreitos do que a reputação sugere:

  • A classe minoritária é tão rara que os lotes não contêm nenhum exemplo dela. Com 0,01% de positivos e lote de 256, a maioria dos lotes não carrega sinal algum. Amostragem balanceada de lote conserta um problema de gradiente, não um problema estatístico.
  • Subamostragem como decisão de computação. Descartar 90% da classe majoritária torna o treino dez vezes mais rápido a um custo pequeno em AP; é uma troca legítima mesmo sem comprar acurácia.
  • A própria perda é o problema. A focal loss reduz o peso dos exemplos fáceis e bem classificados — a esmagadora maioria dos negativos — de modo que o sinal de gradiente venha dos casos difíceis. É um mecanismo diferente de repesar por classe e é o que melhor se sustentou em aprendizado profundo.1

Onde quer que você reamostre, faça dentro da dobra de treino

Sobreamostrar antes da divisão põe cópias da mesma linha dos dois lados dela e o modelo passa a ser pontuado em linhas que memorizou. Medido no capítulo de vazamento: +0,173 de AUC de pura ficção, o maior dos quatro vazamentos daquela página.


4. Especificidades de aprendizado profundo

Para uma rede treinada num conjunto grande, a ordem prática é:

  1. Conserte a métrica. Reporte precisão média e a prevalência, mais aquela entre precisão e recall sobre a qual o seu problema realmente é.
  2. Ajuste o limiar na validação, a partir dos custos.
  3. Pese a perda se o gradiente for genuinamente dominado por negativos — pos_weight, ou focal loss. Lembre que isso descalibra e corrija depois, se alguma coisa consumir as probabilidades.
  4. Balanceie os lotes só se os lotes estiverem vindo sem positivo algum.
  5. Colete mais positivos. Quase sempre a ação de maior valor e quase sempre a que ninguém orça.
import torch

# pos_weight escala o termo positivo da perda; N_neg / N_pos é o ponto de partida usual
loss = torch.nn.BCEWithLogitsLoss(pos_weight=torch.tensor([n_neg / n_pos]))

# focal loss: reduzir o peso dos exemplos fáceis em vez de aumentar o de uma classe inteira
def focal(logits, target, gamma=2.0, alpha=0.25):
    bce = torch.nn.functional.binary_cross_entropy_with_logits(logits, target, reduction='none')
    pt = torch.exp(-bce)                       # a probabilidade atribuída à classe verdadeira
    return (alpha * (1 - pt) ** gamma * bce).mean()

O que os laboratórios ensinam

O reflexo é mudar os dados. As medições dizem para mudar primeiro o relatório e o corte: a métrica, porque acurácia e ROC-AUC vão as duas dizer que um problema desbalanceado vai bem; o limiar, porque é ali que os custos entram e é a única alavanca que move o resultado sem tocar no modelo. Reamostrar é um efeito de terceira ordem que chega com uma conta de calibração junto.



  1. Lin, T.-Y., Goyal, P., Girshick, R., He, K., Dollár, P. Focal Loss for Dense Object Detection, ICCV 2017. Ver também Chawla, N. V. et al. SMOTE: Synthetic Minority Over-sampling Technique, JAIR 2002. E Elkan, C. The Foundations of Cost-Sensitive Learning, IJCAI 2001, de onde vem a correção de prior acima. ↩