返回专辑
·Johan·17 分钟阅读

Focal Loss 与类别不平衡:gamma 调大之后易样本去哪了

从调制因子与梯度贡献出发,说明 gamma/alpha 如何重分配难易样本;对比过采样、类别权重与 focal,用 PR 曲线、混淆矩阵与梯度直方图验收。

Focal Loss 与类别不平衡:gamma 调大之后易样本去哪了

1. 问题不只是少数类样本少,而是梯度预算被谁拿走

一条表面正常的二分类训练曲线,可能掩盖这样的部署结果:正常样本数量远多于缺陷样本,验证损失持续下降,整体准确率也很高,但真正昂贵的漏检没有改善。继续增加训练轮数通常无济于事,因为优化器每一步看到的主要信号仍来自大量已经分对的正常样本。单个易样本的交叉熵很小,数量乘上去之后,总贡献却未必小。

这里有两个容易混在一起的问题。类别不平衡描述不同类别出现频率或业务代价不同;难易不平衡描述同一类别内,已经高置信分对的样本与边界、遮挡、微小目标等困难样本对优化的价值不同。类别权重主要回答“哪个类更重要”,Focal Loss 的调制因子主要回答“当前模型还不会哪些样本”。少数类不必然困难,多数类也不必然容易:一批模糊的多数类负样本完全可能比清晰的少数类正样本更难。

因此,不能把 focal 理解成“给少数类加权”的另一个名字。它根据模型当前给真实类别的概率动态改权,同一个样本会随着参数更新从困难样本变成易样本,其权重也随之衰减。gamma 调大以后,易样本不是从数据加载器里消失了,也不是前向传播被跳过了;它们仍参与计算,只是传到参数更新中的梯度迅速接近零。问题的关键正是:衰减有多快,释放出的梯度预算交给了谁,以及这些被强调的“困难样本”究竟是有效边界样本还是脏标签。

2. 从交叉熵到 focal:先统一 的含义

对多分类任务,令模型 softmax 后第 类概率为 ,真实类别为 ,则

二分类也可用同一记号:正样本时 ,负样本时 。普通交叉熵和 Focal Loss 分别为

其中, 是按难易程度变化的调制因子, 是按类别选择的静态系数。gamma 为零时调制因子恒为一;若同时不使用 alpha,focal 应在数值上退化成普通交叉熵。这是实现必须通过的单元测试,不只是理论备注。

大表示模型已经给真实类较高置信度,通常被称为易样本; 小表示真实类概率低,通常被称为困难样本。这个定义依赖当前模型,而不是数据集里预先写死的标签。训练初期,大多数样本的 都不高;训练后期,边界清楚的样本逐渐右移到高 区域。focal 实际上构造了一个随训练状态变化的软式困难样本挖掘器。

还要注意,低 只能说明“模型目前不相信真实标签”,不能证明标签正确。标注错误、输入损坏、类别定义冲突、严重域外样本同样会得到很低的 。focal 会强调它们,所以困难度观测必须和数据质量观测放在一起。

3. gamma 调大后,易样本的损失权重如何衰减

只看调制因子

就能看到 gamma 的选择性。下面的数值由公式直接计算,不是某个数据集上的经验结果:

0.5010.50.250.03125
0.9010.10.010.00001
0.9910.010.00010.0000000001

时,从 调到 ,这个样本的交叉熵项会再乘一个千分之一。它仍经过骨干网络、仍占显存和计算时间,但在损失求和中的份额几乎被拿走。对数以万计的简单背景,这种衰减可以显著降低其合计影响;对样本本来就不多的分类任务,也可能过早丢掉维持类内结构所需的稳定信号。

gamma 并不设定一个明确的“难样本阈值”,而是连续改变权重曲线。调大 gamma 同时影响中等难度样本和极易样本,只是后者衰减更猛烈。把所有样本简单分成 hard/easy 两桶会丢掉重要信息,更适合按 分位数统计损失及梯度贡献,例如 [0, 0.2)[0.2, 0.5)[0.5, 0.8)[0.8, 0.95)[0.95, 1]

还有一个常见误判:训练损失变得更小,不代表分类器变得更好。gamma 增大本身就会缩小大量样本的损失尺度,两个配置的绝对 loss 不可直接横向比较。应比较相同验证集上的概率排序、混淆结构和分组梯度,而不是把“focal loss 数值更低”当成收益。

4. 调制的是损失,但真正驱动更新的是梯度

只把 当成“交叉熵梯度前的权重”并不精确,因为这个权重本身也依赖模型输出,反向传播还会对它求导。设二分类真实标签的符号为 ,logit 为 ,并写成

对 focal 关于 logit 求导可得

于是梯度幅值为

普通交叉熵的对应梯度幅值是 。两者之比为

这个比值揭示了两个重要事实。第一,focal 的梯度变化不只有表面上的幂函数,括号里的导数项会额外修正中等难度样本。以 、暂不计 为例, 的样本相对普通交叉熵甚至可能略被放大,而 的样本则被强烈压低。第二,当 非常接近一,令 ,有

而交叉熵梯度约为 。也就是说,易样本梯度相对交叉熵大致按 消失。所谓“易样本去哪了”,严格答案是:它们没有离开 batch,而是在反向图中变成了幅值极小的更新信号。

极难样本满足 时,调制因子趋近一,梯度幅值趋近 。focal 不会让无穷困难样本的 logit 梯度无限爆炸,但会改变它们相对其他样本的占比。如果一个 batch 里只有少数错标样本长期处于低 ,其他样本又逐渐被压低,那么参数更新方向仍可能被这些错标点控制。

5. alpha 与 gamma 分工不同,不能用一个补另一个

在二分类原始写法中,常见约定是

因此,代码里无条件给所有样本乘同一个 alpha=0.25 并没有平衡正负类,只是把整体损失和有效学习率缩小到四分之一。多分类场景则通常传入长度为类别数的向量 ,按真实类索引取得 。必须在配置和实验记录里写清是哪一种语义。

alpha 的作用是类级别、静态的:在整个训练过程中,同一类别共享一个系数。它可表达类别频率、漏检代价或二者折中。gamma 的作用是样本级别、动态的:同一类别中的易样本被压低,难样本被保留。同样是少数类,清晰正样本的 focal 权重会下降,边界正样本的权重仍较高;同样是多数类,难负样本也可能得到很大贡献。

两者相乘后,重分配会叠加。若少数类先按逆频率得到很大的 ,又因为模型暂时学不会而长期处于低 ,它会同时获得类别放大和难度保留。这个结果有时符合漏检代价,有时会造成高方差和大量误报。不能因为一个参数叫“平衡因子”、另一个叫“聚焦参数”就认为组合天然安全。

alpha 还涉及归一化契约。将逆频率权重归一到均值为一,和直接使用原始逆频率,类别间相对比例相同,但总体梯度尺度不同。PyTorch 的带权交叉熵在 mean 约简下有自己的分母语义,自写 focal 若直接 loss.mean(),不能假定两者尺度相同。公平消融应固定约简定义,记录每步梯度范数,必要时重新选择学习率,而不是仅替换损失函数名。

6. 过采样、类别权重与 focal 改的是三件不同的事

三种方法都能增加少数类对更新的影响,但作用位置不同。

过采样改变训练时看到的数据分布。少数类样本被更频繁抽中,其增广也有机会产生不同视图;与此同时,重复样本增加过拟合风险,batch 内类别比例和 BatchNorm 统计也会改变。若一个少数类只有几十个独立主体,重复抽样不会创造新的主体多样性。

类别权重不改变样本进入 batch 的概率,而是在目标函数中按类别乘固定系数。它更接近代价敏感学习,不会直接改变 BatchNorm 输入分布,但大权重会增加 batch 间梯度方差。逆频率只是起点,不是业务代价的同义词;漏检与误报的真实成本最终仍要通过阈值和业务流程表达。

focal也不改变采样分布,它根据当前 调整单样本目标。模型变了,权重就变。它适合“大量已经学会的样本仍占据总梯度预算”的场景,而不能为少数类补充外观覆盖,更不能修复标签缺口。

在不考虑增广、归一化层和有限 batch 方差的理想条件下,按某类频率进行过采样与给该类等比例加权,可以得到相似的期望梯度;在真实训练中,它们并不等价。过采样改变了哪些样本共同组成一个 batch,也改变了随机增广次数;类别权重则可能让偶然进入 batch 的少数类样本产生很大的单步更新。focal 又进一步依赖模型状态,无法用一个固定采样概率严格替代。

所以应按单变量建立基线:普通交叉熵、温和类别权重、过采样、focal,然后再判断是否需要组合。若一开始同时启用强过采样、逆频率 alpha 和大 gamma,即使 PR 曲线改善,也无法判断收益来自何处;一旦训练震荡,更无法知道该撤掉哪一层重加权。

7. 一个数值稳定、语义明确的 PyTorch 实现

多分类实现应从 log_softmax 取得 ,不要先做 softmax,再对可能下溢为零的概率取对数。下面的函数支持多分类 alpha 向量,并明确采用“逐样本损失求算术平均”的约简语义:

python
from __future__ import annotations

import torch
import torch.nn.functional as F
from torch import Tensor


def multiclass_focal_loss(
    logits: Tensor,
    targets: Tensor,
    *,
    gamma: float = 2.0,
    alpha: Tensor | None = None,
    reduction: str = "mean",
) -> Tensor:
    if logits.ndim != 2:
        raise ValueError("logits must have shape [N, C]")
    if targets.shape != (logits.shape[0],):
        raise ValueError("targets must have shape [N]")
    if gamma < 0:
        raise ValueError("gamma must be non-negative")

    # 在自动混合精度下也用 FP32 计算 log-probability。
    log_probs = F.log_softmax(logits.float(), dim=1)
    log_pt = log_probs.gather(1, targets[:, None]).squeeze(1)
    pt = log_pt.exp()
    modulating = (1.0 - pt).clamp_min(0.0).pow(gamma)
    loss = -modulating * log_pt

    if alpha is not None:
        if alpha.ndim != 1 or alpha.numel() != logits.shape[1]:
            raise ValueError("alpha must have shape [C]")
        alpha_t = alpha.to(device=logits.device, dtype=loss.dtype)[targets]
        loss = alpha_t * loss

    if reduction == "none":
        return loss
    if reduction == "sum":
        return loss.sum()
    if reduction == "mean":
        return loss.mean()
    raise ValueError(f"unsupported reduction: {reduction}")

若 logits 来自分割模型,形状通常是 [N, C, H, W],需要先把像素维展平,同时正确处理 ignore_index;不能把通道维错误地摊入样本维。多标签任务的各类别彼此独立,应基于 binary_cross_entropy_with_logits 对每个标签计算 ,不能直接复用上面的 softmax 多分类版本。

至少应有下面两类回归测试:

python
torch.manual_seed(7)
logits = torch.randn(32, 5, requires_grad=True)
targets = torch.randint(0, 5, (32,))

got = multiclass_focal_loss(logits, targets, gamma=0.0)
expected = F.cross_entropy(logits.float(), targets)
torch.testing.assert_close(got, expected)

alpha = torch.ones(5)
got_alpha = multiclass_focal_loss(
    logits, targets, gamma=0.0, alpha=alpha
)
torch.testing.assert_close(got_alpha, expected)

第一项保证 gamma=0 的退化关系;第二项防止 alpha 的索引、设备或 dtype 处理出错。还应覆盖极端 logits、空掩码和半精度输入,确认损失与梯度都为有限值。若采用不同的加权平均分母,则测试应针对项目约定写出,而不是强行对齐 CrossEntropyLoss(weight=...) 的默认约简。

8. 梯度直方图要按类别和难度分组

总 loss 只能告诉我们目标函数的标量结果,看不出谁在推动更新。一个低成本诊断是记录每个样本从损失传到 logits 的梯度范数。由于逐样本损失彼此独立,对 loss_vec.sum() 求 logits 梯度后,第 行正好对应第 个样本:

python
def focal_batch_diagnostics(logits, targets, gamma, alpha=None):
    loss_vec = multiclass_focal_loss(
        logits,
        targets,
        gamma=gamma,
        alpha=alpha,
        reduction="none",
    )
    grad_logits, = torch.autograd.grad(
        loss_vec.sum(), logits, retain_graph=True
    )

    with torch.no_grad():
        pt = logits.float().softmax(dim=1).gather(
            1, targets[:, None]
        ).squeeze(1)
        grad_norm = grad_logits.float().norm(dim=1)

    return {
        "loss": loss_vec.detach(),
        "pt": pt,
        "grad_norm": grad_norm,
        "target": targets.detach(),
    }

日志侧应同时画三种分组。第一,按真实类别分组,观察少数类拿到的梯度总量是否增加;第二,按 区间分组,确认高置信易样本确实被压低,而中等难度样本没有一起归零;第三,按正确与错误预测分组,检查错误样本是否垄断更新。直方图最好使用对数横轴,并同时报告每组样本数、梯度范数中位数与范数总和。只画均值会让极少数大值掩盖主体,只画样本数又看不到总贡献。

logit 梯度是干净、便宜的代理量,但不等于每层参数梯度。样本对共享参数的梯度方向可能相互抵消, 含义不同。需要进一步定位时,可以在分类头最后一层上用小规模验证 batch 计算逐样本参数梯度,或按组分别反向得到梯度向量并测余弦相似度。没有必要在每个训练 step 对整个骨干做逐样本梯度,那会显著增加成本。

还要把直方图与样本标识连接起来。持续处于最低 、最高梯度分位的样本应可回查原图和标注。如果其中集中出现错标、裁剪失败或某台设备的坏帧,继续增大 gamma 只是在让优化器更努力地拟合数据管线问题。

9. 消融实验必须控制梯度尺度与数据分布

一套可解释的实验至少包含四条线:

  1. 普通交叉熵,原始训练分布;
  2. 类别权重交叉熵,仍用原始训练分布;
  3. 过采样配普通交叉熵;
  4. focal,先不加 alpha,再单独评估温和的 alpha

gamma 可以从零、一个中等值和一个较大值做粗粒度扫描,但不要把某篇论文的默认值当成跨任务常数。检测器中海量背景候选与图像级分类的样本结构不同,语义分割中的像素相关性又不同。真正要回答的是:在本任务的 分布上,某个 gamma 把多少梯度从哪些桶移到了哪些桶。

所有实验要固定训练/验证划分、增广、优化器、学习率调度、训练预算和随机种子集合。过采样实验只改变训练加载器,验证集必须保持接近部署的自然分布。类别权重和 focal 改变总体梯度尺度时,仅固定名义学习率还不够,应记录分类头及骨干的梯度范数、跳过更新次数和混合精度缩放器状态。如果某个配置频繁触发梯度裁剪,它已经改变了优化过程,不能再把结果简单归因于难易样本重分配。

每条线至少保存:按类别的样本数与有效抽样次数、 分布、分组损失、分组 logit 梯度直方图、验证集预测概率和 checkpoint。保存概率而不只保存最终标签,才能在同一批预测上重算 PR 曲线与不同阈值的混淆矩阵,也能避免为了换一个阈值重复推理。

若少数类验证样本很少,单次 recall 的跳动可能只对应一两个样本。应给关键指标做按独立主体重采样的置信区间,视频帧、同一工件切片等相关样本不能当成独立观测逐帧重采样。否则看似明显的差异可能只是划分或阈值波动。

10. PR 曲线回答排序,混淆矩阵回答操作点

类别严重不平衡时,ROC 曲线的大量真负样本可能让横轴看起来很漂亮,而误报的绝对数量仍不可接受。二分类和一对多评估应优先查看 Precision-Recall 曲线:

PR 曲线评估概率排序,但业务最终运行在某个阈值上。因此每个候选配置至少需要两种比较。第一,在相同阈值下画混淆矩阵,观察损失变化是否导致概率尺度整体漂移;第二,在相同约束下比较,例如固定 precision 下的 recall,或固定每千样本误报数下的召回。只比较各自“最佳 F1 阈值”容易掩盖阈值敏感性,也未必对应业务成本。

多分类混淆矩阵应同时保存原始计数和按真实类别行归一化的版本。原始计数能回答总体错误量,行归一化版本能看每类召回。若某个大 gamma 配置提高罕见类召回,却把相近的正常子类大量推成罕见类,PR 曲线和矩阵会共同暴露这种交换;只看 macro-F1 可能把代价平均掉。

focal 还可能改变概率校准。因为大量高置信易样本的梯度被压低,输出概率未必仍有良好的置信度含义。上线阈值依赖概率时,应补充可靠性图、期望校准误差或 Brier score,并在固定模型后用独立校准集做温度缩放。校准不能改变排序,因此不会修复差的 PR 曲线;它解决的是“分数能否解释为风险”和“阈值能否跨批次稳定复用”。

验收时还应按场景切片:目标尺寸、亮度、设备、地域、遮挡程度、采集批次。总体 PR 改善若完全来自容易场景,而关键困难场景没有改善,就不能证明 focal 达到了目的。反过来,若少数关键切片改善且总体指标略有波动,也需要用明确业务成本判断,而不是机械追求单个平均数。

11. gamma 过大的典型失败模式

第一类失败是脏标签接管梯度。随着正常样本变易,其贡献下降;错标样本始终保持低 ,相对占比越来越高。表现通常是训练后期高梯度样本集合高度稳定,回查后能发现标签冲突。处理方式是修标、降权或采用噪声鲁棒策略,不是继续增大 gamma

第二类失败是中等难度样本也被过早静音。这些样本可能提供类内形状、背景变化和校准所需的密集信号。大 gamma 让训练集中在很窄的一段边界附近,训练 loss 看起来很低,验证 PR 却没有改善,概率还更难校准。梯度直方图会显示高 桶接近零,中间桶也明显塌缩,只剩少量极难样本。

第三类失败是小 batch 下更新方差增大。当每批真正有贡献的样本只剩几个,不同 batch 的更新方向差异变大。总体梯度范数可能并不爆炸,却会呈现尖峰和长时间低值交替。此时可降低 gamma、增大有效 batch、采用更温和的类别系数,或确保 batch 构造包含足够的有效正负样本。

第四类失败是与其他困难样本机制重复。检测器可能已经使用正负样本匹配、候选筛选、在线困难样本挖掘或专门的质量分支;再叠加大 gamma 会让极少数候选获得过高份额。应把匹配后正负数量、候选 IoU/质量分布与 focal 梯度一起观察,而不是只在损失层调参。

第五类失败是任务形式用错。多标签任务误用 softmax focal 会让类别互相排斥;分割任务忘记忽略无效像素会把 padding 当成海量易背景;序列任务把所有 token 平均后,长序列仍可能支配样本级贡献。focal 公式正确不代表约简轴和掩码正确,数据张量的语义必须先核对。

最后,focal 解决不了支持集不足。少数类只有单一设备、单一角度或极少独立主体时,再精细的重加权也只是重复利用已有证据。若错误分析指向覆盖不足,补数据、修增广或重新定义类别边界通常比继续搜索 gamma 更有效。

12. 从“是否使用 focal”落到可审计的验收结论

一个可上线的结论不应是“gamma=2 的 mAP 更高”,而应能回答以下链条:原交叉熵下哪些类别、哪些 区间占据了梯度;focal 把贡献如何重新分配;新增贡献对应真实困难样本还是数据错误;排序能力、固定操作点错误量与概率校准分别发生了什么变化。

具体可按以下条件收口:

  • gamma=0 与未加权交叉熵数值、梯度对齐,极端 logits 下无非有限值;
  • 训练验证划分不变,过采样只作用于训练集,验证集保持部署分布;
  • 梯度直方图按类别、 区间和预测正误分组,并能回查最高梯度样本;
  • PR 曲线在关心的 recall 区域有稳定改善,而不只是某个阈值偶然更好;
  • 原始计数与行归一化混淆矩阵都满足误报、漏检约束;
  • 关键切片的改善可复现,少样本指标给出按独立主体计算的不确定性;
  • alphagamma、约简方式、采样器和学习率作为一个完整优化契约归档;
  • 若概率用于决策,单独检查校准并重新确定部署阈值。

最终选择往往不是“focal 永远优于类别权重”。若问题主要是业务代价不同,温和类别权重配合阈值选择更直接;若问题主要是少数类覆盖不足,过采样加有意义的增广更合适;若问题确实是海量已学会样本淹没了边界样本,focal 才准确命中机制。alpha 决定类别之间如何分预算,gamma 决定每个类别内部如何按当前难度分预算。

所以,调大 gamma 之后,易样本并没有神秘消失。它们仍在前向路径中提供预测,也仍被统计为样本,只是其梯度按高阶速度衰减。真正需要验收的不是“易样本被压下去了没有”,而是空出的优化预算是否流向了值得学习、能够泛化、符合业务代价的困难样本。只有 PR 曲线、混淆矩阵和分组梯度证据指向同一个结论,focal 的收益才算被解释清楚。

← 全部文章

johan's blog