共计 2218 个字符,预计需要花费 6 分钟才能阅读完成。
背景痛点
在深度学习的分类任务中,损失函数的选择至关重要。很多初学者容易犯一个错误:在分类任务中错误地使用均方误差(MSE)损失函数。MSE 在回归任务中表现良好,但在分类任务中却存在几个明显的问题:

- MSE 损失对异常值过于敏感,容易导致训练不稳定
- 当使用 sigmoid 激活函数时,MSE 会导致梯度消失问题,特别是在预测接近 0 或 1 时
- MSE 假设误差服从高斯分布,而分类任务的输出实际上是伯努利分布
另一个常见问题是多标签分类场景下 sigmoid 与 softmax 的混淆。很多同学会错误地在多标签任务中使用 softmax,但实际上:
- softmax 适用于互斥的单标签分类(每个样本只能属于一个类别)
- sigmoid 适用于非互斥的多标签分类(每个样本可以同时属于多个类别)
数学原理
BCE(Binary Cross-Entropy)损失函数源自极大似然估计。假设我们有一个二分类问题,真实标签为 y∈{0,1},模型预测概率为 p,则似然函数可以表示为:
$$
L(y,p) = p^y(1-p)^{1-y}
$$
取对数得到对数似然:
$$
\log L(y,p) = y\log p + (1-y)\log(1-p)
$$
为了将其转化为损失函数(越小越好),我们取负值:
$$
\mathcal{L}_{BCE} = -[y\log p + (1-y)\log(1-p)]
$$
这个公式有几个重要特性:
- 当 y = 1 时,损失为 -log(p),预测越接近 0 损失越大
- 当 y = 0 时,损失为 -log(1-p),预测越接近 1 损失越大
- 对数运算使得预测错误时的惩罚呈指数增长
PyTorch 实战
PyTorch 提供了两种实现 BCE 损失的方式:
nn.BCELoss:输入必须是经过 sigmoid 处理后的概率值(0 到 1 之间)nn.BCEWithLogitsLoss:输入是 logits(未经过 sigmoid),内部会自动应用 sigmoid
推荐使用BCEWithLogitsLoss,因为它更数值稳定(内部使用了 log-sum-exp 技巧)。示例代码:
import torch
import torch.nn as nn
# 正确用法示例
bce_loss = nn.BCEWithLogitsLoss()
# 模型输出 logits(未经 sigmoid)logits = torch.randn(4, 1) # 假设 batch_size=4
# 真实标签(0 或 1)targets = torch.tensor([[1], [0], [1], [1]], dtype=torch.float32)
loss = bce_loss(logits, targets)
对于多标签分类(比如图像中同时包含 ” 猫 ” 和 ” 狗 ”),处理方式如下:
# 多标签分类示例
num_classes = 10 # 假设有 10 个可能的标签
logits = torch.randn(4, num_classes) # 每个样本有 10 个 logits
targets = torch.randint(0, 2, (4, num_classes)).float() # 每个标签独立
loss = nn.BCEWithLogitsLoss()(logits, targets)
高级话题
样本不平衡处理
当正负样本比例严重不平衡时,可以使用 pos_weight 参数:
# 假设正样本是负样本的 1 /10
pos_weight = torch.tensor([10.0])
bce_loss = nn.BCEWithLogitsLoss(pos_weight=pos_weight)
与 Focal Loss 结合
Focal Loss 通过降低易分类样本的权重来解决类别不平衡问题。可以这样实现:
class FocalBCELoss(nn.Module):
def __init__(self, alpha=0.25, gamma=2):
super().__init__()
self.alpha = alpha
self.gamma = gamma
def forward(self, logits, targets):
bce_loss = nn.BCEWithLogitsLoss(reduction='none')(logits, targets)
pt = torch.exp(-bce_loss)
focal_loss = self.alpha * (1-pt)**self.gamma * bce_loss
return focal_loss.mean()
避坑指南
- NaN 问题排查
- 确保输入没有极端的值(比如未归一化的超大 logits)
-
可以尝试梯度裁剪:
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) -
多 GPU 训练
- 确保各 GPU 上的损失正确聚合:
loss = loss.mean() -
使用
DistributedDataParallel而不是DataParallel -
数值稳定性
- 始终优先使用
BCEWithLogitsLoss而不是手动 sigmoid+BCELoss - 可以添加小的 epsilon 防止 log(0):
-[(y+eps)*log(p+eps) + (1-y+eps)*log(1-p+eps)]
结论与思考
BCE 损失函数虽然简单,但在实际应用中仍有许多值得深入探讨的问题:
- 在大语言模型 (LLM) 中,BCE 是否适合作为下一个 token 预测的损失函数?
- 当标签存在噪声时,如何调整 BCE 的鲁棒性?
- BCE 损失与标签平滑 (label smoothing) 技术如何结合使用效果最佳?
这些问题都值得在实践中进一步探索和验证。
