首页 / PyTorch 入门教程 / 损失函数:从 MSE 到交叉熵

PyTorch 入门教程

损失函数:从 MSE 到交叉熵

本教程共 60 篇 · 第 13 篇 · 更新于 2026-08-17 · 约 3 分钟阅读

PyTorch损失函数交叉熵MSE回归分类

本节目标:搞懂损失函数在训练中的角色,会用 PyTorch 内置的损失函数,知道回归和分类分别该选哪个。

损失函数是干嘛的

训练神经网络,说白了就是让模型预测得越来越准。那「准不准」总得有把尺子量吧?这把尺子就是损失函数(loss function)。

你可以把它想成考试判卷:模型交一份答案,损失函数负责打分,分数越低代表答得越好。优化器看到分数后,才知道该往哪个方向调整参数。

所以损失函数的选择,直接决定模型学到什么。同样一个模型,换个损失函数,结果可能天差地别。

统一用法:三步走

好消息是,PyTorch 里所有损失函数都长一个样。它们都是 nn.Module 的子类,用法完全统一:

  1. 实例化一个损失函数对象
  2. 调用它,传入「预测值在前,真实值在后」
  3. 对返回的标量调用 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

拿一组不同大小的误差对比一下:

误差MSELossL1LossSmoothL1Loss
1.01.001.000.50
5.025.005.004.50
10.0100.0010.009.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

BCELossBCEWithLogitsLoss 的标签必须是浮点数 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。先跑起来,再谈优化。