拓冰建站拓冰建站
首页 / 资讯中心 / 正文

类别不平衡分类实战:指标选择、Focal Loss与阈值校准全解析

1. 失衡问题为什么棘手梯度主导、先验偏移与评估陷阱1.1 一个batch算下来负样本把梯度吃掉了先说个最直观的现象。深度学习分类任务里的交叉熵损失默认对每个样本一视同仁不管它是多数类还是少数类进loss的权重都一样。这听起来公平但一旦正负样本比例失衡这种“公平”就变成了灾难。假设一个batch里有128个样本其中正样本1个负样本127个。计算平均loss时这127个负样本的loss占了绝对大头梯度方向基本上由负样本决定。模型学到的特征会变成只要把概率压低loss就能降下来。于是sigmoid输出不断向0偏移最终模型对所有输入都输出接近0的概率。这在逻辑上完全正确因为多数类就是那么多但在业务上等于白训。更隐蔽的还有最后一层偏置项的问题。二分类模型的全连接层如果带了bias训练数据里负样本比例很高时这个bias会被推到很负的位置相当于把整个输出分布平移了。这也是为什么后面我会专门讲阈值校准很多时候模型不是完全没用而是概率分布整体偏移了把决策阈值从0.5往下调一调效果立竿见影。1.2 更隐蔽的问题是评估指标失真如果只用Accuracy来评估失衡场景下的数值简直漂亮得可怕。正负样本1:1000时模型全判负类Accuracy是99.9%可它什么都没学会。这类“全判负但准确率贼高”的模型在实际业务中连测试都过不了。ROC-AUC在极端不均衡下也经常给人错觉。AUC衡量的是排序能力负样本占绝大多数时把负样本排前面很容易正样本虽然排在后面但相对位置没那么差AUC照样能刷到0.9以上。可你要的是把正样本找出来AUC并没有告诉你决策边界附近发生了什么。相比之下PR曲线Precision-Recall曲线对正样本更敏感因为precision和recall都直接围绕正类计算负样本数量的变化不会稀释它们的分辨率。PR-AUC、F1、F2、Balanced Accuracy这些指标才是失衡任务里真正值得盯的东西。所以我的建议是处理正负样本比例失衡之前先把评估指标定下来把验证集的构造方式定下来。指标没定清楚后面做再多优化都是对着错误的方向使劲。2. 先定度量衡失衡分类该看哪些指标为什么我不看准确率2.1 用PR曲线和F-family指标替代Accuracy具体怎么选指标取决于你关心的到底是哪一类错误。二分类的混淆矩阵里有TP、FP、FN、TN四个格子。在正负样本失衡的场景下TN通常巨大而Accuracy把TN也算进去所以它会被淹没在“负类正确”的假象里。F1是precision和recall的调和平均适合两者同等重要的场景。F2给recall更高的权重适合漏检代价更高的场景比如工业缺陷检测、医疗筛查、风控反欺诈。如果误报代价极高而漏检可以容忍那就反过来选F0.5。写代码算这些指标很直接import numpy as np def binary_metrics(y_true, y_pred): tp np.sum((y_pred 1) (y_true 1)) fp np.sum((y_pred 1) (y_true 0)) fn np.sum((y_pred 0) (y_true 1)) tn np.sum((y_pred 0) (y_true 0)) accuracy (tp tn) / max(tp fp fn tn, 1) precision tp / max(tp fp, 1) recall tp / max(tp fn, 1) f1 2 * precision * recall / max(precision recall, 1e-9) f2 5 * precision * recall / max(4 * precision recall, 1e-9) balanced_acc (recall tn / max(tn fp, 1)) / 2 return { accuracy: accuracy, precision: precision, recall: recall, f1: f1, f2: f2, balanced_acc: balanced_acc, }Balanced Accuracy相当于把两个类别的recall分别算出来再取平均它不关心类别数量只关心“每个类各自被召回得怎么样”。当你说不清该用F1还是F2的时候先用Balanced Accuracy和PR-AUC一起看至少不会被Accuracy骗。2.2 验证集要保留原始分布并且分层抽样有个特别容易犯的错训练集做了过采样或SMOTE后忘了验证集和测试集也应该按原始分布来评估。模型在均衡化分布上学的知识最终要回到真实场景里去用如果在验证集上也做重采样那验证结果完全失真。正确做法是先按原始标签比例切分训练集、验证集、测试集并且用分层抽样保证每个集合里正负样本比例与真实分布接近。数据划分完成之后再单独对训练集做重采样、增强之类的操作验证集和测试集始终保持真实分布的原始状态。我用过一个很蠢但很有效的习惯拿到数据先画类别分布图用matplotlib看一眼每个类别到底有多少样本再画一个抽样后的混淆矩阵。如果连训练集里正样本只有几百张都没发现后面一切优化都无从谈起。3. 方案一数据层重采样与类平衡采样器3.1 过采样和欠采样为什么不能无脑用数据层面的重采样是处理失衡最直觉的思路。多数类太多就少用点少数类太少就多来点听起来没毛病实际落地却很容易翻车。随机欠采样是扔掉多数类样本缺点是大量信息被丢弃。比如文本分类里多数类可能包含很多有价值的语义模式删掉一半多数类模型就没见过那些模式上线时稍微换个写法就识别错了。随机过采样则是重复少数类样本相当于让模型反复看同一批少数类数据很快就会过拟合验证集上指标一般测试集上更差。所以现在工程里更常用的是“清洗式欠采样”和“生成式过采样”。清洗式欠采样的思路不是随机删而是找出那些边界上容易引起歧义的样本删掉比如Tomek Links和Edited Nearest Neighbors目的是让两个类别在决策边界上更干净。生成式过采样则是用SMOTE这类方法在特征空间里合成新的少数类样本避免简单的重复。3.2 PyTorch的WeightedRandomSampler最稳的起点如果你用的是PyTorch最省事的方案是给DataLoader加一个WeightedRandomSampler按每个样本的类别权重来采样。这样不需要真的修改数据集只控制采样概率就能让每个batch里少数类出现的频率明显提高。from torch.utils.data import DataLoader, WeightedRandomSampler import torch # train_labels 是训练集的标签数组形状为 (N,) labels torch.as_tensor(train_labels) class_counts torch.bincount(labels) # 每个样本的权重 该样本所属类别样本数的倒数 weights 1.0 / class_counts[labels].float() sampler WeightedRandomSampler( weights, num_sampleslen(labels), replacementTrue, ) train_loader DataLoader( train_dataset, batch_size32, samplersampler, # 注意这里不能同时设置 shuffleTrue )三个容易踩的坑用了sampler之后绝对不能再设shuffleTrue否则PyTorch直接报错。replacementTrue意味着同一个epoch里少数类样本会被重复抽到这相当于在采样层面做了过采样。总觉得重复太多的可以把num_samples设成原始样本数的一半减少同一epoch内的重复。权重不一定非要取倒数的极值。我自己试下来1 / class_count经常把少数类权重抬得过高导致模型在少数类上过度自信。更平滑的设置是sqrt(total / class_count)或者用类别占比的中间值效果往往更稳。3.3 SMOTE对表格特征有效但别用在图像原始像素上SMOTESynthetic Minority Over-sampling Technique的思路是在特征空间里找少数类样本的近邻然后在样本和近邻之间做线性插值生成新的少数类样本。它非常适配表格数据、低维特征但在图像领域要格外小心。直接在原始像素上做插值大概率会生成一些不像真实物体的图像模型学了反而更乱。给一个适用于表格数据的SMOTE函数import numpy as np from sklearn.neighbors import NearestNeighbors def smote_interpolate(X, y, minority_label1, k5, num_syntheticNone): X_min X[y minority_label] if num_synthetic is None: num_synthetic len(X_min) nn NearestNeighbors(n_neighborsk 1).fit(X_min) new_samples [] for _ in range(num_synthetic): i np.random.randint(len(X_min)) x X_min[i] indices nn.kneighbors(x.reshape(1, -1), return_distanceFalse)[0][1:] neighbor X_min[np.random.choice(indices)] lam np.random.uniform(0.1, 0.9) synth x lam * (neighbor - x) new_samples.append(synth) return np.vstack(new_samples)用的时候注意SMOTE之前要把数据标准化否则不同特征的尺度差异会让欧氏距离失真SMOTE只对训练集做验证集和测试集不能碰SMOTE生成的样本样本间相关性强容易让模型过拟合必要时要配合更强的正则化。如果你做的是NLPSMOTE在原始文本上没有意义但可以在句向量特征层做插值效果看任务不能保证一定提升。图像任务我建议少在像素层合成直接上类特异增强会更安全。4. 方案二损失函数重设计Focal Loss与OHEM4.1 Focal Loss的核心逻辑让难分类样本主导梯度Focal Loss最初是目标检测里为了解决前景背景极端失衡提出的后来被广泛用到各种分类任务里。它的公式长这样FL -alpha_t * (1 - pt)^gamma * log(pt)其中pt是模型对正确类别的预测概率。如果模型已经把这个样本分对了且概率很高pt接近1(1 - pt)^gamma接近0这个样本的loss被压得很低。反过来如果样本分错了pt很小(1 - pt)^gamma接近1loss保持较高水平。放在失衡场景下多数的负样本模型很快就能分对Focal Loss自动把它们“静音”让梯度更多流向少数类和难例。PyTorch的实现很简单import torch import torch.nn as nn import torch.nn.functional as F class FocalLoss(nn.Module): def __init__(self, alphaNone, gamma2.0, reductionmean): super().__init__() self.alpha alpha self.gamma gamma self.reduction reduction def forward(self, logits, targets): ce F.cross_entropy(logits, targets, reductionnone) pt torch.exp(-ce) focal (1 - pt) ** self.gamma * ce if self.alpha is not None: alpha_t self.alpha[targets] focal alpha_t * focal if self.reduction mean: return focal.mean() return focal.sum()参数经验gamma从1.0开始试一般2.0是常见值但有些人任务里gamma1.0更好因为太难样本也可能是噪声压得太狠容易把噪声也当宝贝。alpha如果是类别频率倒数经常太极端我会先用sqrt(neg_count / pos_count)这种温和的值再逐步微调。另外Focal Loss和常规CE Loss对学习率敏感度不同用Focal Loss时如果发现训练不稳先调小学习率别一上来就堆别的技巧。4.2 OHEM在线困难样本挖掘OHEMOnline Hard Example Mining思路也很朴素计算一个batch里所有样本的loss只挑loss最大的那top-k个回传梯度其余样本当背景。这样模型不会把精力浪费在一堆已经学得很好的多数类样本上。def ohem_loss(logits, targets, keep_ratio0.3): loss F.cross_entropy(logits, targets, reductionnone) k max(int(loss.numel() * keep_ratio), 1) topk_loss, _ torch.topk(loss, k) return topk_loss.mean()但OHEM在极端失衡场景有一个明显缺陷如果一个batch里绝大多数都是负样本top-k挑出来可能全是难负样本正样本根本没机会被选进来学习。所以实战里OHEM更常和采样策略结合先把batch里的正负样本控制在一个合理比例比如1:3再做困难样本挑选这样既保证了正样本参与训练又让模型重点啃难例。对比来看Focal Loss适合“全局平滑地压低易分样本”OHEM适合“直接砍掉多数类里太容易的样本”。两者可以叠加但要小心过拟合难例我自己的体会是先用Focal Loss跑通再加OHEM看收益不要一上来全上。5. 方案三阈值校准与代价敏感决策真正的性价比之王5.1 0.5不是真理别把默认阈值当金科玉律很多人在二分类任务里模型输出概率后直接拿0.5当门槛大于0.5判正类小于0.5判负类。但0.5只是数学上的中立值和业务目标没有任何关系。正负样本失衡时模型输出的概率分布整体向0偏最佳决策边界往往在0.1到0.3之间有时候甚至在0.05附近。光调阈值能带来多大收益我举一个真实例子。某个风控模型默认0.5阈值下F1只有0.31我把验证集上所有样本的输出概率按不同阈值扫了一遍在0.12处F1到了0.67翻了一倍还多。模型一丁点没改只是把决策门槛挪了挪。这个方案投入产出比实在太高却经常被忽略。5.2 阈值搜索代码在验证集上搜索最优阈值是标准做法用验证集而不是测试集避免对测试集过拟合。import numpy as np # valid_probs: 模型在验证集上的正类概率 # valid_labels: 验证集真实标签 def search_threshold(valid_probs, valid_labels, beta2): best_th 0.5 best_score -1.0 for th in np.arange(0.01, 0.99, 0.01): pred (valid_probs th).astype(int) tp np.sum((pred 1) (valid_labels 1)) fp np.sum((pred 1) (valid_labels 0)) fn np.sum((pred 0) (valid_labels 1)) precision tp / max(tp fp, 1) recall tp / max(tp fn, 1) # F_beta: beta 1 时更看重 recall score (1 beta**2) * precision * recall / max(beta**2 * precision recall, 1e-9) if score best_score: best_score score best_th th return best_th, best_score这个搜索逻辑本身很简单关键在设定beta。默认0.5阈值往往偏保守导致recall极低搜索阈值后你会发现阈值可以大幅下调模型依然能保持可以接受的precision。5.3 用期望成本替代单一F值业务场景里漏检和误报的代价往往不一样。工业质检中漏掉一个缺陷品可能造成批量退货而误报一个良品只是多一次人工复检。这种情况下光优化F2还不够应该直接把代价矩阵放进阈值选择里。假设漏检代价是200误检代价是1那么可以把分类结果折算成总代价cost_fn 200.0 cost_fp 1.0 min_cost float(inf) best_th 0.5 for th in np.arange(0.01, 0.99, 0.01): pred (valid_probs th).astype(int) fn np.sum((pred 0) (valid_labels 1)) fp np.sum((pred 1) (valid_labels 0)) cost fn * cost_fn fp * cost_fp if cost min_cost: min_cost cost best_th th这样选出来的阈值直接对应业务上的最小化代价比机械地追求F1最高点更合理。在做这个之前先和业务方确认清楚漏检一次损失多少误检一次损失多少。这个沟通本身就价值很大。6. 方案四两阶段训练与难例调度6.1 两阶段训练为什么有效很多失衡分类任务的最佳实践并不是从头到尾用同一种采样策略而是分阶段来。第一个阶段让模型在类平衡或重采样的数据上充分训练先把少数类的特征学好第二个阶段切回原始数据分布用较小的学习率微调让模型重新校准真实场景里的先验概率。第一阶段可以理解成“配眼镜”让模型看清少数类到底长什么样第二阶段是“回到真实世界”适应真实分布下负样本占绝大多数的情况。只在重采样数据上训练模型的概率输出会偏向少数类只在原始数据上训练模型又对少数类“视而不见”。两阶段结合既保证了特征学习又保留了分布校准。这个套路在目标检测里很常见比如先用平衡的采样策略训练RPN再在全量数据上微调。分类任务同样适用。6.2 一个可落地的epoch调度模板用PyTorch框架来说训练循环可以按epoch分段切换total_epochs 30 warmup_epochs 15 finetune_epochs 15 for epoch in range(total_epochs): if epoch warmup_epochs: # 阶段一类平衡采样 loader balanced_train_loader optimizer.param_groups[0][lr] base_lr else: # 阶段二原始分布微调 loader original_train_loader optimizer.param_groups[0][lr] base_lr * 0.1 for batch in loader: # 常规训练逻辑 ...需要注意两个细节。第一阶段二的学习率必须显著降低否则会把第一阶段学到的少数类特征打乱我自己常用的比例是1/10。第二阶段二刚开始时验证集指标大概率会掉一点因为模型突然重新面对大量负样本可能短暂地“懵”一会儿。不要急着回滚给它几个epoch的时间恢复往往后半段会回升并超过第一阶段的峰值。6.3 从易到难的课程式难例调度除了训练流程分两段还可以在整个训练过程中动态调节“难例比例”。举个例子OHEM的keep_ratio可以设计成一个随epoch衰减的曲线前期保留80%的样本做训练后期只保留20%的困难样本。相当于是让模型先练基础题再集中刷难题这叫课程学习Curriculum Learning的思路。keep_ratio 0.8 * (1 - epoch / total_epochs) 0.2不过要提醒一句难例挖掘在标签噪声较多的数据集上非常危险。因为标签标错的样本loss通常也很大OHEM或Focal Loss会把它们当作难例重点学习模型越训越错。如果怀疑数据里有噪声先把清洗和纠错做在前面再考虑难例挖掘。7. 方案五类特异增强与生成式补充7.1 少数类专属增强增强要“偏心”数据增强是图像分类的标配但常规做法是对所有类别做同样的增强。失衡场景下更好的做法是给不同类别分配不同的增强强度少数类用更强的增强多数类用更弱的增强甚至不用增强。原因很简单增强的本质是一种先验知识注入它告诉模型“这些变换不影响语义”。对少数类注入更强的变换等于在有限样本上制造更多合法变体降低过拟合。比如工业质检里缺陷样本只有几百张可以给缺陷类多做旋转、缩放、颜色扰动正常样本有几万张只需轻微裁剪。PyTorch里可以用一个简单的条件判断实现def get_transform(label, is_trainTrue): if not is_train: return base_transform if label 1: # 少数类强增强 return strong_augment_transform else: # 多数类普通增强 return base_augment_transform需要注意这只是一种工程技巧不能替代对数据分布的理解。强增强虽然能扩展样本数量但如果增强变换让少数类样本偏离真实分布比如质检图像被过度旋转导致缺陷形态失真反而会伤害模型。7.2 Mixup/CutMix在类别失衡时的正确用法Mixup的思路是把两个样本按比例混合标签也按同样的比例混合。CutMix在图像上更进一步把一张图的某块区域贴到另一张图上。默认对所有样本做Mixup时多数类因为数量多混合后模型看到的仍然多数类信息占主导。想让Mixup对少数类更友好可以先把少数类在一个batch内的占比提上去再合成Mixup样本这样合成样本的标签分布更均衡。def mixup_batch(x, y, alpha0.2): lam np.random.beta(alpha, alpha) idx torch.randperm(x.size(0)) mixed_x lam * x (1 - lam) * x[idx] mixed_y lam * y (1 - lam) * y[idx] # 软标签 return mixed_x, mixed_y这段话可能有人觉得不新鲜但实战中的差异往往来自细节先重整batch里的类比例再做Mixup比直接全batch Mixup在少数类上的recall要好不少。还有一点Mixup训练时模型看到的是软标签验证推理阶段不需要Mixup直接用普通前向即可否则指标会被带偏。7.3 GAN和扩散模型生成样本听起来很美落地要谨慎这两年用扩散模型生成训练数据很火看上去是解决样本稀缺的终极方案生成几万张少数类样本训练集瞬间平衡。我也跟风做过一次在某个缺陷检测项目里用扩散模型生成了几千张缺陷图期望模型recall大幅提升结果测试集上的AUC没有任何提升有些指标还变差了。原因很直接生成数据的分布和真实缺陷分布并不完全一致模型学到了一些“生成痕迹”比如背景纹理和缺陷形态的虚假相关性一到真实环境就露馅。生成数据更适合用来做预训练、做负例补充、或者辅助难例的多样性分析而不是直接当正样本灌进训练集。如果你手上少数类样本少到只有几十张生成模型也救不了太多因为它同样需要足够多的真实样本才能学到可用的条件分布。我现在的优先级是先把采样、损失函数、阈值这些方案做扎实再考虑生成式补充。生成样本可以作为最后一公里锦上添花而不是雪中送炭的主力。8. 选型路线图不同失衡程度下怎么做组合8.1 先看失衡比例和业务约束不是所有失衡都需要处理。正负样本1:2到1:10之间常规训练加上合理的评估指标就够了额外引入过采样反而可能过拟合。到了1:10到1:100需要用重采样或类别权重帮助模型见到足够多的少数类。到了1:100到1:1000单靠采样不够了要上Focal Loss、阈值校准可能还要配合两阶段训练。到了1:10000以上先别急着训练考虑任务是否应该重新定义比如改成异常检测、排序问题或者召回K评估。失衡程度建议起点进阶方案备选思路1:1 ~ 1:10标准训练 CE Loss基本不用特殊处理关注指标选择即可1:10 ~ 1:100WeightedRandomSampler / 类别权重阈值校准SMOTE、Focal Loss1:100 ~ 1:1000加权采样 Focal Loss两阶段训练Mixup、难例调度1:1000以上异常检测 / 排序建模重新定义任务召回K、PR-AUC8.2 我推荐的优先级顺序这套方法论我使用过很多轮如果非要说一个“套路”那就是按代价从低到高逐个叠加先把验证集口径和指标定义清楚分层抽样保留真实分布。用WeightedRandomSampler或类别权重做一个加权CE Loss的baseline。在验证集上做阈值搜索看F2或业务代价能不能明显提升。如果还不行换Focal Loss调整gamma和alpha。再不行试两阶段训练先均衡再微调。最后才考虑生成样本和复杂增强。施行过程中每次只加一个改动并做消融对比。这个经验我从多次实战里总结出来最多的时候我把所有技巧叠到一起验证集F1确实高得离谱但换了一个地区的数据后直接崩盘原因是模型对少数类里一小批增强样本过拟合了反而丢了泛化能力。后来改成一次只动一个变量才发现收益最大的往往是最不起眼的阈值搜索和验证集修正Focal Loss和生成样本只是锦上添花。现在写代码之前我也会提醒自己先花半小时把数据分布和评估指标看清楚。很多失衡问题解决不了真不是缺技巧而是从一开始就没想清楚“什么样的结果才叫好”。把这个问题想明白方案自然就定下来了。
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门