欢迎光临
我们一直在努力

损失函数大汇总(六)(Focal Loss附公式推导与代码)

损失函数大汇总(六)(Focal Loss附公式推导与代码)

一、引言

前面写了 Binary Cross Entropy Loss。BCE 在二分类和多标签分类中非常常用,但它有一个很现实的问题:当正负样本极不平衡时,大量容易分类的样本会主导损失,真正困难的样本反而学得不够。

例如目标检测中,绝大多数候选框其实都是背景;异常检测中,正常样本远多于异常样本;缺陷识别中,真正困难的少数类样本常常被大量简单样本淹没。这个时候,单纯使用 BCE 往往不够理想。

Focal Loss 就是为了解决这类问题提出的。它的核心思想很直接:降低简单样本的损失权重,把训练重点更多放到难分类样本上。

本文就详细整理一下 Focal Loss 的定义、由来、推导、性质、使用场景和代码实现。

限于笔者水平,文中如有疏漏或错误,欢迎留言交流。

二、Focal Loss

Focal Loss 是在交叉熵基础上改造出来的一种损失函数,最早主要用于密集目标检测任务。它特别适合处理下面两类问题:

  • 类别极不平衡
  • 简单样本太多,困难样本太少

它的思路不是推翻 BCE,而是在 BCE 前面加一个调制因子,让模型对“已经分得很对的样本”少关注一些,对“还没分好的样本”多关注一些。

1. 从 BCE 出发

先回顾一下二分类 BCE 的单样本形式:

LBCE=−[ylog⁡(y^)+(1−y)log⁡(1−y^)]
\\mathcal{L}_{\\mathrm{BCE}} = -\\left[y\\log(\\hat{y}) + (1-y)\\log(1-\\hat{y})\\right]
LBCE=[ylog(y^)+(1y)log(1y^)]

其中:

  • y∈{0,1}y \\in \\{0,1\\}y{0,1} 表示真实标签;
  • y^∈(0,1)\\hat{y} \\in (0,1)y^(0,1) 表示预测为正类的概率。

把它拆开看:

  • y=1y=1y=1 时:

LBCE=−log⁡(y^)
\\mathcal{L}_{\\mathrm{BCE}} = -\\log(\\hat{y})
LBCE=log(y^)

  • y=0y=0y=0 时:

LBCE=−log⁡(1−y^)
\\mathcal{L}_{\\mathrm{BCE}} = -\\log(1-\\hat{y})
LBCE=log(1y^)

这个形式本身没有问题,但它会带来一个现象:大量简单样本即使损失很小,数量一多,依然会占据总损失的大头。

2. 数学定义

为了写得更紧凑,先定义:

pt={y^,y=11−y^,y=0
p_t =
\\begin{cases}
\\hat{y}, & y=1 \\\\
1-\\hat{y}, & y=0
\\end{cases}
pt={y^,1y^,y=1y=0

这个 ptp_tpt 可以理解为:模型对真实类别给出的预测概率。

于是 BCE 可以统一写成:

LBCE=−log⁡(pt)
\\mathcal{L}_{\\mathrm{BCE}} = -\\log(p_t)
LBCE=log(pt)

Focal Loss 在它前面乘了一个调制因子 (1−pt)γ(1-p_t)^\\gamma(1pt)γ,于是得到:

LFL=−(1−pt)γlog⁡(pt)
\\mathcal{L}_{\\mathrm{FL}} = -(1-p_t)^\\gamma \\log(p_t)
LFL=(1pt)γlog(pt)

其中:

  • γ≥0\\gamma \\ge 0γ0 称为聚焦参数(focusing parameter);
  • (1−pt)γ(1-p_t)^\\gamma(1pt)γ 用来压低简单样本的损失。

进一步地,为了处理类别不平衡,通常再加入一个类别平衡系数 αt\\alpha_tαt,得到更常见的形式:

LFL=−αt(1−pt)γlog⁡(pt)
\\mathcal{L}_{\\mathrm{FL}} = -\\alpha_t (1-p_t)^\\gamma \\log(p_t)
LFL=αt(1pt)γlog(pt)

其中:

αt={α,y=11−α,y=0
\\alpha_t =
\\begin{cases}
\\alpha, & y=1 \\\\
1-\\alpha, & y=0
\\end{cases}
αt={α,1α,y=1y=0

这里:

  • α∈[0,1]\\alpha \\in [0,1]α[0,1] 用来调节正负样本权重;
  • γ\\gammaγ 用来调节“降低简单样本权重”的强度。

3. 这个式子到底在做什么

Focal Loss 最核心的地方就在这一项:

(1−pt)γ
(1-p_t)^\\gamma
(1pt)γ

它的作用可以直接分情况理解。

(1)当样本很容易分类时

如果一个样本已经分得很好,比如真实类别概率 pt=0.95p_t=0.95pt=0.95,那么:

1−pt=0.05
1-p_t = 0.05
1pt=0.05

如果 γ=2\\gamma=2γ=2,那么:

(1−pt)γ=0.052=0.0025
(1-p_t)^\\gamma = 0.05^2 = 0.0025
(1pt)γ=0.052=0.0025

这个系数会非常小,于是这个样本的损失会被大幅压低。

(2)当样本很难分类时

如果一个样本分得很差,比如真实类别概率 pt=0.2p_t=0.2pt=0.2,那么:

1−pt=0.8
1-p_t = 0.8
1pt=0.8

如果 γ=2\\gamma=2γ=2,那么:

(1−pt)γ=0.82=0.64
(1-p_t)^\\gamma = 0.8^2 = 0.64
(1pt)γ=0.82=0.64

这个系数仍然不小,因此困难样本的损失会被保留下来。

所以 Focal Loss 的本质就是:

  • 简单样本:损失压低
  • 困难样本:损失保留
  • 训练重点自动向困难样本倾斜

4. Focal Loss 与 BCE 的关系

Focal Loss 可以直接看成是 BCE 的加权版本。

因为 BCE 写成统一形式就是:

LBCE=−log⁡(pt)
\\mathcal{L}_{\\mathrm{BCE}} = -\\log(p_t)
LBCE=log(pt)

而 Focal Loss 是:

LFL=−αt(1−pt)γlog⁡(pt)
\\mathcal{L}_{\\mathrm{FL}} = -\\alpha_t (1-p_t)^\\gamma \\log(p_t)
LFL=αt(1pt)γlog(pt)

所以它等于:

LFL=αt(1−pt)γLBCE
\\mathcal{L}_{\\mathrm{FL}} = \\alpha_t (1-p_t)^\\gamma \\mathcal{L}_{\\mathrm{BCE}}
LFL=αt(1pt)γLBCE

也就是说,Focal Loss 并没有改变交叉熵的基本方向,而是在每个样本前面乘了一个动态权重。

这个权重由两部分组成:

  • αt\\alpha_tαt:处理类别不平衡
  • (1−pt)γ(1-p_t)^\\gamma(1pt)γ:处理难易样本不平衡

特别地,当 γ=0\\gamma = 0γ=0 时:

(1−pt)0=1
(1-p_t)^0 = 1
(1pt)0=1

这时 Focal Loss 退化成:

LFL=−αtlog⁡(pt)
\\mathcal{L}_{\\mathrm{FL}} = -\\alpha_t \\log(p_t)
LFL=αtlog(pt)

如果再取 αt=1\\alpha_t = 1αt=1,那就进一步退化为标准 BCE。

所以可以把 Focal Loss 看成是 BCE 的一个推广版本。

5. 参数 γ\\gammaγ 的作用

γ\\gammaγ 是 Focal Loss 里最关键的参数。

(1)γ=0\\gamma = 0γ=0

这时 Focal Loss 就退化成加权 BCE,不再区分简单样本和困难样本。

(2)γ\\gammaγ 较小

例如 γ=1\\gamma=1γ=1,简单样本的损失会被压低,但程度还不算特别强。

(3)γ\\gammaγ 较大

例如 γ=2\\gamma=2γ=2 或更大时,容易样本的损失会被压得更明显,模型会更集中地去学那些困难样本。

一般来说:

  • γ\\gammaγ 越大,越强调困难样本;
  • 但如果取得过大,可能会导致大量样本的损失被压得太低,训练不够稳定。

在很多实际工作里,γ=2\\gamma=2γ=2 是一个比较常见的起点。

6. 参数 α\\alphaα 的作用

α\\alphaα 主要用来解决类别数量不平衡的问题。

设正类较少、负类很多。如果直接训练,模型可能更容易偏向负类。这时可以给少数类更大的权重。

定义:

αt={α,y=11−α,y=0
\\alpha_t =
\\begin{cases}
\\alpha, & y=1 \\\\
1-\\alpha, & y=0
\\end{cases}
αt={α,1α,y=1y=0

如果正类更少,就可以把 α\\alphaα 设得大一些,让正样本在损失中占更高比重。

所以:

  • α\\alphaα 主要解决类别不平衡
  • γ\\gammaγ 主要解决样本难易程度不平衡

这两个参数解决的问题不完全一样。

7. 梯度层面的直观理解

Focal Loss 最重要的效果,不只是让简单样本损失变小,更关键的是:让简单样本对参数更新的贡献也变小。

因为一个样本的损失被 (1−pt)γ(1-p_t)^\\gamma(1pt)γ 压低之后,它反向传播时产生的梯度也会同步变弱。于是训练过程中:

  • 已经分得很对的样本,基本不会再反复主导更新;
  • 那些还分不清的样本,会继续提供更强的优化信号。

这也是 Focal Loss 在难样本挖掘上很有效的原因。

如果只从直观上记一句话,那就是:

BCE 是所有样本都学,Focal Loss 是重点学难样本。

8. 函数特性

(1)对简单样本自动降权

这是 Focal Loss 最核心的特点。随着 ptp_tpt 增大,(1−pt)γ(1-p_t)^\\gamma(1pt)γ 会迅速减小,从而使简单样本的损失变小。

(2)保留交叉熵的基本性质

它并没有脱离交叉熵的框架,仍然是在提高真实类别概率、压低错误类别概率。

(3)适合类别极不平衡场景

特别是在负样本极多、正样本极少时,Focal Loss 往往比普通 BCE 更有效。

(4)对超参数更敏感

相比 BCE,Focal Loss 多了 α\\alphaαγ\\gammaγ 两个关键超参数,因此调参会更重要。

9. 使用场景与局限性

Focal Loss 主要适合下面这些任务:

  • 正负样本极不平衡的二分类任务
  • 多标签分类任务
  • 目标检测中的前景/背景分类
  • 缺陷检测、异常检测、故障检测
  • 医学图像中的病灶识别
  • 样本中“简单负类很多、困难正类很少”的任务
优势
  • 能有效降低大量简单样本的影响;
  • 对类别不平衡问题更友好;
  • 能让模型更关注困难样本;
  • 在检测、异常识别等任务中很常用。
缺点
  • 需要额外调节 α\\alphaαγ\\gammaγ
  • 当数据本身并不失衡时,未必比 BCE 更好;
  • 如果 γ\\gammaγ 太大,可能导致训练信号过弱;
  • 对标签噪声有时比较敏感,因为被误标的样本往往会被当成“困难样本”重点学习。

10. 多分类版本的写法

虽然 Focal Loss 最常见的是二分类或多标签版本,但它也可以扩展到多分类任务。

对于多分类任务,设真实类别对应的预测概率为 ptp_tpt,则多分类 Focal Loss 常写成:

LFL=−αt(1−pt)γlog⁡(pt)
\\mathcal{L}_{\\mathrm{FL}} = -\\alpha_t (1-p_t)^\\gamma \\log(p_t)
LFL=αt(1pt)γlog(pt)

这个形式和二分类统一表达式其实是一样的,只不过这里的 ptp_tpt 来自 Softmax 输出,而不是 Sigmoid 输出。

所以从本质上看,Focal Loss 并不是只能用于二分类,而是“交叉熵 + 难样本调制”的思想可以推广到多分类场景。

11. 代码实现

下面先用 NumPy 实现一个二分类版本的 Focal Loss。这里假设输入已经是概率,而不是 logits。

import numpy as np

def focal_loss_binary(y_true, y_pred, alpha=0.25, gamma=2.0, eps=1e-12):
"""
计算二分类 Focal Loss

参数:
y_true — 真实标签,取值为 0 或 1
y_pred — 预测为正类的概率,范围为 (0, 1)
alpha — 正类权重系数
gamma — 聚焦参数
eps — 防止 log(0) 的微小常数

返回:
平均 Focal Loss
"""
y_true = np.asarray(y_true, dtype=np.float32)
y_pred = np.asarray(y_pred, dtype=np.float32)

y_pred = np.clip(y_pred, eps, 1.0 eps)

pt = np.where(y_true == 1, y_pred, 1 y_pred)
alpha_t = np.where(y_true == 1, alpha, 1 alpha)

loss = alpha_t * ((1 pt) ** gamma) * np.log(pt)
return np.mean(loss)

def sigmoid(x):
"""
Sigmoid 函数
"""

x = np.asarray(x)
return 1.0 / (1.0 + np.exp(x))

示例使用

logits = np.array([2.0, 1.0, 0.5, 3.0, 3.5, 0.2])
y_true = np.array([1, 0, 1, 0, 1, 1])

y_pred = sigmoid(logits)
loss = focal_loss_binary(y_true, y_pred, alpha=0.25, gamma=2.0)

print("Predicted Probabilities:", y_pred)
print("Focal Loss:", loss)

12. PyTorch中的实现

下面给出一个简单的 PyTorch 版本实现。这里直接输入 logits,在损失函数内部完成 Sigmoid 和 Focal Loss 计算。

import torch
import torch.nn as nn
import torch.nn.functional as F

class BinaryFocalLoss(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: 原始输出,形状与 targets 一致
targets: 真实标签,取值为 0 或 1,float 类型
"""

probs = torch.sigmoid(logits)
probs = torch.clamp(probs, min=1e-6, max=11e-6)

pt = torch.where(targets == 1, probs, 1 probs)
alpha_t = torch.where(targets == 1, self.alpha, 1 self.alpha)

loss = alpha_t * ((1 pt) ** self.gamma) * torch.log(pt)

if self.reduction == 'mean':
return loss.mean()
elif self.reduction == 'sum':
return loss.sum()
else:
return loss

示例使用

import torch

criterion = BinaryFocalLoss(alpha=0.25, gamma=2.0)

logits = torch.tensor([2.0, 1.0, 0.5, 3.0, 3.5, 0.2], dtype=torch.float32)
targets = torch.tensor([1, 0, 1, 0, 1, 1], dtype=torch.float32)

loss = criterion(logits, targets)
print("Focal Loss:", loss.item())

如果放到模型训练中,一般写法如下:

import torch
import torch.nn as nn
import torch.optim as optim

model = nn.Linear(10, 1)
criterion = BinaryFocalLoss(alpha=0.25, gamma=2.0)
optimizer = optim.SGD(model.parameters(), lr=0.01)

x = torch.randn(8, 10)
y = torch.randint(0, 2, (8, 1)).float()

logits = model(x)
loss = criterion(logits, y)

optimizer.zero_grad()
loss.backward()
optimizer.step()

print("Loss:", loss.item())

13. Focal Loss 和 BCE 怎么选

这两个损失函数并不是谁一定比谁强,更重要的是任务特点。

如果任务数据比较平衡,或者没有明显的“简单样本淹没困难样本”的问题,那么普通 BCE 往往已经足够,而且更简单、稳定。

但如果你遇到下面这些现象:

  • 正负样本极不平衡;
  • 大量容易分类的背景样本占据训练主导;
  • 少数困难样本学不动;
  • 模型对少数类召回率偏低;

那么 Focal Loss 通常值得尝试。

可以简单理解为:

  • BCE:标准方案
  • Focal Loss:更偏向困难样本和不平衡场景的方案

14. 一个简单理解

Focal Loss 可以简单理解为:

已经学会的样本少管一点,还没学会的样本多管一点。

它不是让模型“忽略简单样本”,而是让简单样本别再反复占用太多训练资源。这样一来,模型就能把更多注意力放在真正难分的样本上。

三、小结

Focal Loss 是在交叉熵基础上发展出来的一种损失函数,特别适合类别不平衡和困难样本较少的场景。它通过引入调制因子 (1−pt)γ(1-p_t)^\\gamma(1pt)γ,自动降低简单样本的损失贡献,让模型更关注难样本。

如果把它和 BCE 放在一起看,那么最核心的区别就是:

  • BCE 对所有样本一视同仁;
  • Focal Loss 会主动降低容易样本的影响。

在目标检测、异常识别、缺陷检测、多标签分类等任务中,Focal Loss 都是一个很值得考虑的选择。

四、参考文献

  • Lin, T.-Y., Goyal, P., Girshick, R., He, K., & Dollár, P. Focal Loss for Dense Object Detection. IEEE Transactions on Pattern Analysis and Machine Intelligence, 2020.
  • Goodfellow, I., Bengio, Y., & Courville, A. Deep Learning. MIT Press, 2016.
  • Bishop, C. M. Pattern Recognition and Machine Learning. Springer, 2006.
赞(0)
未经允许不得转载:171主机测评 » 损失函数大汇总(六)(Focal Loss附公式推导与代码)
分享到: 更多 (0)

评论 抢沙发

  • 昵称 (必填)
  • 邮箱 (必填)
  • 网址