ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

PyTorch损失函数详解:BCELoss与BCEWithLogitsLoss的对比与应用

PyTorch损失函数详解:BCELoss与BCEWithLogitsLoss的对比与应用 先说结论如果你在 PyTorch 里做二分类、多标签分类或者类似的任务几乎绕不开BCELoss和BCEWithLogitsLoss这两个名字。很多初学者包括我自己刚入门的时候都踩过“为什么我用了 BCELoss 就报错”“为什么训练出来 loss 是 nan”“为什么两个函数算出来的结果不一样”这类坑。这篇博客我就把这两个函数从原理到实现、从参数到坑点一次性讲透。先说清楚它们各自能干什么BCELoss是“输入已经过 sigmoid 的概率值”之后计算二分类交叉熵BCEWithLogitsLoss则是把 sigmoid 层和 BCELoss 合并成一个操作直接吃模型的原始 logits 输出。后者在数值稳定性和工程实现上都更推荐。这篇文章适合谁看刚接触 PyTorch 的初学者正在做图像分类、文本分类、多标签任务的同学以及那些想知道“为什么官方文档推荐用 BCEWithLogitsLoss”但没搞懂深层原因的人。1. 先搞清楚 BCE 到底在计算什么1.1 从交叉熵到二元交叉熵在聊代码之前我们先从数学上把这两个函数的身世捋清楚。深度学习里的分类问题最常用的损失函数是交叉熵Cross Entropy。它衡量的是模型预测的概率分布和真实标签分布之间的差异。多分类任务里我们用的是 categorical cross entropy对应 PyTorch 里的CrossEntropyLoss。而二分类任务也就是标签只有 0 和 1 的情况交叉熵会退化成一种更简洁的形式这就是二元交叉熵Binary Cross Entropy。它的公式长这样[ L -\frac{1}{N} \sum_{i1}^{N} \left[ y_i \cdot \log(p_i) (1 - y_i) \cdot \log(1 - p_i) \right] ]其中 (y_i) 是真实标签0 或 1(p_i) 是模型预测为正类标签为 1的概率。单个样本的 loss 可以拆开看当 (y1) 时loss 是 (-\log(p))预测概率 (p) 越接近 1loss 越小当 (y0) 时loss 是 (-\log(1-p))预测概率 (p) 越接近 0loss 越小。这个公式的直觉特别简单模型对真实标签越自信loss 越低越不自信甚至完全搞反loss 越高。本质上它是在做“最大似然估计”让模型在训练数据上的似然概率最大化。1.2 为什么二分类不能用 MSE你可能会问回归任务用均方误差MSE用得好好的分类任务为什么非要换成交叉熵我用一个例子说明。假设一个二分类模型对某个样本输出 sigmoid 后的概率是 0.9真实标签是 1。如果使用 MSE计算出的梯度会很小因为预测值已经比较接近目标了。但问题在于当模型完全分错时——比如预测 0.1真实标签是 1——MSE 的梯度在某些情况下依然不够大因为 sigmoid 函数在两端饱和梯度趋近于 0这就导致模型学习速度极慢。而交叉熵配合 sigmoid 时梯度里会天然消掉 sigmoid 的导数项从而让误差越大、梯度越大模型学得越快。这就是“梯度消失”问题在损失函数层面的一个解法。后面我们会看到BCEWithLogitsLoss在数学上正是通过合并 sigmoid 和交叉熵把这个优势发挥到极致。这里先埋一个伏笔当你看到很多人说“BCEWithLogitsLoss比BCELoss数值更稳定”时本质上就是因为它在计算过程中规避了 sigmoid 函数在两端饱和导致的精度损失问题后面会详细展开。2. BCELoss最朴素的二分类损失2.1 手动实现一个 BCELossnn.BCELoss是 PyTorch 里最直接的实现。它要求输入是一个已经经过 sigmoid 激活的概率值范围在 (0, 1) 之间并且输入和 target 的 shape 必须一致。我们先用 Numpy 手动实现一遍帮助你建立直觉import numpy as np def binary_cross_entropy(y_true, y_pred, eps1e-7): # y_pred 是 sigmoid 之后的概率值 y_pred np.clip(y_pred, eps, 1 - eps) # 防止 log(0) 导致 nan loss -np.mean( y_true * np.log(y_pred) (1 - y_true) * np.log(1 - y_pred) ) return loss y_true np.array([1, 0, 1, 0]) y_pred np.array([0.9, 0.1, 0.8, 0.3]) print(binary_cross_entropy(y_true, y_pred))这个脚本输出约等于 0.25。为什么需要eps做裁剪因为如果某一项预测概率正好是 0 或 1log(0)会直接变成负无穷计算结果就是 nan。这是所有实现 CE 类损失函数都必须处理的细节PyTorch 内部同样有对应的保护机制。2.2 BCELoss 的使用姿势在 PyTorch 中BCELoss最简单的用法如下import torch import torch.nn as nn loss_fn nn.BCELoss() # 模拟一个 batch共 4 个样本每个样本只有一个输出节点 logits torch.tensor([[0.9], [0.1], [0.8], [0.3]], dtypetorch.float32) probs torch.sigmoid(logits) # 必须是 sigmoid 之后的概率 targets torch.tensor([[1.], [0.], [1.], [0.]], dtypetorch.float32) loss loss_fn(probs, targets) print(loss.item())请注意 target 的数据类型是torch.float32不是torch.long。很多初学者在这里栽跟头把标签定义成了整数型然后直接丢给BCELoss就会遇到奇怪的报错。2.3 输入 shape、dtype 和 device 的硬性要求BCELoss对输入输出的要求非常“苛刻”总结起来有四点第一input和target的 shape 必须完全一致。要么都是(N,)要么都是(N, C)要么都是(N, C, H, W)总之 broadcast 的情况虽然偶尔能跑通但结果往往不是你想要的建议一开始就保持一致。第二target必须是浮点类型即 0.0 和 1.0而不是整型 0 和 1。这一点和CrossEntropyLoss完全不同——后者要求 target 是torch.long类型的类别索引而BCELoss的 target 本质上是“概率值”虽然实践中只有 0 和 1 两种取值。第三input必须经过 sigmoid 激活取值范围要在 0 到 1 之间严格来说不能包含 0 和 1。如果你直接把 logits 丢进来loss 大概率会算出一个奇怪的值而且训练会非常不稳定。第四input和target要在同一个 device 上。GPU 训练时尤其容易忽略这一点一个在 CPU 一个在 GPU 会直接报错。注意nn.BCELoss的默认 reduction 是mean也就是对整个 batch 所有元素求平均。你可以通过reductionsum改成求和或用reductionnone得到每个样本单独的结果。这三个模式我们在第 4 章统一实验对比。3. BCEWithLogitsLoss为什么官方推荐它的底层逻辑3.1 一句话理解它做了什么事BCEWithLogitsLoss从名字就能看出来它把Sigmoid层和BCELoss合并成了一个操作。也就是说你用这个函数时模型的最后一层不需要额外加 sigmoid 激活函数直接把原始 logits 丢进去就行了loss_fn nn.BCEWithLogitsLoss() logits torch.tensor([[2.0], [-2.0], [1.5], [-0.8]]) # 不需要过 sigmoid targets torch.tensor([[1.], [0.], [1.], [0.]]) loss loss_fn(logits, targets) print(loss.item())同样的 logits如果你先用torch.sigmoid(logits)再喂给nn.BCELoss()得到的结果在数学上应该完全一致。这就会引出一个很自然的疑问既然结果一样为什么要多此一举融合起来3.2 数值稳定性的数学原理这个问题的关键在于“数值稳定性”。我们先看原始的 BCE 公式[ L -\left[ y \cdot \log(\sigma(x)) (1 - y) \cdot \log(1 - \sigma(x)) \right] ]其中 (\sigma(x) \frac{1}{1 e^{-x}}) 是 sigmoid 函数(x) 是 logits。问题出在哪当 (x) 是一个很大的负数时比如 (-100)(\sigma(x)) 会非常接近 0。在计算机里这可能会被舍入成精确的 0。随后 (\log(0)) 就变成了负无穷再乘上系数loss 就变成了 nan。反过来当 (x) 是一个很大的正数比如 100 时(\sigma(x)) 会非常接近 11 - sigma(x)可能被舍入成 0同样会导致 nan。BCEWithLogitsLoss在实现上不是先算 sigmoid 再取 log而是直接做了一个数学恒等变形。这里的关键步骤是[ \log(1 - \sigma(x)) \log\left(1 - \frac{1}{1 e^{-x}}\right) \log\left(\frac{e^{-x}}{1 e^{-x}}\right) ]再配合 LogSumExp 技巧最终 PyTorch 内部会使用类似以下稳定形式计算[ \max(x, 0) - x \cdot y \log(1 \exp(-|x|)) ]这个公式在 (x) 极大或极小时都不会出现中间变量被舍入为 0 的情况因此数值上是稳定的。这个设计思想是所有现代深度学习框架的通用做法你可以理解成它牺牲了一点点公式的“直观性”换取了计算过程中的稳定性。3.3 手工验证两者等价我们用一个小实验验证BCELoss配合sigmoid和BCEWithLogitsLoss的结果一致import torch import torch.nn as nn torch.manual_seed(42) logits torch.randn(8, 1) * 10 # 故意用较大的值容易触发数值问题 targets torch.randint(0, 2, (8, 1)).float() # 方法1BCEWithLogitsLoss 直接吃 logits loss_fn_1 nn.BCEWithLogitsLoss() loss1 loss_fn_1(logits, targets) # 方法2手动 sigmoid 后接 BCELoss probs torch.sigmoid(logits) loss_fn_2 nn.BCELoss() loss2 loss_fn_2(probs, targets) print(fBCEWithLogitsLoss: {loss1.item():.10f}) print(fSigmoid BCELoss: {loss2.item():.10f}) print(f差异: {abs(loss1.item() - loss2.item()):.2e})在我本机的运行结果里两者在小数值上几乎完全一致差异大概在 (10^{-8}) 量级但如果 logits 的绝对值特别大手写 sigmoid 后再算 BCELoss 的那条路径会有更高概率出现 nan。提示这也是为什么很多开源代码在最后一层不加nn.Sigmoid()而是直接用nn.BCEWithLogitsLoss。除了数值稳定训练结束后做推理时才临时加 sigmoid也是为了让训练/推理的解耦更干净——训练时模型输出 logits推理时再激活逻辑更清晰。3.4 与多标签分类的关系讲到这里顺带提一个高频场景多标签分类。比如一张图片里同时有“人”“车”“树”三个标签每个标签都是独立的二分类问题。这时候模型的输出头是多个节点每个节点代表一个标签是否出现。在这种情况下PyTorch 官方的推荐做法依然是用BCEWithLogitsLoss因为它的内部实现会自动对每个输出节点独立计算二分类交叉熵然后求平均。这就是为什么你会看到很多目标检测、多标签分类的项目里都在用它。4. 参数细节与实操对照实验4.1 两个函数的完整参数对比BCELoss和BCEWithLogitsLoss在参数层面有很多相似之处但有一个关键差异。先看这组对照表参数BCELossBCEWithLogitsLoss作用weight支持支持对每个样本/通道的 loss 加权size_average已弃用已弃用老版本控制是否求平均reduce已弃用已弃用老版本控制是否降维reduction支持支持mean/sum/nonepos_weight不支持支持正样本加权处理类别不平衡注意看最后一行的pos_weight这个参数只有BCEWithLogitsLoss才有。它专门用来解决正负样本数量不平衡的问题公式变为[ L -\left[ pos_weight \cdot y \cdot \log(\sigma(x)) (1 - y) \cdot \log(1 - \sigma(x)) \right] ]直白地说当正样本太少时把pos_weight设置成大于 1 的数相当于人为放大了正样本预测错误的惩罚让模型更重视正样本的学习。举个例子一个数据集里正样本占 10%负样本占 90%那么你设置pos_weight9就是一个非常常见的初始化选择它尽量让正负样本的累积损失贡献接近 1:1。4.2 实操演示三个 reduction 模式的差异我写一段代码把三种reduction模式的结果完整打印出来方便你直观理解import torch import torch.nn as nn logits torch.tensor([[1.5], [-0.5], [2.0], [-1.0]]) targets torch.tensor([[1.], [0.], [1.], [0.]]) # none: 返回每个样本各自的 loss loss_none nn.BCEWithLogitsLoss(reductionnone)(logits, targets) print(reductionnone:) print(loss_none) # mean: 所有样本 loss 的均值 loss_mean nn.BCEWithLogitsLoss(reductionmean)(logits, targets) print(freductionmean: {loss_mean.item():.4f}) # sum: 所有样本 loss 的求和 loss_sum nn.BCEWithLogitsLoss(reductionsum)(logits, targets) print(freductionsum: {loss_sum.item():.4f}) # 验证 mean 等于 none 求平均 print(f验证: none.mean() {loss_none.mean().item():.4f}) print(f验证: none.sum() {loss_none.sum().item():.4f})前向传播时默认用mean计算 loss 用于反向传播。验证 loss 是否合理时我会用reductionnone逐样本检查特别适合在调试时找出“哪些样本让模型非常困惑”。4.3 完整训练循环中的正确用法把理论放在一边我们看一个更接近实战的代码片段——用它可以跑通一个最简单的二分类训练循环import torch import torch.nn as nn import torch.optim as optim # 定义一个极简的模型最后一层没有 sigmoid model nn.Sequential( nn.Linear(10, 20), nn.ReLU(), nn.Linear(20, 1) # 输出 logits ) loss_fn nn.BCEWithLogitsLoss() optimizer optim.Adam(model.parameters(), lr0.001) # 造一点假数据 x torch.randn(16, 10) y torch.randint(0, 2, (16, 1)).float() for epoch in range(3): optimizer.zero_grad() logits model(x) loss loss_fn(logits, y) loss.backward() optimizer.step() print(fEpoch {epoch1}, Loss: {loss.item():.6f})在这个训练循环里模型最后一层没有任何激活函数输出直接是 logits。这非常关键——千万不要在模型里加了nn.Sigmoid()然后又用BCEWithLogitsLoss那就是“双重 sigmoid”会让模型极难收敛。反过来如果你用的是nn.BCELoss那模型最后一层就必须输出 sigmoid 之后的结果。注意推理阶段如果你需要用 0-1 之间的概率做阈值判断记得对 logits 手动torch.sigmoid()。此时不需要保留 sigmoid 的梯度放在torch.no_grad()环境里执行即可。4.4 处理正负样本不平衡时的实操样本不平衡在业务场景里太常见了比如点击率预测、异常检测、医疗图像中的罕见病症识别。用BCEWithLogitsLoss的pos_weight参数是成本最低的解法# 假设训练集中正样本 500 个负样本 4500 个 pos_weight torch.tensor([4500.0 / 500.0]) # 9.0 loss_fn nn.BCEWithLogitsLoss(pos_weightpos_weight)注意pos_weight的 shape 要能够 broadcast 到输出维度。对于单标签二分类一般是一个长度为 1 的 Tensor对于多标签分类可以传一个与标签数量相同的 Tensor为每个标签单独设置权重。设置pos_weight之后正样本的分类错误会被放大相当于你人为告诉模型“正样本很稀有看到了就必须抓住。”实际使用中pos_weight不需要完全严格等于负正样本比也可以把它当超参去调。5. 踩坑记录与排查技巧5.1 最常见的四个报错和解决方案这一章都是实打实的经验教训我按踩坑频率排序。第一个坑target类型用成torch.long。报错信息一般是这样的RuntimeError: result type Float cant be cast to the desired output type Long。解决办法很简单对标签做.float()转换。我见到过有人把.float()写在模型输出上这不对应该写在 target 上。第二个坑shape 不匹配。假设你的模型输出是(batch_size, 1)target 却是(batch_size,)在部分 PyTorch 版本里可能会自动广播但不报错算出来的 loss 数值却是错的。建议在 loss 前加一行断言assert logits.shape targets.shape, fshape mismatch: {logits.shape} vs {targets.shape}第三个坑模型里已经包含了 sigmoid然后又用BCEWithLogitsLoss。症状是 loss 一直在下降但永远下不到一个理想范围或者收敛非常慢。解决办法是检查模型最后一层去掉nn.Sigmoid()。第四个坑使用未经torchvision.transforms归一化处理的数据训练一开始 loss 就变成 nan。这不是损失函数本身的问题而是输入数据里可能包含 NaN 或者极大值导致 logits 发散进而触发数值不稳定。排查技巧是在 loss 计算后加一个torch.isnan(loss)断言快速定位出问题的是前向传播还是反向传播。5.2 如何判断你的 loss 是否正常很多初学者看到 loss 在 0.7 左右起伏就以为模型学崩了。我提供一个快速判断基线对于二分类问题随机初始化模型的 loss 大约在ln(2) ≈ 0.693附近。如果你的模型训练刚开始 loss 远低于 0.693说明初始化就偏向某类反而要检查一下是不是样本不均衡或者初始化有问题。在训练过程中如果 loss 在 0.3 以下稳步下降说明模型在学习如果 loss 直接跳到 nan 或无穷大往往是因为学习率过大导致梯度爆炸。这时候可以先降低学习率然后再考虑用clip_grad_norm_这类梯度裁剪手段。一个很实用的小技巧在训练集上抽样几百条数据用reductionnone逐条查看 loss把那些 loss 特别高的样本打印出来。这个习惯帮我解决过很多“看起来 loss 正常但准确率差”的疑难杂症。5.3 logits 与 probability 混淆的深水区最后分享一个我在多标签分类任务里经常遇到的问题。有人会把BCEWithLogitsLoss的输入误认为是“概率”于是提前对模型输出做了torch.softmax(dim1)。这里有两个错误多标签分类的每个类别是独立二分类类别之间互斥概率和为 1 的假设不成立应该用 sigmoid而不是 softmaxBCEWithLogitsLoss内部已经有 sigmoid你只需要确保输入是未激活的 logits 即可。如果你非要在模型里加 sigmoid那就改用BCELoss。这两种配置在数学上是等价的但在数值稳定性和速度上BCEWithLogitsLoss更优。提示写自定义网络时我习惯用命名来区分变量。模型输出命名为logits激活后命名为probs。这个小习惯能省下大量排查“到底传进去的是不是概率”的时间。写在最后的一个小技巧根据我个人的使用经验如果你正在搭建一个新项目二分类或多标签分类的损失函数可以直接无脑选nn.BCEWithLogitsLoss模型最后一层不加 sigmoidtarget 记得.float()。这个组合覆盖了 95% 的常见场景训练稳定、代码简洁、不容易踩数值坑。另外补充一点你可能会在某些古老的教程里看到nn.Sigmoid()加nn.BCELoss()的写法这本身没有错只是从 2020 年之后的 PyTorch 版本实践来看BCEWithLogitsLoss已经成为社区主流。如果你在维护旧代码看到这种写法也不用急着改只要模型能正常收敛两种方案都可以接受。以后遇到任何关于二分类损失函数“算出来是 nan”“预测概率全在 0.5 附近”“正负样本不收敛”的问题优先检查三件事logits 有没有重复 sigmoid、target 是不是 float 类型、pos_weight有没有设置。这三步排查完90% 的坑都能被填平。
返回列表