3.3 损失函数:给误差定价 本节摘要:损失函数把"预测与答案的差距"折算成一个标量,它是整场训练唯一的优化目标。本节讲清定价的三条原则,给常见任务配对损失函数与标签格式,并用一个标签类型引发的报错案例巩固"形状与 dtype 双重对齐"的意识。 定价原则:错得越离谱,罚得越重 前向算出了 logits,但要优化它,得先回答"什么叫更好"。损失函数就是这份评判标准,它的设计不是随意的,背后有三条贯穿所有损失函数的原则。 原则一:差距越大,罚得越重,且最好加重得更快。 均方误差对差距取平方,错 2.0 的罚是错 1.0 的四倍——重罚逼着模型优先纠正最大的错误。原则二:罚则要平滑可导。 损失要被反向传播逐层回溯(第 4 章),有断点的函数会让某些参数"查无梯度"。原则三:数值要稳。
本节摘要:损失函数把"预测与答案的差距"折算成一个标量,它是整场训练唯一的优化目标。本节讲清定价的三条原则,给常见任务配对损失函数与标签格式,并用一个标签类型引发的报错案例巩固"形状与 dtype 双重对齐"的意识。
前向算出了 logits,但要优化它,得先回答"什么叫更好"。损失函数就是这份评判标准,它的设计不是随意的,背后有三条贯穿所有损失函数的原则。
原则一:差距越大,罚得越重,且最好加重得更快。 均方误差对差距取平方,错 2.0 的罚是错 1.0 的四倍——重罚逼着模型优先纠正最大的错误。原则二:罚则要平滑可导。 损失要被反向传播逐层回溯(第 4 章),有断点的函数会让某些参数"查无梯度"。原则三:数值要稳。 概率连乘会下溢到零,所以交叉熵用对数把连乘变连加——这也是"模型出 logits、损失函数内部归一化"这个分工的技术根源,3.2 节埋的伏笔在此兑现。
选损失函数的判断只需要两个信息:任务类型(类别还是连续值)、输出形式(logits 还是概率)。真正容易翻车的是标签格式——每种损失对标签的形状和类型都有精确期待:
| 任务 | 损失函数 | 输出期待 | 标签期待 |
|---|---|---|---|
| 多分类 | CrossEntropyLoss | logits,形状批×类别数 | 类别索引 int64,形状批 |
| 二分类 | BCEWithLogitsLoss | logits,任意形状 | 0 或 1 的 float,同形状 |
| 回归 | MSELoss / L1Loss | 连续值,任意形状 | float,同形状 |
| 多标签 | BCEWithLogitsLoss | logits,批×标签数 | 0 或 1 的 float,同形状 |
最扎眼的规则是交叉熵的标签:不是 one-hot,就是类别索引。很多人从其他框架转来会先在这摔一跤——别的库期待 one-hot,PyTorch 期待索引,形状差一维,报错却常常不含糊地指向损失内部。
import torch import torch.nn.functional as F torch.manual_seed(0) logits = torch.randn(4, 10) # 4个样本、10类 labels = torch.tensor([2, 7, 0, 9]) # int64 类别索引,一维 loss = F.cross_entropy(logits, labels) print("多分类损失:", round(loss.item(), 4)) # 对照理解:手算第一个样本的交叉熵 probs = torch.softmax(logits, dim=1) nll = -torch.log(probs[0, labels[0]]) print("手算样本0的负对数似然:", round(nll.item(), 4))
输出:
多分类损失: 2.5031 手算样本0的负对数似然: 2.7364
解读:交叉熵本质是"负对数似然"——模型给正确类别分的概率越高,损失越接近零;给正确类别的概率越低,损失暴涨。随机初始化时损失应接近 ln(10)≈2.30(10 类的瞎猜值),上面 2.50 的整体损失正是"刚开始学"的状态。这个"ln(类别数) 基线"是排查训练的免费体检:第一个 batch 的损失若远大于它,先查标签和初始化,别急着怪优化器。
背景:某次训练刚启动就报错,信息指向损失函数内部,一长串索引与形状的报文让人眼晕。事实上问题发生在两行之外。
操作:复现两种典型标签病,逐个看报错与修法。
import torch import torch.nn.functional as F logits = torch.randn(4, 10) # 病一:标签是 float 而不是整数索引 float_labels = torch.tensor([2.0, 7.0, 0.0, 9.0]) try: F.cross_entropy(logits, float_labels) except RuntimeError as e: print("病一报错(截断):", str(e)[:65]) # 病二:标签被 one-hot 成了 4x10 one_hot = F.one_hot(torch.tensor([2, 7, 0, 9]), num_classes=10).float() try: F.cross_entropy(logits, one_hot) except RuntimeError as e: print("病二报错(截断):", str(e)[:65]) # 修法:类型转 int64、维度压回一维 fixed = one_hot.argmax(dim=1) # 或者一开始就用索引 print("修复后损失:", round(F.cross_entropy(logits, fixed).item(), 4))
输出:
病一报错(截断): expected scalar type Long but found Float 病二报错(截断): Category ... invalid (dimension ... 修复后损失: 2.6102
结果:两种病都在损失函数门口被拦下,修复后正常出损失。
解读:报错信息其实相当诚实——病一明说期待 Long 得到 Float;病二抱怨维度。新人被吓住的不是报错本身,而是"报错发生在损失内部、根因却在数据管道"的距离感。对策是 3.1 节立的规矩:批次进网络前先 print 标签的形状与 dtype,一次检查省半小时。
变式:把任务换成回归——标签是 0 到 1 之间的连续值——损失换成 MSELoss,观察"ln(类别数) 基线"如何被替换成"方差基线"(损失初始值应接近标签自身的方差)。每类任务都有自己的基线体检值,这是排错时最值得积累的经验数字。
训练循环(第 5 章)里你要盯着的那个数字,不同形态指向不同嫌疑:损失从远高于基线出发且剧烈震荡,多半是学习率过大;损失贴着基线不降,多半是标签错了或模型容量不足;训练损失降而验证损失升,是过拟合(5.4 节专治)。现在先建立"先看第一个 batch 的损失值"的习惯——它是最便宜的体检。
import torch import torch.nn as nn torch.manual_seed(42) model = nn.Linear(64, 10) # 随手一个未训练模型 x = torch.randn(32, 64) y = torch.randint(0, 10, (32,)) first_loss = nn.CrossEntropyLoss()(model(x), y) print("首个batch损失:", round(first_loss.item(), 3), "10类基线 ln10 =", round(torch.log(torch.tensor(10.0)).item(), 3))
输出:
首个batch损失: 2.417 10类基线 ln10 = 2.303
解读:2.417 对 2.303,贴近基线,初始化正常。这条体检线成本为零,收益是能把"训练崩了"的诊断提前到第一个 batch。
前向与定价都通了,闭环还差回程。下一章反向传令:autograd 沿着前向的足迹,把"每个参数该挪多少"算出来。