标签平滑为什么能防止模型过拟合?

来源:IPIPP.com作者:清原小日向头衔:网络博主
导读:本期聚焦于清原小日向创作的《标签平滑为什么能防止模型过拟合?》,敬请观看详情。训练集准确率高但验证集开始下降,模型却对每个预测都给出极高置信度,这是过拟合的典型信号。标签平滑通过改变分类任务的目标标签分布来缓解这个问题:它不再要求正确类别的概率为1,而是给正确类别一个略低于1的目标,例如0.9,并把被削减的概率均匀分配给其他类别。这样做可以限制模型输出极端logits,减少对单个训练样本的依赖。本文从交叉熵损失和软标签的角度解释标签平滑的数学原理,分析它如何影响梯度分布和特征聚类,并给出PyTorch与TensorFlow的可用实现。文章还会讨论平滑系数的选择范围、类别不平衡场景下的风险,以及标签平滑与知识蒸馏、Mixup等方法的关系,帮助读者在不同任务中判断是否应该启用这项轻量级正则化技术。

在图像分类或文本分类任务中,模型经过几轮训练后,验证集准确率不再上升,但训练集上的预测置信度却接近100%,这通常是过拟合和过度自信的典型表现。标签平滑正是针对这一问题设计的一种轻量级正则化技术,它不修改网络结构,也不增加额外参数,只改变训练时使用的目标标签分布。

标签平滑为什么能防止模型过拟合?

一、标签平滑的核心机制:从 one-hot 到软标签

分类任务中最常见的标签形式是 one-hot 编码。假设类别数为 K,某个样本的真实类别为第 i 类,那么它的标签向量中只有第 i 个位置是 1,其余位置全部是 0。这种硬标签会推动模型在训练时不断增大正确类对应的 logits,同时压低错误类对应的 logits,最终使得 softmax 输出趋近于 1 和 0 的极端分布。

标签平滑的做法是把硬标签改造成一种软标签。给定平滑系数 epsilon,通常取 0.1,正确类别的目标值从 1 调整为 1 - epsilon + epsilon / K,其他类别则获得 epsilon / K。当类别数 K 等于 10、epsilon 等于 0.1 时,正确类别目标变为 0.91,其余 9 个类别各自获得 0.01。也可以简化理解为正确类别目标下降为 0.9,剩余概率 0.1 被均分给所有错误类别。

这种变化看似很小,但它改变了模型优化的目标。模型不再被要求对某个训练样本输出无限大的正确类 logits,而是只需要让正确类 logits 明显高于错误类即可。由于错误类也获得了少量概率,模型必须学习在多个类别之间保持一定的区分度,这有助于提升泛化能力。

二、标签平滑为什么能抑制过拟合?

从交叉熵损失的角度看,原始交叉熵损失可以分解为正确类的负对数概率。当标签为硬标签时,该损失只关心正确类的预测概率,错误类上的概率变化不会直接产生梯度。标签平滑引入的 epsilon / K 会作用在每一个错误类上,使损失函数中包含对所有类别 logits 的约束。这样一来,模型无法只通过不断放大正确类 logits 来降低损失,因为它同时也要为错误类分配一定的概率,从而抑制 logits 的无限制增长。

从梯度角度看,硬标签训练时,错误类的梯度相对较小或趋近于零,而正确类的梯度会持续推动参数更新。标签平滑相当于给所有错误类增加了一个小的梯度信号,使参数更新更加均衡。这种均衡可以理解为一种正则化,它减少了模型对单个类别或少量训练样本的依赖,与权重衰减、dropout 等方法在目标上类似,但实现位置不同。

从特征表示角度看,标签平滑还能让同类样本在特征空间中更加紧凑,同时让不同类别之间保持稳定的间隔。有实验表明,使用标签平滑训练的模型,其最终层特征向量的类内距离更小,类别中心的分布更加均匀。这种几何属性在迁移学习或模型蒸馏时也有帮助。

三、在 PyTorch 和 TensorFlow 中实现标签平滑

PyTorch 标准库没有直接提供标签平滑版本的交叉熵损失,但可以通过组合 log_softmaxnll_loss 的方式手动实现。下面给出一个常用实现:

import torch
import torch.nn as nn
import torch.nn.functional as F

class LabelSmoothingCrossEntropy(nn.Module):
    def __init__(self, epsilon=0.1, reduction='mean'):
        super().__init__()
        self.epsilon = epsilon
        self.reduction = reduction

    def forward(self, logits, target):
        log_probs = F.log_softmax(logits, dim=-1)
        nll_loss = -log_probs.gather(dim=-1, index=target.unsqueeze(1)).squeeze(1)
        smooth_loss = -log_probs.mean(dim=-1)
        loss = (1 - self.epsilon) * nll_loss + self.epsilon * smooth_loss
        if self.reduction == 'mean':
            return loss.mean()
        elif self.reduction == 'sum':
            return loss.sum()
        return loss

这段代码中,target 传入的是类别索引而不是 one-hot 向量。它先计算每个类别对应的对数概率,再取出正确类别的负对数似然,最后与所有类别平均对数概率的负值按 epsilon 加权求和。这样得到的损失在数学上与标签平滑交叉熵等价。

TensorFlow 的实现更加直接,tf.keras.losses.CategoricalCrossentropy 原生支持 label_smoothing 参数:

import tensorflow as tf

loss_fn = tf.keras.losses.CategoricalCrossentropy(
    from_logits=True,
    label_smoothing=0.1
)

# logits shape: (batch_size, num_classes)
# labels shape: (batch_size, num_classes), one-hot
loss = loss_fn(labels, logits)

使用时需要注意,from_logits=True 表示模型最后一层输出的是未经过 softmax 的 logits,标签仍使用 one-hot 形式。框架内部会自动将标签转换为软标签,因此不需要手动修改数据集。PyTorch 用户如果不想自己编写损失函数,也可以先把 one-hot 标签转换为软标签,再使用标准交叉熵损失,但这种方式会增加额外的内存开销。

四、经验调参:平滑系数怎么选?

平滑系数 epsilon 并不是越大越好。常用的范围是 0.05 到 0.2。对于类别数较多、噪声较大或训练样本较少的任务,可以适当提高到 0.1 到 0.15;对于类别数很少、数据质量较高且模型容量不足的任务,过大的 epsilon 会让模型过于保守,导致正确类别的预测概率偏低,影响准确率。

类别不平衡场景需要特别小心。标签平滑默认把被削减的概率均分给所有错误类别,这会进一步提升长尾类别在损失中的权重。如果数据集中某些类别的样本极少,均匀分配可能让模型把更多注意力放在那些稀有类别上,表面上可能改善某些指标,但也可能引入新的噪声。此时可以考虑使用类别频率加权的标签平滑,或者在数据采样和损失加权上同时进行调整。

另一个常见误区是认为标签平滑只能用于图像分类。实际上,它在文本分类、语音识别、推荐系统的多分类任务中同样有效。对于二分类问题,标签平滑也可以使用,只是概率分配方式更简单,正确类别目标从 1 降为 1 - epsilon / 2,错误类别从 0 提升为 epsilon / 2

五、标签平滑与知识蒸馏、Mixup 的异同

标签平滑生成的软标签在形式上与知识蒸馏中的教师软标签很相似,但两者来源不同。知识蒸馏使用一个已经训练好的教师模型来生成软标签,这些软标签包含了类别之间的相似性信息;标签平滑则完全根据均匀分布构造软标签,没有利用任何额外模型,因此计算成本更低,但也不包含数据本身的类别关系。

Mixup 和 CutMix 等数据增强方法在输入空间进行插值,并同步构造软标签,而标签平滑只修改标签空间。它们可以叠加使用,但叠加后需要重新调整 epsilon 和 Mixup 的混合系数,否则可能过度软化目标,削弱模型对核心特征的辨别能力。实践中,如果已经使用了知识蒸馏或 Mixup,可以先把标签平滑系数设置得稍小一些,再通过验证集微调。

总体来看,标签平滑是一种实现成本低、副作用小的正则化手段。它不会改变推理过程,推理时仍使用普通的 softmax 输出。只有在模型明显过拟合或对预测置信度有较高要求时,关闭标签平滑可能更合适。对于大多数中小规模分类任务,开启一个较小的标签平滑通常能带来更稳定的收敛曲线和更好的泛化效果。

标签平滑过拟合交叉熵损失修改时间:2026-08-22 20:48:01

免责声明:​ 已尽一切努力确保本网站所含信息的准确性。网站内容多为原创整理与精心编撰,观点力求客观中立。本站旨在免费分享,内容仅供个人学习、研究或参考使用。若引用了第三方作品,版权归原作者所有。如内容涉及您的权益,请联系我们处理。
内容垂直聚焦
专注技术核心技术栏目,确保每篇文章深度聚焦于实用技能。从代码技巧到架构设计,为用户提供无干扰的纯技术知识沉淀,精准满足专业提升需求。
知识结构清晰
覆盖从开发到部署的全链路。AI、前端、编程、数据库、服务器、建站、系统层层递进,构建清晰学习路径,帮助用户系统化掌握开发与运维所需的核心技术。
深度技术解析
拒绝泛泛而谈,深入技术细节与实践难点。无论是数据库优化还是服务器配置,均结合真实场景与代码示例进行剖析,致力于提供可直接应用于工作的解决方案。
专业领域覆盖
精准对应开发生命周期。从前端界面到后端编程,从数据库操作到服务器运维,形成完整闭环,一站式满足全栈工程师和运维人员的技术需求。
即学即用高效
内容强调实操性,步骤清晰、代码完整。用户可根据教程直接复现和应用于自身项目,显著缩短从学习到实践的距离,快速解决开发中的具体问题。
持续更新保障
专注既定技术方向进行长期、稳定的内容输出。确保各栏目技术文章持续更新迭代,紧跟主流技术发展趋势,为用户提供经久不衰的学习价值。