共计 2386 个字符,预计需要花费 6 分钟才能阅读完成。
1. 背景痛点
在二元分类任务中,损失函数的选择直接影响到模型的训练效果。一个常见的误区是直接使用均方误差(MSE)作为损失函数。MSE 损失在回归问题中表现良好,但在分类任务中存在以下问题:

- MSE 损失对概率输出的惩罚不对称,导致模型训练效率低下
- 当预测概率接近 0 或 1 时,MSE 的梯度会变得非常小,造成梯度消失问题
手动实现二元交叉熵(BCE)损失时,开发者经常遇到数值稳定性问题。最常见的是当预测概率 p 接近 0 或 1 时,log(p) 或 log(1-p) 会趋向于负无穷大,导致数值溢出。
2. 数学原理
2.1 从极大似然估计推导
BCE 损失可以从极大似然估计的角度推导出来。对于二元分类问题,我们希望最大化观测数据的似然概率。假设我们有:
- 真实标签 y ∈ {0,1}
- 预测概率 p = P(y=1|x)
则似然函数可以表示为:
$$
L(p) = p^y(1-p)^{1-y}
$$
为了简化计算,我们通常取负对数似然作为损失函数:
$$
L = -[y\cdot\log(p) + (1-y)\cdot\log(1-p)]
$$
2.2 信息论解释
从信息论角度看,交叉熵衡量的是两个概率分布之间的差异。在二元分类中,我们希望最小化预测分布 p 与真实分布 y 之间的交叉熵:
$$
H(y,p) = -\mathbb{E}_y[\log p]
$$
交叉熵可以分解为熵和 KL 散度,其中熵是固定值,因此最小化交叉熵等价于最小化 KL 散度。
3. PyTorch 实现对比
3.1 基础手动实现
def manual_bce_loss(y_pred, y_true, eps=1e-12):
"""
手动实现 BCE 损失,添加 epsilon 防止数值溢出
Args:
y_pred: 预测概率 [batch_size]
y_true: 真实标签 [batch_size]
eps: 极小值防止 log(0)
Returns:
BCE 损失值
"""
y_pred = torch.clamp(y_pred, eps, 1. - eps) # 截断到 [eps, 1-eps]
loss = -(y_true * torch.log(y_pred) + (1 - y_true) * torch.log(1 - y_pred))
return loss.mean()
3.2 nn.BCELoss
PyTorch 内置的 BCELoss 使用简单,但需要注意输入需要预先通过 sigmoid 激活:
import torch.nn as nn
bce_loss = nn.BCELoss()
sigmoid = nn.Sigmoid()
y_pred = model(x) # 原始输出
prob = sigmoid(y_pred) # 需要手动 sigmoid
loss = bce_loss(prob, y_true)
3.3 nn.BCEWithLogitsLoss
这是更推荐的实现方式,它整合了 sigmoid 和 BCE 损失,具有更好的数值稳定性:
bce_logits_loss = nn.BCEWithLogitsLoss()
y_pred = model(x) # 原始输出
loss = bce_logits_loss(y_pred, y_true) # 自动处理 sigmoid
4. 性能实验
我们在 MNIST 数据集上进行了二分类实验(区分数字 5 和非 5),比较三种实现方式的训练速度和收敛情况:
- 手动实现:训练速度较慢,需要额外处理数值稳定性
- nn.BCELoss:速度中等,需要手动 sigmoid
- nn.BCEWithLogitsLoss:训练速度最快,收敛最稳定
通过可视化 sigmoid 函数的梯度,可以观察到在饱和区(p 接近 0 或 1)时,梯度接近于 0,这解释了为什么使用 MSE 损失在分类任务中效果不佳。
5. 避坑指南
5.1 多标签分类
在多标签分类任务中,可以使用 pos_weight 参数调整正样本的权重:
# 假设正样本是负样本的 3 倍少
pos_weight = torch.tensor([3.0])
bce_logits_loss = nn.BCEWithLogitsLoss(pos_weight=pos_weight)
5.2 类别不平衡
对于类别不平衡问题,可以使用 weight 参数调整每个样本的权重:
# 假设类别 0 的权重是 1,类别 1 的权重是 2
weight = torch.tensor([1.0, 2.0])
bce_logits_loss = nn.BCEWithLogitsLoss(weight=weight)
5.3 混合精度训练
在使用混合精度训练时,需要注意数值精度问题:
- 避免在损失计算中使用 fp16,可能导致精度不足
- 建议保持损失计算部分为 fp32
6. 代码规范
所有 PyTorch 代码遵循 Google 代码风格规范:
- 张量操作明确标注 shape 变化
- 关键步骤添加注释
- 变量命名有意义
示例规范代码:
# 输入形状: [batch_size, num_classes]
logits = model(inputs)
# 计算损失前确保 y_true 形状匹配
targets = targets.view(-1, 1).float() # [batch_size] -> [batch_size, 1]
loss = bce_loss(logits, targets)
7. 延伸思考
- 对于严重类别不平衡问题,可以尝试 Focal Loss 作为 BCE 的替代,它对难样本赋予更高权重:
$$
FL(p_t) = -\alpha_t(1-p_t)^\gamma\log(p_t)
$$
-
Label smoothing 技术可以缓解过拟合问题,通过将硬标签(0 或 1)替换为软标签(如 0.1 或 0.9),使模型不会对预测过度自信。
-
考虑 BCE 损失在不同任务中的变体,如多标签分类中如何调整阈值,排序任务中如何与 NDCG 等指标结合等。
总结
BCE 损失函数是二元分类任务的核心组件,理解其数学原理和实现细节对于构建高效模型至关重要。PyTorch 提供了多种实现方式,其中 BCEWithLogitsLoss 是最推荐的做法。在实际应用中,还需要考虑类别不平衡、数值稳定性等问题,并针对具体任务进行调整。
