深度学习损失函数实战指南:从理论到代码实现
1. 损失函数:模型学习的“导航仪”
如果你刚开始接触深度学习,可能会觉得损失函数是个挺抽象的概念。我第一次接触时也是一头雾水,直到我把它想象成开车时的导航仪,才真正理解了它的作用。想象一下,你开车去一个陌生的地方,导航仪会实时告诉你“偏离路线了”或者“正在最优路线上”。损失函数干的就是这个活儿——它不停地告诉模型:“嘿,你现在的预测离正确答案还有多远。”
这个“距离”的量化,就是损失函数的核心。它不是一个固定的公式,而是一整套工具箱。不同的任务,就像去不同的目的地(比如去市中心购物和去郊外爬山),需要选择不同的导航策略(比如是优先考虑时间最短,还是避开拥堵)。在深度学习中,回归任务(预测一个连续值,比如房价、温度)和分类任务(判断一张图片是猫还是狗)就是两种完全不同的“目的地”,自然需要不同的“导航仪”——也就是不同的损失函数。
选择对了损失函数,模型训练就能事半功倍,收敛又快又好;选错了,可能就像用步行导航去规划跨国航线,怎么调参都感觉不对劲,模型性能死活上不去。我刚开始做图像分类时,曾经在一个二分类问题上错误地使用了多分类交叉熵损失,结果模型训练时损失值震荡得厉害,准确率一直卡在50%左右(跟瞎猜一样),排查了好久才发现是损失函数用错了。这个坑踩过之后,我就特别重视损失函数的选择了。
所以,这篇指南的目的,就是帮你把这个“导航仪”工具箱彻底搞清楚。我们不只讲每个工具长什么样(数学公式),更要讲清楚它最适合在什么路况下用(适用场景),并且手把手带你把它装到你的“车”(模型)上(代码实现)。你会发现,一旦理解了损失函数,你对模型训练过程的掌控力会大大增强。
2. 回归任务:预测一个具体的数
回归任务的目标是预测一个连续值。比如,根据房屋面积、地段预测房价,或者根据历史数据预测明天的气温。这时候,我们的损失函数核心是衡量“预测值”和“真实值”之间的数值差距。
2.1 均方误差:最常用的“尺子”
均方误差,简称MSE,这可能是你最早接触、也最常用的损失函数。它的思想非常直接:把所有预测值和真实值之间的差,先平方(去掉负号,同时放大大的误差),再求平均。
数学公式很简单:MSE = (1/n) * Σ(y_true - y_pred)^2
我习惯把它理解为一种“严厉的考官”。因为平方项的存在,如果一个预测错得特别离谱(比如误差是10),那么它对整体损失的“贡献”是100;而一个普通的小错误(误差是1),贡献只有1。这意味着MSE对异常值(那些偏离特别大的数据点)非常敏感。模型会拼命去修正那些错得最离谱的点,有时甚至会以牺牲其他大多数点的准确性为代价。
什么时候用MSE最合适? 当你的数据里噪声比较小,没有太多极端异常值,并且你认为大的误差应该受到远比小误差更严厉的惩罚时,MSE是很好的选择。它在很多理论推导中也特别方便,因为平方函数处处可导,性质良好。
下面是一个用NumPy和PyTorch分别实现的例子,你可以直观地感受一下:
import numpy as np
import torch
import torch.nn as nn
# 真实值和预测值
y_true_np = np.array([3.0, -0.5, 2.0, 7.0])
y_pred_np = np.array([2.5, 0.0, 2.0, 8.0]) # 最后一个预测偏差较大
# NumPy 手动实现
def mse_loss_np(y_true, y_pred):
return np.mean((y_true - y_pred) ** 2)
print("NumPy MSE:", mse_loss_np(y_true_np, y_pred_np))
# PyTorch 实现(更贴近实际训练)
y_true_torch = torch.tensor([3.0, -0.5, 2.0, 7.0])
y_pred_torch = torch.tensor([2.5, 0.0, 2.0, 8.0])
loss_fn = nn.MSELoss() # 直接调用PyTorch内置函数
loss = loss_fn(y_pred_torch, y_true_torch) # 注意参数顺序:预测值在前,真实值在后
print("PyTorch MSE Loss:", loss.item())
运行这段代码,你会看到一个具体的损失值。试着把最后一个预测值从8.0改成12.0,再看看损失值的变化,你就能切身感受到MSE对大幅误差的“严厉惩罚”了。
2.2 平均绝对误差与Huber Loss:应对“坏数据”的稳健选择
如果你的数据里可能混入了一些异常值(比如传感器偶尔的故障读数,或者数据录入时的错误),MSE的“严厉”就会变成缺点。它会导致模型变得不稳定。这时,平均绝对误差(MAE,也叫L1 Loss)就派上用场了。
MAE的公式是:MAE = (1/n) * Σ|y_true - y_pred|
MAE就像一位“公平的裁判”,它只关心误差的绝对值大小。误差是10,贡献就是10;误差是1,贡献就是1。它对异常值的敏感度远低于MSE。在存在异常值的数据集上,使用MAE训练的模型通常会更加稳健。
但是,MAE也有自己的问题:它在零点处不可导(因为绝对值函数在零点有个“尖”),这在利用梯度下降法优化时可能会带来一些小麻烦,虽然现代深度学习框架都能很好地处理。另一个特点是,MAE的梯度大小是恒定的(正负1),这可能导致模型在损失值已经很小时收敛变慢。
有没有一种损失函数能取二者之长呢?有的,这就是Huber Loss。它是我处理回归问题,尤其是数据质量不确定时的首选。
Huber Loss的设计很巧妙:它设定一个阈值 δ(delta)。当预测误差小于δ时,它采用MSE的行为(二次函数),保证在误差小时收敛更精确、更快;当误差大于δ时,它切换成MAE的行为(线性函数),以减少异常值的过大影响。
它的公式分段表示如下:
- 如果 |y_true - y_pred| <= δ:
Loss = 0.5 * (y_true - y_pred)^2 - 如果 |y_true - y_pred| > δ:
Loss = δ * |y_true - y_pred| - 0.5 * δ^2
你可以把δ看作一个“敏感度”旋钮。δ设得越大,损失函数就越像MAE(更鲁棒);δ设得越小,就越像MSE(对误差更敏感)。通常,δ=1.0是一个不错的起点。
import numpy as np
def huber_loss_np(y_true, y_pred, delta=1.0):
"""
手动实现Huber Loss
"""
error = y_true - y_pred
abs_error = np.abs(error)
# 核心:分段函数
quadratic_part = np.minimum(abs_error, delta)
linear_part = abs_error - quadratic_part
loss = 0.5 * quadratic_part ** 2 + delta * linear_part
return np.mean(loss)
# 测试数据,故意加入一个异常值
y_true = np.array([1.0, 2.0, 3.0, 4.0, 5.0])
y_pred = np.array([1.1, 1.9, 3.2, 4.1, 10.0]) # 最后一个预测是离谱的异常
mse = np.mean((y_true - y_pred) ** 2)
mae = np.mean(np.abs(y_true - y_pred))
huber = huber_loss_np(y_true, y_pred, delta=1.0)
print(f"MSE (对异常值敏感): {mse:.4f}")
print(f"MAE (相对稳健): {mae:.4f}")
print(f"Huber Loss (delta=1.0): {huber:.4f}")
# 使用PyTorch内置的Huber Loss
import torch.nn as nn
huber_loss_torch = nn.HuberLoss(delta=1.0)
loss_torch = huber_loss_torch(torch.tensor(y_pred, dtype=torch.float32),
torch.tensor(y_true, dtype=torch.float32))
print(f"PyTorch Huber Loss: {loss_torch.item():.4f}")
运行代码,对比三个损失值。你会发现MSE的值被那个“10.0”的异常预测拉得很高,而MAE和Huber Loss的值则相对温和很多。这就是Huber Loss的实用之处——在保持可导性和优化效率的同时,获得了鲁棒性。
3. 分类任务:判断属于哪一类
分类任务与回归任务有本质区别。模型通常输出一个概率分布,表示输入样本属于各个类别的可能性。损失函数的核心变成了衡量两个概率分布(模型预测分布 vs 真实标签分布)之间的差异。
3.1 交叉熵损失:分类任务的“黄金标准”
交叉熵是分类任务中毋庸置疑的基石。要理解它,可以打个比方:真实标签分布是一份“标准答案”(对于一张猫的图片,猫类概率为1,狗类概率为0),模型预测分布是你的“答卷”。交叉熵衡量的是你按照自己的“答卷”去估计“标准答案”的编码长度,所需的额外比特数。这个值越小,说明你的“答卷”越接近“标准答案”。
对于二分类任务(比如判断邮件是垃圾邮件还是正常邮件),我们常用二元交叉熵。假设真实标签y_true是0或1,模型预测出它是正类的概率为p,那么损失函数为:
Loss = - [y_true * log(p) + (1 - y_true) * log(1 - p)]
这个公式很优雅:当真实标签是1时,只有前半部分生效,我们希望p越大越好(log(p)越大,负得越少);当真实标签是0时,只有后半部分生效,我们希望p越小越好(log(1-p)越大,负得越少)。
对于多分类任务(比如识别手写数字0-9),我们使用多元交叉熵。模型会通过一个Softmax层输出一个概率向量(所有类别概率之和为1)。损失函数是:
Loss = - Σ (y_true_i * log(y_pred_i))
其中y_true_i是样本属于第i类的真实概率(通常是one-hot编码,即正确类别为1,其余为0),y_pred_i是模型预测的第i类的概率。
import torch
import torch.nn as nn
import torch.nn.functional as F
# 示例1:二分类交叉熵 (BCE Loss)
# 假设我们有一个批量大小为3的二分类任务
sigmoid_output = torch.tensor([0.8, 0.2, 0.6]) # 模型经过Sigmoid后的输出,表示正类概率
labels = torch.tensor([1.0, 0.0, 1.0]) # 真实标签
# 使用内置函数,注意输入需要是浮点型,且形状一致
bce_loss = nn.BCELoss()
loss_bce = bce_loss(sigmoid_output, labels)
print(f"Binary Cross-Entropy Loss: {loss_bce.item():.4f}")
# 更常见的做法:模型输出未归一化的logits,使用BCEWithLogitsLoss(更数值稳定)
logits = torch.tensor([1.5, -1.2, 0.5]) # Sigmoid前的原始输出
bce_logits_loss = nn.BCEWithLogitsLoss()
loss_bce_logits = bce_logits_loss(logits, labels)
print(f"BCEWithLogits Loss: {loss_bce_logits.item():.4f}")
# 示例2:多分类交叉熵 (CrossEntropyLoss)
# 假设一个3分类任务,批量大小为2
# 模型输出的是未经过Softmax的logits (raw scores)
logits = torch.tensor([[2.0, 0.5, -1.0], # 第一个样本的logits
[0.5, 1.5, 0.2]]) # 第二个样本的logits
# 真实标签是类别索引,不是one-hot
targets = torch.tensor([0, 2]) # 第一个样本属于第0类,第二个属于第2类
ce_loss = nn.CrossEntropyLoss() # 这个函数内部包含了Softmax和交叉熵计算
loss_ce = ce_loss(logits, targets)
print(f"Multi-class Cross-Entropy Loss: {loss_ce.item():.4f}")
# 手动验证一下过程,加深理解
# 对logits做Softmax得到概率
probs = F.softmax(logits, dim=1)
print("Predicted probabilities:\n", probs)
# 根据真实标签索引取出对应概率
selected_probs = probs[torch.arange(2), targets] # probs[0,0] 和 probs[1,2]
print("Probabilities of true classes:", selected_probs)
# 计算负对数似然
manual_loss = -torch.log(selected_probs).mean()
print(f"Manual calculation loss: {manual_loss.item():.4f}")
多分类交叉熵是实践中使用最广泛的损失函数之一。nn.CrossEntropyLoss这个PyTorch函数非常方便,它接受的是未归一化的logits和类别索引,内部帮你完成了Softmax和交叉熵计算,并且做了数值稳定性处理,是你训练分类网络时最常打交道的伙伴。
3.2 Focal Loss:解决“简单样本”与“类别不平衡”的双重难题
标准的交叉熵损失有一个潜在问题:它对所有样本“一视同仁”。但在目标检测等任务中,一张图片里可能包含大量背景(负样本),只有少数几个物体(正样本),这就是严重的类别不平衡。而且,大部分背景区域是“简单样本”(模型很容易判断为背景),它们虽然单个损失小,但数量巨大,累积起来会淹没掉少数但重要的正样本的梯度。
Focal Loss的提出就是为了解决这个问题。它的核心思想是:降低那些已经被模型很好分类的“简单样本”的权重,让模型更专注于学习那些难分的样本。
它的公式是在标准交叉熵基础上加了一个调制因子:FL(p_t) = -α_t * (1 - p_t)^γ * log(p_t)
p_t:对于正样本,是模型预测其为正类的概率;对于负样本,是模型预测其为负类的概率。p_t越大,说明模型预测得越有信心(越简单)。(1 - p_t)^γ:这就是调制因子。当样本被分对且p_t很大(接近1)时,(1-p_t)接近0,整个调制因子就接近0,从而大幅降低该样本的损失权重。γ(gamma)是一个可调参数,越大,对简单样本的抑制就越强。α_t:一个用于平衡正负样本权重的因子,可以缓解类别不平衡。通常为正样本设置一个较小的α(如0.25),负样本设置较大的α(如0.75)。
简单来说,Focal Loss让模型“看不起”那些它已经会了的题,逼着它去啃“硬骨头”。
import torch
import torch.nn as nn
import torch.nn.functional as F
class FocalLoss(nn.Module):
"""
二分类Focal Loss的实现
"""
def __init__(self, alpha=0.25, gamma=2.0, reduction='mean'):
super(FocalLoss, self).__init__()
self.alpha = alpha
self.gamma = gamma
self.reduction = reduction
def forward(self, inputs, targets):
# inputs: 模型输出的logits (未经Sigmoid)
# targets: 真实标签 (0或1)
BCE_loss = F.binary_cross_entropy_with_logits(inputs, targets, reduction='none')
pt = torch.exp(-BCE_loss) # pt = p if y=1, else pt = 1-p
# 计算alpha_t
alpha_t = targets * self.alpha + (1 - targets) * (1 - self.alpha)
# Focal Loss
FL = alpha_t * (1 - pt) ** self.gamma * BCE_loss
if self.reduction == 'mean':
return FL.mean()
elif self.reduction == 'sum':
return FL.sum()
else:
return FL
# 模拟一个严重类别不平衡的场景:100个负样本,10个正样本
batch_size = 110
logits = torch.randn(batch_size) * 0.5 # 随机logits
# 前100个为负样本(标签0),后10个为正样本(标签1)
targets = torch.cat([torch.zeros(100), torch.ones(10)])
# 标准交叉熵损失
bce_loss = nn.BCEWithLogitsLoss(reduction='none')
loss_bce = bce_loss(logits, targets)
print(f"Standard BCE Loss (mean): {loss_bce.mean().item():.4f}")
print(f" - Loss on negative samples (first 100): {loss_bce[:100].mean().item():.4f}")
print(f" - Loss on positive samples (last 10): {loss_bce[100:].mean().item():.4f}")
# Focal Loss
focal_loss_fn = FocalLoss(alpha=0.25, gamma=2.0)
loss_focal = focal_loss_fn(logits, targets)
print(f"\nFocal Loss (gamma=2, alpha=0.25): {loss_focal.item():.4f}")
# 我们可以看看调制因子(1-pt)^gamma的效果
with torch.no_grad():
probs = torch.sigmoid(logits)
pt = torch.where(targets == 1, probs, 1 - probs) # 计算pt
modulating_factor = (1 - pt) ** 2.0
print(f"\nModulating factor (1-pt)^gamma for a few samples:")
print(f" Sample 1 (easy negative, prob={probs[0]:.3f}): factor = {modulating_factor[0]:.5f}")
print(f" Sample 105 (hard positive, prob={probs[104]:.3f}): factor = {modulating_factor[104]:.5f}")
运行这段代码,你可以观察到,对于预测概率很高(很确信)的简单负样本,其调制因子非常小(例如0.0001),导致其损失贡献几乎被忽略。而对于预测概率不高(难分)的正样本,调制因子较大(例如0.5),保留了其大部分损失贡献。这样,在反向传播时,难样本的梯度就占据了主导,有效缓解了类别不平衡带来的问题。
4. 特殊任务与进阶损失函数
除了回归和分类这两大主流任务,深度学习还有很多细分领域,它们有自己独特的挑战,也催生了一些专用的损失函数。
4.1 Dice Loss与IoU Loss:图像分割的“专属武器”
图像分割任务(比如把医学影像中的肿瘤区域抠出来)的评价指标通常是交并比,也就是IoU。一个很自然的想法是:能不能直接用IoU作为损失函数来优化呢?这就是IoU Loss的初衷。而Dice Loss与IoU Loss在数学上高度相关,都特别适合处理图像分割中常见的前景-背景像素极度不平衡的问题(一张图里大部分是背景,肿瘤只占一小部分)。
Dice系数的定义是:Dice = (2 * |A ∩ B|) / (|A| + |B|),其中A是真实分割区域(Ground Truth),B是预测区域。Dice系数衡量的是两个集合的重叠程度,取值范围[0, 1],值越大越好。Dice Loss就是 1 - Dice。
它的优点在于,它对集合的大小不敏感。即使目标区域很小,只要预测区域和真实区域重叠得好,Dice系数依然可以很高。而像交叉熵这类逐像素计算的损失,可能会因为背景像素太多而完全忽略掉小目标。
import torch
import torch.nn as nn
import torch.nn.functional as F
def dice_loss(pred, target, smooth=1e-6):
"""
计算二分类Dice Loss
pred: 模型预测的概率图 (经过Sigmoid,值在0~1之间),形状 [N, H, W]
target: 真实二值分割图 (0或1),形状 [N, H, W]
smooth: 平滑项,防止分母为0
"""
# 将预测图展平
pred_flat = pred.contiguous().view(-1)
target_flat = target.contiguous().view(-1)
# 计算交集
intersection = (pred_flat * target_flat).sum()
# 计算 Dice 系数
dice = (2. * intersection + smooth) / (pred_flat.sum() + target_flat.sum() + smooth)
# 返回 Dice Loss
return 1 - dice
# 模拟一个简单的分割任务:预测一个5x5图像中的小方块
batch_size = 2
height = width = 5
# 真实标签:中间一个3x3的方块为前景(1),其余为背景(0)
target = torch.zeros(batch_size, height, width)
target[:, 1:4, 1:4] = 1.0
# 模拟模型预测:预测得比较好,但不够精确
pred = torch.zeros(batch_size, height, width)
pred[:, 1:4, 1:4] = 0.8 # 预测了方块,但置信度0.8
pred[:, 2, 2] = 0.9 # 中心点置信度更高
print("Target (one sample):")
print(target[0])
print("\nPrediction (one sample, probability):")
print(pred[0])
# 计算Dice Loss
loss_dice = dice_loss(pred, target)
print(f"\nDice Loss: {loss_dice.item():.4f}")
# 对比一下二值交叉熵损失
# 注意:BCE需要sigmoid输入,我们的pred已经是概率,但为了公平对比,我们假设logits是通过logit函数反推的(仅作演示)
# 这里我们直接对概率图计算BCE,实际中应对logits计算BCEWithLogitsLoss
loss_bce = F.binary_cross_entropy(pred, target)
print(f"Binary Cross-Entropy Loss: {loss_bce.item():.4f}")
# 假设另一个模型预测了整个图像都是背景(全0),这是很坏的情况
pred_bad = torch.zeros_like(pred)
loss_dice_bad = dice_loss(pred_bad, target)
loss_bce_bad = F.binary_cross_entropy(pred_bad, target)
print(f"\n--- Bad Prediction (all background) ---")
print(f"Dice Loss: {loss_dice_bad.item():.4f} (接近1,很差)")
print(f"BCE Loss: {loss_bce_bad.item():.4f}")
你会发现,当预测完全错误(全背景)时,Dice Loss会直接飙升至接近1(最差情况),而BCE Loss由于大部分背景像素都预测正确了(背景为0预测也为0),其损失值反而可能没那么大。这直观地展示了Dice Loss对于前景-背景不平衡问题的针对性:它更关注于前景区域预测的准确性。
4.2 对比损失与三元组损失:让模型学会“区分”
在人脸识别、图像检索、语义相似度等任务中,我们并不直接分类,而是希望模型学习一个“特征空间”,在这个空间里,同类样本的特征彼此靠近,不同类样本的特征彼此远离。这就需要用到度量学习中的损失函数,对比损失和三元组损失是其中的经典代表。
对比损失的思想是成对比较。给定一个样本对(两张图片)和它们的标签(是否属于同一类),损失函数鼓励同类样本的特征向量距离小,异类样本的特征向量距离大,并且要大于一个预设的边界值margin。
三元组损失则更进一步,它同时考虑三个样本:一个锚点样本(Anchor)、一个正样本(与Anchor同类)、一个负样本(与Anchor不同类)。损失函数的目标是,让锚点与正样本的距离,比锚点与负样本的距离,至少小一个margin。这样学习到的特征区分度更强。
import torch
import torch.nn as nn
import torch.nn.functional as F
def contrastive_loss(feat1, feat2, label, margin=1.0):
"""
对比损失实现
feat1, feat2: 样本对的特征向量 [batch_size, feature_dim]
label: 1表示同类,0表示不同类
margin: 边界值
"""
euclidean_dist = F.pairwise_distance(feat1, feat2, p=2) # 计算欧氏距离
# 同类样本,损失就是距离;异类样本,损失是max(0, margin - distance)
loss_same = label * torch.pow(euclidean_dist, 2)
loss_diff = (1 - label) * torch.pow(torch.clamp(margin - euclidean_dist, min=0.0), 2)
loss = torch.mean(loss_same + loss_diff)
return loss
def triplet_loss(anchor, positive, negative, margin=1.0):
"""
三元组损失实现
anchor, positive, negative: 锚点、正样本、负样本的特征向量 [batch_size, feature_dim]
margin: 边界值
"""
pos_dist = F.pairwise_distance(anchor, positive, p=2)
neg_dist = F.pairwise_distance(anchor, negative, p=2)
# 核心公式:让 pos_dist - neg_dist + margin <= 0
loss = torch.mean(torch.clamp(pos_dist - neg_dist + margin, min=0.0))
return loss
# 模拟数据
batch_size = 4
feat_dim = 128
# 生成特征向量
feat1 = torch.randn(batch_size, feat_dim)
feat2 = torch.randn(batch_size, feat_dim)
# 随机生成标签:0或1
labels = torch.randint(0, 2, (batch_size,)).float()
print("Sample Pair Labels:", labels)
loss_contrastive = contrastive_loss(feat1, feat2, labels, margin=1.0)
print(f"Contrastive Loss: {loss_contrastive.item():.4f}")
# 三元组损失示例
anchor = torch.randn(batch_size, feat_dim)
positive = anchor + torch.randn(batch_size, feat_dim) * 0.1 # 正样本靠近锚点
negative = torch.randn(batch_size, feat_dim) # 负样本随机
loss_triplet = triplet_loss(anchor, positive, negative, margin=0.5)
print(f"\nTriplet Loss: {loss_triplet.item():.4f}")
# 计算一下距离,看看是否符合预期
with torch.no_grad():
pos_d = F.pairwise_distance(anchor, positive, p=2).mean()
neg_d = F.pairwise_distance(anchor, negative, p=2).mean()
print(f"Avg distance (Anchor-Positive): {pos_d.item():.4f}")
print(f"Avg distance (Anchor-Negative): {neg_d.item():.4f}")
print(f"Difference (Pos - Neg): {(pos_d - neg_d).item():.4f}")
在实际应用中,构造“困难三元组”(即那些pos_dist很大或neg_dist很小的三元组)对训练至关重要,因为简单的三元组(已经满足neg_dist > pos_dist + margin)产生的损失为0,对模型没有贡献。这就是“困难样本挖掘”策略。
5. 代码实战:在真实训练流程中集成损失函数
了解了这么多理论,最后我们来看一个完整的、贴近实战的例子:训练一个简单的卷积神经网络在CIFAR-10数据集上进行图像分类,并对比使用不同损失函数(标准交叉熵 vs Label Smoothing交叉熵)的效果。
Label Smoothing 是一种常用的正则化技术,用于缓解模型对训练标签的“过度自信”。它把原始的one-hot标签(如[0, 0, 1, 0])稍微“软化”,给非目标类别分配一个很小的概率(如ε),目标类别概率变为1 - ε(如[0.01, 0.01, 0.97, 0.01])。这可以防止模型在训练集上过度拟合,有时能提升模型的泛化能力。
import torch
import torch.nn as nn
import torch.optim as optim
import torchvision
import torchvision.transforms as transforms
from torch.utils.data import DataLoader
import matplotlib.pyplot as plt
# 1. 定义带Label Smoothing的交叉熵损失
class LabelSmoothingCrossEntropy(nn.Module):
def __init__(self, smoothing=0.1, reduction='mean'):
super().__init__()
self.smoothing = smoothing
self.reduction = reduction
def forward(self, logits, targets):
# logits: [N, C]
# targets: [N] 类别索引
num_classes = logits.size(-1)
# 将标签转换为one-hot,并应用平滑
with torch.no_grad():
targets_onehot = torch.zeros_like(logits)
targets_onehot.scatter_(1, targets.unsqueeze(1), 1)
smoothed_targets = targets_onehot * (1 - self.smoothing) + self.smoothing / num_classes
# 计算交叉熵
log_probs = F.log_softmax(logits, dim=-1)
loss = - (smoothed_targets * log_probs).sum(dim=-1)
if self.reduction == 'mean':
return loss.mean()
elif self.reduction == 'sum':
return loss.sum()
else:
return loss
# 2. 准备数据
transform = transforms.Compose([
transforms.ToTensor(),
transforms.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))
])
trainset = torchvision.datasets.CIFAR10(root='./data', train=True, download=True, transform=transform)
trainloader = DataLoader(trainset, batch_size=128, shuffle=True, num_workers=2)
testset = torchvision.datasets.CIFAR10(root='./data', train=False, download=True, transform=transform)
testloader = DataLoader(testset, batch_size=128, shuffle=False, num_workers=2)
# 3. 定义一个简单的CNN模型
class SimpleCNN(nn.Module):
def __init__(self):
super().__init__()
self.conv1 = nn.Conv2d(3, 32, 3, padding=1)
self.pool = nn.MaxPool2d(2, 2)
self.conv2 = nn.Conv2d(32, 64, 3, padding=1)
self.fc1 = nn.Linear(64 * 8 * 8, 256)
self.fc2 = nn.Linear(256, 10)
self.relu = nn.ReLU()
self.dropout = nn.Dropout(0.3)
def forward(self, x):
x = self.pool(self.relu(self.conv1(x)))
x = self.pool(self.relu(self.conv2(x)))
x = x.view(-1, 64 * 8 * 8)
x = self.relu(self.fc1(x))
x = self.dropout(x)
x = self.fc2(x)
return x
# 4. 训练函数
def train_one_epoch(model, dataloader, criterion, optimizer, device):
model.train()
running_loss = 0.0
correct = 0
total = 0
for inputs, labels in dataloader:
inputs, labels = inputs.to(device), labels.to(device)
optimizer.zero_grad()
outputs = model(inputs)
loss = criterion(outputs, labels)
loss.backward()
optimizer.step()
running_loss += loss.item()
_, predicted = outputs.max(1)
total += labels.size(0)
correct += predicted.eq(labels).sum().item()
epoch_loss = running_loss / len(dataloader)
epoch_acc = 100. * correct / total
return epoch_loss, epoch_acc
# 5. 测试函数
def evaluate(model, dataloader, criterion, device):
model.eval()
running_loss = 0.0
correct = 0
total = 0
with torch.no_grad():
for inputs, labels in dataloader:
inputs, labels = inputs.to(device), labels.to(device)
outputs = model(inputs)
loss = criterion(outputs, labels)
running_loss += loss.item()
_, predicted = outputs.max(1)
total += labels.size(0)
correct += predicted.eq(labels).sum().item()
epoch_loss = running_loss / len(dataloader)
epoch_acc = 100. * correct / total
return epoch_loss, epoch_acc
# 6. 主训练循环对比
device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
epochs = 5 # 为了演示,只训练5个epoch
loss_functions = {
'Standard CE': nn.CrossEntropyLoss(),
'Label Smoothing CE (ε=0.1)': LabelSmoothingCrossEntropy(smoothing=0.1)
}
history = {name: {'train_loss': [], 'train_acc': [], 'val_loss': [], 'val_acc': []} for name in loss_functions}
for loss_name, criterion in loss_functions.items():
print(f"\n=== Training with {loss_name} ===")
model = SimpleCNN().to(device)
optimizer = optim.Adam(model.parameters(), lr=0.001)
for epoch in range(epochs):
train_loss, train_acc = train_one_epoch(model, trainloader, criterion, optimizer, device)
val_loss, val_acc = evaluate(model, testloader, criterion, device)
history[loss_name]['train_loss'].append(train_loss)
history[loss_name]['train_acc'].append(train_acc)
history[loss_name]['val_loss'].append(val_loss)
history[loss_name]['val_acc'].append(val_acc)
print(f'Epoch {epoch+1:2d}: Train Loss: {train_loss:.4f}, Acc: {train_acc:.2f}% | Val Loss: {val_loss:.4f}, Acc: {val_acc:.2f}%')
# 7. 绘制对比图
plt.figure(figsize=(12, 5))
plt.subplot(1, 2, 1)
for name in loss_functions:
plt.plot(history[name]['train_loss'], label=f'{name} Train')
plt.plot(history[name]['val_loss'], '--', label=f'{name} Val')
plt.xlabel('Epoch')
plt.ylabel('Loss')
plt.title('Training and Validation Loss')
plt.legend()
plt.grid(True)
plt.subplot(1, 2, 2)
for name in loss_functions:
plt.plot(history[name]['train_acc'], label=f'{name} Train')
plt.plot(history[name]['val_acc'], '--', label=f'{name} Val')
plt.xlabel('Epoch')
plt.ylabel('Accuracy (%)')
plt.title('Training and Validation Accuracy')
plt.legend()
plt.grid(True)
plt.tight_layout()
plt.show()
运行这段代码(需要一些时间下载CIFAR-10数据集并训练),你会看到两个损失函数下模型训练损失和验证准确率的曲线。通常,使用Label Smoothing的损失函数在验证集上的准确率曲线会更平滑,有时最终泛化性能也略好,因为它起到了正则化的作用,防止模型对训练标签过于确信。这个例子展示了如何将我们讨论的损失函数理论,集成到一个真实的、端到端的模型训练流程中。在实际项目中,你可以像切换模块一样,轻松地尝试MSE、Huber、Focal Loss等,观察它们对你特定任务和数据的影响,这是优化模型性能的一个非常直接且有效的手段。
更多推荐


所有评论(0)