Focal Loss深度解析:从交叉熵原理到PyTorch实现与调参实战
先聊点实际的Focal Loss这名字搞目标检测的人应该都不陌生。RetinaNet靠它一战成名YOLOv8的损失函数里也挂着它的影子甚至不少做长尾分类、分割任务的朋友都在用它。但真让你解释清楚它到底干了什么、为什么要用一个看起来有点绕的公式、参数γ和α到底怎么调很多人是说不透的。这篇文章我打算把Focal Loss从头到尾拆一遍。包括它要解决的痛点、公式每一步的含义、PyTorch怎么正确实现、在YOLO这类框架里是怎么用的以及我调参踩过的坑。内容会尽量详细数学只用到高中水平例子也用大白话讲目标是让你看完之后不光会用还能在面试、组会、技术评审的时候把这事讲明白。1. Focal Loss要解决的问题两个让人头疼的失衡先说一个我经常在训练检测模型时遇到的场景。一张街景图里可能有几十辆车但真正是行人的区域只有几百个像素。在目标检测里候选框一生成就是几万个其中绝大多数是背景。你用普通的交叉熵损失去训练模型很容易被大量的“背景框”带偏学到一堆“这个区域不是目标”的结论忽略真正难分的那些正样本。1.1 类别不平衡正样本太少了怎么办第一个失衡是类别不平衡。正负样本比例可能达到1:1000甚至更夸张。如果直接用交叉熵负样本背景在损失函数里贡献的梯度绝对量非常大模型优化方向基本被背景主导正样本那点信号完全被淹没了。常见思路是给正负样本加权重。比如正样本乘以权重α负样本乘以1-α这个思想后来在Focal Loss里被保留成一个可调参数。但你很快会发现单纯靠α只能把正负样本的量掰回来解决不了下一个问题。1.2 难易样本失衡大量“送分题”盖过了“难题”第二个失衡很多人会忽略但我觉得它才是Focal Loss真正想解决的——难易样本失衡。你可以把损失函数想象成一个老师批改作业。普通交叉熵对所有题目一视同仁超级简单的送分题模型预测概率0.99的负样本和极其刁钻的难题预测概率0.4但确实是目标的正样本对损失的贡献是按照预测值来的。送分题虽然单条损失小架不住数量巨大——几万个简单负样本累积起来的梯度照样能淹没几十个难正样本的贡献。更麻烦的是哪怕你把类别平衡因子α加上也只会让所有负样本统一降低权重但简单负样本和困难负样本之间依然没区分。模型花大量力气反复优化那些已经能很好分类的简单样本真正的边界样本反而学不到位。Focal Loss的核心思路就是解决这个让损失函数自动把重点放在难分样本上自动降低简单样本的权重。它不是通过重采样或OHEM那样专门挑样本而是在损失函数层面上用数学方式给每个样本算一个“困难度权重”。2. 从交叉熵到Focal Loss一个公式的演变过程理解Focal Loss最好的方式是把它看成交叉熵的一个“升级补丁”。如果你想透彻掌握它先要把交叉熵的底子打好。2.1 交叉熵是怎么算的问题出在哪以二分类为例标准交叉熵长这样[ CE(p, y) -[y \log(p) (1-y) \log(1-p)] ]这里 (y) 是真实标签1或0(p) 是模型预测为正类的概率。为了方便表达何恺明在论文里定义了一个变量 (p_t)当真实标签 (y1) 时(p_t p)当真实标签 (y0) 时(p_t 1-p)于是交叉熵就可以简写成[ CE(p_t) -\log(p_t) ]这个式子非常漂亮地统一了正负样本的写法。(p_t) 越大说明模型对这个样本分得越对对应的 (-\log(p_t)) 就越小(p_t) 越小说明模型分错了或者分得不确定损失就越大。问题出在哪假设模型对某个简单负样本输出 (p0.9)所有 (p_t) 都很高。对一批简单样本算梯度时每个样本虽然损失小但量大。累积起来简单样本的梯度贡献在总梯度中占比极高把困难的、信息量大的样本信号给稀释了。2.2 加一个调制因子Focal Loss的核心思想Focal Loss在交叉熵基础上引入了一个调制因子 ((1 - p_t)^\gamma)公式变成[ FL(p_t) -(1 - p_t)^\gamma \log(p_t) ]其中 (\gamma) 是聚焦参数论文里默认取2。这个调制因子的作用很直接当样本被分得很好、(p_t) 接近1的时候((1-p_t)^\gamma) 会变得非常小损失被压得很低当样本被分得很差、(p_t) 比较小的时候调制因子接近1损失基本保留。你可以具体感受一下数值变化。假设 (\gamma2)我列一个简单的对照表样本类型(p_t)交叉熵损失调制因子 ((1-p_t)^2)Focal Loss简单负样本0.90.1050.010.00105一般样本0.50.6930.250.173难分样本0.12.3030.811.865你会发现简单样本的损失被压缩到原来的1/100而难样本只被压缩到原来的81%左右。这样一比较难样本在总损失里的“话语权”就大大提升了。这个设计的本质相当于给每道题乘上一个“难度加权系数”。送分题权重趋近于0难题权重接近1。模型就能把更多精力放在真正需要优化的边界样本上。2.3 α平衡因子把正负样本的权重也一起调了只加调制因子依然有一个问题虽然简单负样本的权重被压低了但负样本数量实在太多累积起来还是可能压过正样本。于是论文里又加了一个 (\alpha) 平衡因子得到完整版[ FL(p_t) -\alpha_t (1 - p_t)^\gamma \log(p_t) ]这里的 (\alpha_t) 作用和加权交叉熵里的 (\alpha) 一样。当真实标签 (y1) 时(\alpha_t \alpha)当 (y0) 时(\alpha_t 1-\alpha)。(\alpha) 一般取0.25意思是正样本权重0.25负样本权重0.75。注意这不是说重视负样本而是负样本数量实在太多(\alpha) 会把正样本的相对重要性明显放大通常配合 (\gamma2) 使用效果比较好。理解了这两层后你会发现Focal Loss并不是什么高深的魔法。它就是在交叉熵外面套了两层权重一层的 (\alpha_t) 控制类别平衡一层 ((1-p_t)^\gamma) 控制难易样本平衡。两个机制互不冲突各管各的。3. 数学推导和代码实现PyTorch配套公式看懂了还不够实际写代码时有很多细节容易出错。我见过不少人直接把论文公式抄进PyTorch结果数值不稳定或梯度爆炸。这一节我把代码细节讲透。3.1 公式的每一项代表什么在写代码前先明确一下输入。通常我们用模型输出的 logits未经过sigmoid的值而不是概率 (p)。原因很简单logits配合BCEWithLogitsLoss或CrossEntropyLoss在数值上更稳定直接对概率做log很容易出现log(0)的问题。假设二分类场景预测概率 (p \sigma(z))(z) 是logits。那么[ pt p \cdot y (1-p) \cdot (1-y) ][ \log(pt) \log(p) \cdot y \log(1-p) \cdot (1-y) ]再根据标签算出对应的 (\alpha_t)然后套公式就行。3.2 从零实现一个正确的Focal Loss下面我给出一个在工程项目里验证过多次的PyTorch实现。它支持二分类和多分类两种模式也支持梯度裁剪前的稳定计算。import torch import torch.nn as nn import torch.nn.functional as F class FocalLoss(nn.Module): def __init__(self, alpha0.25, gamma2.0, reductionmean): super().__init__() self.alpha alpha self.gamma gamma self.reduction reduction def forward(self, logits, targets): # logits: [N, num_classes] 或 [N]targets: [N]类别索引 num_classes logits.size(-1) if num_classes 1: # 多分类 log_probs F.log_softmax(logits, dim-1) probs torch.exp(log_probs) # targets 转 one-hot targets_one_hot F.one_hot(targets, num_classesnum_classes).float() # 提取每个样本target类别的 log prob 和 prob log_pt (log_probs * targets_one_hot).sum(dim-1) pt (probs * targets_one_hot).sum(dim-1) else: # 二分类 logits logits.squeeze(-1) log_pt F.binary_cross_entropy_with_logits( logits, targets.float(), reductionnone ) probs torch.sigmoid(logits) pt probs * targets (1 - probs) * (1 - targets) focal_weight (1 - pt) ** self.gamma # alpha 处理 if num_classes 1: alpha_t targets_one_hot * self.alpha (1 - targets_one_hot) * (1 - self.alpha) alpha_t alpha_t.sum(dim-1) else: alpha_t targets * self.alpha (1 - targets) * (1 - self.alpha) loss alpha_t * focal_weight * log_pt if self.reduction mean: return loss.mean() elif self.reduction sum: return loss.sum() else: return loss有几个地方需要特别提醒如果num_classes是1且你传入的logits是[N, 1]一定要先squeeze(-1)否则binary_cross_entropy_with_logits的形状检查会报错。在二分类场景里pt算出来的是每个样本属于真实类别的预测概率用法和论文完全一致。log_pt由于用了binary_cross_entropy_with_logits内部会把logits和targets对齐不用手写sigmoid和log。我一般建议在项目里直接用这种实现而不是抄网上那些只针对二分类的简化版本。因为做检测任务时分类分支往往是多分类比如COCO的80类二分类版本容易在改造时引入bug。3.3 三种容易出现bug的细节epsilon、alpha的shape、reduction很多人写完Focal Loss测试的时候发现loss是NaN九成是这几个原因。第一个是log(0)问题。如果你直接用概率 (p) 去算 (log(p))而 (p) 是sigmoid输出在某些数值极端情况下可能等于0。处理办法是使用log_softmax或binary_cross_entropy_with_logits它们内部已经做了数值稳定处理。如果自己实现log(pt)记得加一个eps1e-7之类的极小值。第二个是alpha的形状。这是我最常帮人排查的问题。如果你的alpha是一个标量0.25而targets是一个batch的向量直接相乘不会有问题。但如果你用类别级别的alpha比如各类别频率不同alpha是一个长度为num_classes的向量就必须先通过targets索引取出每个样本对应的alpha再做乘法。否则形状对不上结果会完全错掉。第三个是reduction的选择。论文里的sum是常规做法但实际训练中我更推荐mean。原因很简单sum模式下loss的数值会随着batch size变化如果batch size设置大了loss直接翻倍学习率就得跟着调。而mean可以保持loss量级稳定看到曲线时更容易判断收敛状态。4. 实战在目标检测和YOLO系模型里怎么用Focal Loss最出名的应用场景就是目标检测。但你想在YOLO这类框架里用好它不能停在“loss函数换一下”这个层面。4.1 RetinaNet为什么靠它翻身在RetinaNet之前主流单阶段检测器比如SSD精度上不去最重要的原因就是正负样本极端不平衡。两阶段检测器靠RPN的候选框筛选机制把负样本控制在较小规模单阶段检测器面对的是密集网格上几万个预测框正样本占比微乎其微。何恺明团队提出RetinaNet时一个核心贡献就是证明了不需要复杂的两阶段筛选只要把分类损失从交叉熵换成Focal Loss单阶段模型就能超过两阶段精度。当时这个结论对检测领域影响很大因为它让大家意识到精度瓶颈不只在网络结构损失函数设计同样重要。在RetinaNet里Focal Loss被用在分类分支上回归分支仍然用Smooth L1 Loss。这一点很关键Focal Loss是专门为分类问题设计的不能随意套到回归任务上。4.2 YOLOv8等检测框架里的Focal Loss形态YOLOv8、YOLOv9这些现代检测器虽然对外宣称用的还是BCE Loss但很多实现里已经融合了Focal Loss的思路。比如分类分支直接用BCEWithLogitsLoss本质上就是 (\gamma0) 的特殊Focal Loss。而YOLOv8在DFLDistribution Focal Loss损失里也借鉴了Focal的思想对预测分布中概率较高的部分加强约束。如果你要手动在YOLO的自定义训练里接入Focal Loss一般改分类分支就行# 原来 loss_cls nn.BCEWithLogitsLoss()(pred_cls, target_cls) # 改成 loss_cls FocalLoss(alpha0.25, gamma2.0)(pred_cls, target_cls)但要注意YOLO的target_cls是0/1矩阵表示每个位置是否包含某类目标。上面我写的FocalLoss是硬标签版本如果你的标签是soft label比如用了标签平滑代码里的targets就不能直接用整数索引需要换成一维或二维的浮点标签格式。否则one-hot处理会出错。4.3 怎么把损失曲线画出来观察收敛效果很多人问YOLOv8怎么画损失函数曲线。其实原理很简单训练过程中把每个iteration或每个epoch的loss记下来最后用matplotlib画出来就行。我习惯记三个量分类损失、回归损失、总损失。import matplotlib.pyplot as plt # 假设 train_loss 是训练时每个 epoch 记录的列表 def smooth_curve(values, beta0.9): smoothed [] last values[0] for v in values: last beta * last (1 - beta) * v smoothed.append(last) return smoothed epochs range(1, len(train_loss) 1) plt.plot(epochs, smooth_curve(train_loss), labeltrain_loss) plt.plot(epochs, smooth_curve(val_loss), labelval_loss) plt.xlabel(epoch) plt.ylabel(loss) plt.legend() plt.grid(True) plt.savefig(loss_curve.png)观察曲线时有个容易忽略的细节Focal Loss的绝对数值比普通CE要小很多这是正常的因为简单样本的损失被压制了。你更应关注的是loss是否在持续下降以及val_loss和train_loss的差距。如果val_loss在某个epoch之后开始反弹说明过拟合了如果train_loss一直不降很可能γ太大梯度被压得过小模型学不动。5. 参数怎么调γ、α的选择和踩坑记录Focal Loss有两个超参数(\gamma) 和 (\alpha)。它们看着简单实际调起来比想象中麻烦因为二者是相互影响的。5.1 γ2和α0.25的经验从哪来论文给的默认值是 (\gamma2)、(\alpha0.25)。这个组合是在COCO数据集上大量实验调出来的不是拍脑袋定的。它的道理在于在 (\gamma2) 时简单样本的损失衰减已经非常厉害大到再增加 (\gamma)难样本的权重也会被压得过低导致整体学习速度变慢。(\alpha0.25) 则是在正负样本比例极端的情况下配合 (\gamma2) 找到的平衡点。如果你换到一个正负样本比例差不多的数据集这个值大概率不是最优的。我的建议是先用论文默认参数跑一版把训练曲线画出来作为baseline再根据效果调 (\gamma) 和 (\alpha)。不要上来就改参数否则你根本不知道是数据问题还是损失函数问题。5.2 调参失败现场γ太大、α乱设我自己调参踩过的坑可以列出来给你参考。有一次我在一个细粒度分类任务上直接用 (\gamma5)结果模型训了十几个epochaccuracy完全不动。原因是 (\gamma) 太大时所有样本的损失都太小梯度也随之变小模型参数更新很慢。排查了半天才意识到是Focal Loss的锅。后来把 (\gamma) 降到1.5训练曲线就恢复正常了。还有一次把 (\alpha) 设成0.1本意是更重视正样本。实际正好相反因为 (\alpha_t) 会把正样本的分类权重压到很低模型最后倾向于把所有样本都预测成背景。这个经验告诉我(\alpha) 虽然叫正样本权重但它是在整体损失的尺度上起作用的不能拍脑袋设必须结合负样本比例来看。如果你用的是二分类建议先固定 (\gamma2)然后在一组候选 (\alpha)0.1、0.2、0.25、0.3上用验证集指标做对比。多分类场景更复杂一些因为每个类别都有自己的频率一般用alpha 1/class_frequency归一化或者干脆用固定的0.25先跑。5.3 什么时候不能用Focal LossFocal Loss不是银弹。我在三种场景下见过它水土不服回归任务。上面提过它只适用于分类分支。回归输出是连续值没有概率意义上的难易之分硬套Focal Loss只会干扰回归收敛。类别分布均匀的多分类任务。如果每个类别的样本数差不多样本难易程度也比较均衡Focal Loss带来的提升很小甚至会让loss曲线变难调。普通交叉熵或者label smoothing效果更好。噪声标签比较多的数据集。Focal Loss会把难分样本的权重放大但难分样本并不一定是“值得学”的样本它也可能是标注错误的样本。我之前在一个标注质量很差的业务数据集上尝试Focal Loss不仅没有提升反而让模型开始拟合那些错误标注的难例。这个坑需要特别留意。6. 几个常见问题速查与个人体会最后这部分我把平时被问得最多的问题整理成一套速查都是比较实际的点。6.1 Focal Loss与OHEM的差别OHEMOnline Hard Example Mining的思路是先计算每个样本的损失取损失最高的那部分样本参与反向传播。Focal Loss则是给所有样本计算一个连续权重不丢样本只是把简单样本的权重压低。这种区别带来的影响很直观OHEM会让模型只关注困难样本但训练初期如果困难样本里有很多噪声模型容易被带偏Focal Loss则始终保留全部样本的信息只是把注意力向难例倾斜。所以Focal Loss训练起来通常更稳收敛曲线也更平滑。6.2 对噪声标签会不会更敏感会尤其是当你把 (\gamma) 调得比较大时。因为噪声标签通常就是那些模型学不动、损失居高不下的“难例”Focal Loss会进一步放大它们的权重。我现在的习惯是如果数据质量不是很有把握先用普通CE训练一版观察哪些样本的损失一直很高人工抽看一批确认不是标注问题后再切换到Focal Loss。这个过程花的时间不长但能避免很多返工。6.3 我的实操建议最后说几个稳定复用的经验。如果你刚接触Focal Loss我建议按这个顺序来先用普通交叉熵训练一个epoch作为baseline明确当前瓶颈到底是类别不平衡还是难例太多。切换Focal Loss时固定 (\alpha0.25)先把 (\gamma) 从1.0到2.0之间扫一遍每次只改一个参数。观察训练曲线确保loss不是直接躺平(\gamma) 太大也不是震荡太厉害(\gamma) 太小。调 (\alpha)。如果正样本recall偏低别急着加大 (\alpha)先看是不是 (\gamma) 过大了。我自己这几年用下来最大的心得是Focal Loss不是一个“装上就涨点”的组件它解决的是特定的失效模式——大量简单样本淹没难例梯度。你的任务里如果确实存在这个问题它的效果会非常明显如果不存在它可能只是个锦上添花甚至帮倒忙的东西。理解了这一点比背下公式更有价值。
上一篇/下一篇内容由系统自动关联
返回资讯列表 →