先聊点实际的: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):
- 当真实标签 (y=1) 时,(p_t = p)
- 当真实标签 (y=0) 时,(p_t = 1-p)
于是交叉熵就可以简写成:
[ CE(p_t) = -\log(p_t) ]
这个式子非常漂亮地统一了正负样本的写法。(p_t) 越大,说明模型对这个样本分得越对,对应的 (-\log(p_t)) 就越小;(p_t) 越小,说明模型分错了或者分得不确定,损失就越大。
问题出在哪?假设模型对某个简单负样本输出 (p=0.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,损失基本保留。
你可以具体感受一下数值变化。假设 (\gamma=2),我列一个简单的对照表:
| 样本类型 | (p_t) | 交叉熵损失 | 调制因子 ((1-p_t)^2) | Focal Loss |
|---|---|---|---|---|
| 简单负样本 | 0.9 | 0.105 | 0.01 | 0.00105 |
| 一般样本 | 0.5 | 0.693 | 0.25 | 0.173 |
| 难分样本 | 0.1 | 2.303 | 0.81 | 1.865 |
你会发现,简单样本的损失被压缩到原来的1/100,而难样本只被压缩到原来的81%左右。这样一比较,难样本在总损失里的“话语权”就大大提升了。
这个设计的本质,相当于给每道题乘上一个“难度加权系数”。送分题权重趋近于0,难题权重接近1。模型就能把更多精力放在真正需要优化的边界样本上。
2.3 α平衡因子:把正负样本的权重也一起调了
只加调制因子,依然有一个问题:虽然简单负样本的权重被压低了,但负样本数量实在太多,累积起来还是可能压过正样本。于是论文里又加了一个 (\alpha) 平衡因子,得到完整版:
[ FL(p_t) = -\alpha_t (1 - p_t)^\gamma \log(p_t) ]
这里的 (\alpha_t) 作用和加权交叉熵里的 (\alpha) 一样。当真实标签 (y=1) 时,(\alpha_t = \alpha);当 (y=0) 时,(\alpha_t = 1-\alpha)。(\alpha) 一般取0.25,意思是正样本权重0.25,负样本权重0.75。注意这不是说重视负样本,而是负样本数量实在太多,(\alpha) 会把正样本的相对重要性明显放大,通常配合 (\gamma=2) 使用效果比较好。
理解了这两层后,你会发现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, alpha=0.25, gamma=2.0, reduction='mean'): 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_classes=num_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(), reduction='none' ) 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),记得加一个eps=1e-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,本质上就是 (\gamma=0) 的特殊Focal Loss。而YOLOv8在DFL(Distribution Focal Loss)损失里,也借鉴了Focal的思想,对预测分布中概率较高的部分加强约束。
如果你要手动在YOLO的自定义训练里接入Focal Loss,一般改分类分支就行:
# 原来 loss_cls = nn.BCEWithLogitsLoss()(pred_cls, target_cls) # 改成 loss_cls = FocalLoss(alpha=0.25, gamma=2.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, beta=0.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), label='train_loss') plt.plot(epochs, smooth_curve(val_loss), label='val_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的经验从哪来
论文给的默认值是 (\gamma=2)、(\alpha=0.25)。这个组合是在COCO数据集上大量实验调出来的,不是拍脑袋定的。它的道理在于:
- 在 (\gamma=2) 时,简单样本的损失衰减已经非常厉害,大到再增加 (\gamma),难样本的权重也会被压得过低,导致整体学习速度变慢。
- (\alpha=0.25) 则是在正负样本比例极端的情况下,配合 (\gamma=2) 找到的平衡点。如果你换到一个正负样本比例差不多的数据集,这个值大概率不是最优的。
我的建议是:先用论文默认参数跑一版,把训练曲线画出来作为baseline,再根据效果调 (\gamma) 和 (\alpha)。不要上来就改参数,否则你根本不知道是数据问题还是损失函数问题。
5.2 调参失败现场:γ太大、α乱设
我自己调参踩过的坑可以列出来给你参考。
有一次我在一个细粒度分类任务上直接用 (\gamma=5),结果模型训了十几个epoch,accuracy完全不动。原因是 (\gamma) 太大时,所有样本的损失都太小,梯度也随之变小,模型参数更新很慢。排查了半天才意识到是Focal Loss的锅。后来把 (\gamma) 降到1.5,训练曲线就恢复正常了。
还有一次把 (\alpha) 设成0.1,本意是更重视正样本。实际正好相反,因为 (\alpha_t) 会把正样本的分类权重压到很低,模型最后倾向于把所有样本都预测成背景。这个经验告诉我:(\alpha) 虽然叫正样本权重,但它是在整体损失的尺度上起作用的,不能拍脑袋设,必须结合负样本比例来看。
如果你用的是二分类,建议先固定 (\gamma=2),然后在一组候选 (\alpha)(0.1、0.2、0.25、0.3)上用验证集指标做对比。多分类场景更复杂一些,因为每个类别都有自己的频率,一般用alpha = 1/class_frequency归一化,或者干脆用固定的0.25先跑。
5.3 什么时候不能用Focal Loss
Focal Loss不是银弹。我在三种场景下见过它水土不服:
回归任务。上面提过,它只适用于分类分支。回归输出是连续值,没有概率意义上的难易之分,硬套Focal Loss只会干扰回归收敛。
类别分布均匀的多分类任务。如果每个类别的样本数差不多,样本难易程度也比较均衡,Focal Loss带来的提升很小,甚至会让loss曲线变难调。普通交叉熵或者label smoothing效果更好。
噪声标签比较多的数据集。Focal Loss会把难分样本的权重放大,但难分样本并不一定是“值得学”的样本,它也可能是标注错误的样本。我之前在一个标注质量很差的业务数据集上尝试,Focal Loss不仅没有提升,反而让模型开始拟合那些错误标注的难例。这个坑需要特别留意。
6. 几个常见问题速查与个人体会
最后这部分,我把平时被问得最多的问题整理成一套速查,都是比较实际的点。
6.1 Focal Loss与OHEM的差别
OHEM(Online 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时,固定 (\alpha=0.25),先把 (\gamma) 从1.0到2.0之间扫一遍,每次只改一个参数。
- 观察训练曲线,确保loss不是直接躺平((\gamma) 太大)也不是震荡太厉害((\gamma) 太小)。
- 调 (\alpha)。如果正样本recall偏低,别急着加大 (\alpha),先看是不是 (\gamma) 过大了。
我自己这几年用下来,最大的心得是:Focal Loss不是一个“装上就涨点”的组件,它解决的是特定的失效模式——大量简单样本淹没难例梯度。你的任务里如果确实存在这个问题,它的效果会非常明显;如果不存在,它可能只是个锦上添花甚至帮倒忙的东西。理解了这一点,比背下公式更有价值。