共计 1798 个字符,预计需要花费 5 分钟才能阅读完成。
背景痛点:损失函数为什么重要?
在 BP 神经网络(Backpropagation Neural Network)的训练过程中,损失函数(Loss Function)就像导航仪一样,负责告诉模型当前预测结果与真实值之间的差距。但实际应用中常遇到两个典型问题:

- 梯度消失(Vanishing Gradient):当使用 Sigmoid 激活函数配合均方误差时,深层网络的梯度会指数级减小,导致权重更新停滞。
- 过拟合(Overfitting):模型在训练集上损失值持续下降,但测试集表现变差,说明损失函数未能有效约束模型复杂度。
技术对比:交叉熵 vs 均方误差
数学特性比较
-
均方误差(MSE/Mean Squared Error):
$$L_{MSE} = \frac{1}{N}\sum_{i=1}^N(y_i – \hat{y}_i)^2$$
适用于回归任务,对异常值敏感,输出层通常无激活函数。 -
交叉熵损失(CrossEntropy):
$$L_{CE} = -\sum_{i=1}^C y_i\log(\hat{y}_i)$$
专为分类任务设计,与 Softmax 激活函数天然配对,能处理概率分布差异。
适用场景
- MSE 更适合:
- 房价预测等回归问题
-
输出值范围无约束的场景
-
CrossEntropy 更适合:
- 图像分类等离散输出任务
- 需要概率解释的场景(如医疗诊断)
核心实现:PyTorch 代码实战
交叉熵损失实现
import torch
import torch.nn as nn
# 模拟 10 个样本的 3 分类问题
outputs = torch.randn(10, 3) # [batch_size, num_classes]
labels = torch.randint(0, 3, (10,)) # 真实类别索引
# 关键点 1:无需手动 Softmax,CrossEntropyLoss 内部包含
criterion = nn.CrossEntropyLoss()
loss = criterion(outputs, labels)
# 反向传播演示
loss.backward() # 自动计算梯度
均方误差实现
# 回归任务示例
true_values = torch.randn(10, 1) # 真实值
predictions = torch.randn(10, 1) # 预测值
mse_loss = nn.MSELoss()
loss = mse_loss(predictions, true_values)
# 学习率影响演示
optimizer = torch.optim.SGD(model.parameters(), lr=0.01) # 过高学习率会导致震荡
调优指南:三大实战技巧
- 学习率动态调整
- 使用
torch.optim.lr_scheduler.ReduceLROnPlateau根据损失值自动调节 -
交叉熵损失通常需要比 MSE 更小的初始学习率
-
正则化项添加
# L2 正则化(权重衰减)optimizer = torch.optim.Adam(model.parameters(), weight_decay=1e-4) # L1 正则化需要手动实现 l1_lambda = 0.001 l1_norm = sum(p.abs().sum() for p in model.parameters()) loss = criterion(outputs, labels) + l1_lambda * l1_norm -
类别不平衡处理
# 假设类别权重比为 [1, 2, 5](第二类是第三类的 2.5 倍)weights = torch.tensor([1, 2.5, 1]) criterion = nn.CrossEntropyLoss(weight=weights)
避坑建议:两个典型错误
- 分类任务误用 MSE
- 会导致梯度更新方向与分类目标不一致
-
表现:准确率卡在随机猜测水平
-
激活函数不匹配
- Softmax 输出配 MSE 损失:概率计算与误差衡量方式冲突
- Sigmoid 输出配 CrossEntropy:需要改用
BCEWithLogitsLoss
延伸思考:开放性问题
- 电商推荐系统中,如何设计损失函数同时优化点击率和购买率?
- 对比 Adam 优化器 +SGD 优化器在相同损失函数下的收敛曲线差异
- 当验证集损失下降但准确率不升时,可能是什么原因?
个人实践心得
最近在医疗影像分类项目中,发现 Dice Loss 比传统 CrossEntropy 更适合小目标分割。建议大家在选定损失函数前,先用小样本测试不同损失函数的梯度更新行为。调参时可以先用 torchviz 可视化计算图,确保反向传播路径符合预期。
正文完
