共计 2372 个字符,预计需要花费 6 分钟才能阅读完成。
为什么需要 BCE 损失函数
在二分类和多标签任务中,模型需要输出每个类别的概率值(0 到 1 之间)。二元交叉熵(Binary Cross-Entropy, BCE)通过衡量预测概率与真实标签的差异,成为这类任务的标配损失函数。举个通俗的例子:当预测癌症是否复发时,我们不仅需要判断 ” 是 / 否 ”(二分类),还可能同时预测 ” 转移风险 ”、” 化疗敏感性 ” 等多标签属性——BCE 损失函数能优雅地处理这些场景。

与多分类交叉熵(CrossEntropyLoss)不同,BCE 的每个输出节点独立计算损失,这使得它特别适合标签不互斥的场景(比如一张图片可以同时包含 ” 猫 ” 和 ” 阳光 ” 两个标签)。
数学原理拆解
基本公式推导
BCE 损失函数的数学表达式为:
$$
L = -\frac{1}{N}\sum_{i=1}^N [y_i\cdot\log(p_i) + (1-y_i)\cdot\log(1-p_i)]
$$
其中:
– $y_i$ 是真实标签(0 或 1)
– $p_i$ 是预测概率(通过 sigmoid 函数输出)
– $N$ 是样本数量
这个公式的直观理解是:当真实标签 $y_i=1$ 时,损失由 $\log(p_i)$ 决定(预测概率越接近 1 损失越小);当 $y_i=0$ 时,损失由 $\log(1-p_i)$ 决定。
数值稳定性问题
实际计算时会遇到两个典型问题:
1. 当 $p_i$ 接近 0 时,$\log(p_i)$ 趋向负无穷
2. 当 $p_i$ 接近 1 时,$\log(1-p_i)$ 趋向负无穷
PyTorch 的 BCEWithLogitsLoss 通过将 sigmoid 和 log 运算合并,采用以下等价形式避免数值溢出:
$$
L = \frac{1}{N}\sum_{i=1}^N \max(p_i, 0) – p_i\cdot y_i + \log(1 + e^{-|p_i|})
$$
这种实现方式称为 log-sum-exp 技巧,既保持数学等价性,又提高了数值稳定性。
PyTorch 实战指南
基础使用对比
import torch
import torch.nn as nn
# 方式 1:需要手动添加 sigmoid 层(容易出错)loss_fn1 = nn.BCELoss()
output = torch.sigmoid(model(input))
loss1 = loss_fn1(output, target)
# 方式 2:推荐!内置 sigmoid 和稳定化计算
loss_fn2 = nn.BCEWithLogitsLoss()
output = model(input) # 直接输出 logits
loss2 = loss_fn2(output, target)
完整多标签分类示例
import torch
from torch import nn, optim
from torch.utils.data import DataLoader
class MultiLabelModel(nn.Module):
def __init__(self, input_dim: int, output_dim: int):
super().__init__()
self.linear = nn.Linear(input_dim, output_dim)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.linear(x) # 输出 logits
# 假设有 10 个特征和 5 个标签
dataset = ... # 自定义数据集
model = MultiLabelModel(10, 5)
optimizer = optim.Adam(model.parameters())
# 处理类别不平衡(假设正样本比负样本少)pos_weight = torch.tensor([2.0, 1.5, 3.0, 1.0, 2.5]) # 每个标签的正样本权重
loss_fn = nn.BCEWithLogitsLoss(pos_weight=pos_weight)
for epoch in range(100):
for x, y in DataLoader(dataset, batch_size=32):
optimizer.zero_grad()
logits = model(x)
loss = loss_fn(logits, y.float()) # 注意 target 需要是 float 类型
loss.backward()
optimizer.step()
常见问题与解决方案
输入值范围检查
使用 BCELoss 时必须确保输入值在 (0,1) 区间:
# 错误示范(会导致 NaN)output = model(input)
loss = nn.BCELoss()(output, target) # 缺少 sigmoid
# 正确做法
torch.sigmoid(output).clamp(min=1e-8, max=1-1e-8) # 添加微小偏移避免 log(0)
多标签任务注意事项
- 每个标签独立计算 sigmoid(不要用 softmax!)
- 预测阶段需要手动对 logits 取 sigmoid:
with torch.no_grad(): probs = torch.sigmoid(model(input)) predictions = (probs > 0.5).float() # 按阈值划分
显存优化技巧
对于极多标签场景(如超过 1000 个标签):
1. 使用混合精度训练
2. 梯度累积减少 batch size
3. 对不活跃标签采用稀疏矩阵
思考与实践
- 当正负样本比例为 1:100 时,pos_weight 应该设置为何值?如何通过训练数据自动计算?
- 在多标签任务中,为什么不能将 BCE 损失简单拆分为多个二分类任务?
- 尝试用 BCEWithLogitsLoss 实现一个新闻多标签分类器(数据集可用 Reuters-21578)
希望通过这篇指南,你能避开 BCE 使用中的那些 ” 坑 ”,在实战中游刃有余。如果有任何实现问题,欢迎在评论区交流讨论——毕竟每个深度学习实践者都曾在损失函数上栽过跟头,这正是我们成长的必经之路。
