损失函数:从 MSE 到交叉熵
本教程共 60 篇 · 第 13 篇 · 更新于 2026-08-17 · 约 3 分钟阅读
本节目标:搞懂损失函数在训练中的角色,会用 PyTorch 内置的损失函数,知道回归和分类分别该选哪个。
损失函数是干嘛的
训练神经网络,说白了就是让模型预测得越来越准。那「准不准」总得有把尺子量吧?这把尺子就是损失函数(loss function)。
你可以把它想成考试判卷:模型交一份答案,损失函数负责打分,分数越低代表答得越好。优化器看到分数后,才知道该往哪个方向调整参数。
所以损失函数的选择,直接决定模型学到什么。同样一个模型,换个损失函数,结果可能天差地别。
统一用法:三步走
好消息是,PyTorch 里所有损失函数都长一个样。它们都是 nn.Module 的子类,用法完全统一:
- 实例化一个损失函数对象
- 调用它,传入「预测值在前,真实值在后」
- 对返回的标量调用
backward()
import torch
import torch.nn as nn
criterion = nn.CrossEntropyLoss()
# 预测值在前,标签在后
loss = criterion(predictions, targets)
# loss 是标量,直接反向传播
loss.backward()
Note参数顺序别写反。所有损失函数都是
loss(预测值, 真实值)。写反了能报错还算幸运,就怕它不报错。
至于函数内部怎么算,暂时不用管。先记住每种损失对「输入形态」的要求,这是新手最容易踩坑的地方。
回归:MSE 与它的兄弟们
先看最简单的场景:预测一个连续数字,比如房价、温度。这类任务叫回归(regression)。
MSELoss:均方误差
最经典的回归损失,算的是「预测值与真实值之差的平方,再取平均」。
criterion = nn.MSELoss()
predictions = torch.tensor([2.5, 0.5, 2.0, 8.0])
targets = torch.tensor([3.0, -0.5, 2.0, 7.0])
loss = criterion(predictions, targets)
print(f"MSE Loss: {loss.item():.4f}") # 0.5625
# 手动验证一下
manual = ((predictions - targets) ** 2).mean()
print(f"手动计算: {manual.item():.4f}") # 0.5625
误差 0.5 平方后是 0.25,误差 1.0 平方后是 1.0。看出门道了吗?MSE 会放大较大的误差,所以它对离群点(outlier)特别敏感。数据干净时它是首选。
L1Loss:平均绝对误差
L1 不做平方,直接取绝对值再平均。大误差不会被放大,对离群点更皮实。
criterion = nn.L1Loss()
# predictions 和 targets 沿用上面 MSE 例子的数据
loss = criterion(predictions, targets)
print(f"L1 Loss: {loss.item():.4f}") # 0.6250
同样一组数,MSE 报 0.5625,L1 报 0.6250。哪个更好没有定论,看你的数据里有没有捣乱的离群点。
SmoothL1Loss:两头占便宜
误差小时用平方(梯度平滑),误差大时用绝对值(抗离群点)。目标检测里回归边框位置,用的就是它。
criterion = nn.SmoothL1Loss()
# 还是上面那组数据
loss = criterion(predictions, targets)
print(f"SmoothL1 Loss: {loss.item():.4f}") # 0.2813
拿一组不同大小的误差对比一下:
| 误差 | MSELoss | L1Loss | SmoothL1Loss |
|---|---|---|---|
| 1.0 | 1.00 | 1.00 | 0.50 |
| 5.0 | 25.00 | 5.00 | 4.50 |
| 10.0 | 100.00 | 10.00 | 9.50 |
看见没?MSE 把 5 的误差放大成 25,SmoothL1 则温和得多。
分类:交叉熵一家
分类任务才是深度学习的绝对主力。想象模型输出三个分数,代表「猫、狗、鸟」的得分。怎么把这些分数变成损失?
CrossEntropyLoss:多分类标配
交叉熵损失(cross-entropy loss)是 PyTorch 里最常用的损失函数。它内部自动做了三件事:Softmax 归一化、取对数、取负值。
所以它要的是「原始分数」(logits),不是概率!
criterion = nn.CrossEntropyLoss()
# 模型输出的原始分数,形状 (batch_size, num_classes)
# 注意:不要手动做 Softmax!
predictions = torch.tensor([
[2.0, 0.5, 0.3], # 样本1,得分最高的是类别0
[0.1, 3.0, 0.2], # 样本2,得分最高的是类别1
[0.2, 0.1, 4.0], # 样本3,得分最高的是类别2
])
# 标签是整数类别索引,形状 (batch_size,)
targets = torch.tensor([0, 1, 2])
loss = criterion(predictions, targets)
print(f"Loss: {loss.item():.4f}") # 0.1640
模型猜得越准、越自信,损失就越小。
Tip传预测值时千万别先
softmax。CrossEntropyLoss 内部已经做过了,再做一次反而坏事。这是新手最常犯的错,没有之一。
BCEWithLogitsLoss:二分类首选
二分类任务(是/否、垃圾邮件判断)用二元交叉熵。PyTorch 给了两个版本:
BCELoss:要求输入已经是 Sigmoid 之后的概率(0~1),直接传原始分数会数值不稳定BCEWithLogitsLoss:内部自动做 Sigmoid,数值更稳,推荐直接用
criterion = nn.BCEWithLogitsLoss()
# 直接传原始分数,不用手动 Sigmoid
predictions = torch.tensor([2.0, -1.0, 0.5, -3.0])
targets = torch.tensor([1.0, 0.0, 1.0, 0.0]) # 注意:浮点标签
loss = criterion(predictions, targets)
print(f"Loss: {loss.item():.4f}") # 0.2407
Warning
BCELoss和BCEWithLogitsLoss的标签必须是浮点数1.0/0.0。传整数张量会直接报类型错。
多标签分类(一张图同时属于多个类别)也用它,每个标签独立判断,代码一行不用改。
NLLLoss:拆开的交叉熵
CrossEntropyLoss 内部其实就是 LogSoftmax + NLLLoss。只有当你需要中间步骤的对数概率时(比如做 Beam Search),才需要手动拆开用 NLLLoss。平时不用管它。
reduction 参数:怎么汇总
所有损失函数都带一个 reduction 参数,控制一批样本的损失怎么合并:
'mean'(默认):取平均'sum':求和'none':不合并,返回每个样本各自的损失
'none' 最实用,你可以给不同样本手动加权:
per_sample = nn.MSELoss(reduction='none')(predictions, targets)
weights = torch.tensor([1.0, 1.0, 2.0, 2.0]) # 后两个样本权重更高
weighted_loss = (per_sample * weights).mean()
另外,类别不平衡时可以用 CrossEntropyLoss(weight=...) 给少数类加权;语义分割里用 ignore_index 忽略边界像素。这两招实战中很常用。
自定义损失函数
内置的不够用?两种方式。
简单场景写个函数就行:
import torch.nn.functional as F
def focal_loss(predictions, targets, gamma=2.0, alpha=0.25):
# 让模型专注难分类的样本,目标检测常用
ce = F.cross_entropy(predictions, targets, reduction='none')
pt = torch.exp(-ce) # 预测正确的概率
return (alpha * (1 - pt) ** gamma * ce).mean()
复杂一点、需要带可学习参数的,就继承 nn.Module 写个类。这样它能和内置损失一样参与训练。
常见坑速查
- 交叉熵前多做了一次 Softmax → 损失降不下去
BCEWithLogitsLoss传了整数标签 → 直接报类型错- 累加 loss 时忘了
.item()→ 计算图越攒越大,显存爆掉
# 错误:loss 是张量,一直持有计算图
total_loss += loss
# 正确:取成 Python 数字再累加
total_loss += loss.item()
Tip选型口诀:多分类用
CrossEntropyLoss,二分类和多标签用BCEWithLogitsLoss,干净数据的回归用MSELoss,有离群点用SmoothL1Loss。先跑起来,再谈优化。