损失函数深度实战:从交叉熵、标签平滑到 Focal Loss 与多任务损失的训练工程全解

损失函数是模型训练信号的源头,也是"经验风险最小化"这一机器学习第一性原理的唯一落点。它被反复引用,却极少有文章从数值稳定性、类别不平衡、软标签冲突、多任务尺度失衡等生产视角做工程化展开。本文与本站《RMSNorm》《相对位置编码》《交叉注意力》《残差连接》《学习率调度》《梯度累积》《词嵌入》《激活函数》共同构成 Transformer 内部机制与训练工程深度解系列,补全"训练目标"这一环。

一、第一性原理:为什么需要"替代损失"

监督学习的目标是最小化期望风险 $R(f)=\mathbb{E}_{(x,y)\sim\mathcal{D}}[\ell(f(x),y)]$。真实分布不可知,只能用经验风险 $\frac{1}{N}\sum_i \ell(f(x_i),y_i)$ 近似。这里的 $\ell$ 就是损失函数。

理想的 0/1 损失 $\ell_{0/1}=\mathbb{I}[f(x)\neq y]$ 直观但不可导、不连续,梯度处处为 0 或不存在,无法用梯度下降优化。因此所有可训练损失都是替代损失(surrogate loss):用平滑、可微的函数去逼近 0/1 损失的"优化意图"。理解这一点,就能解释为什么 SVM 用 hinge、回归用平方、分类用交叉熵——它们都不是"目标本身",而是"可优化的代理"。

损失 形式 可微性 适用 缺陷
0/1 loss $\mathbb{I}[\hat y \neq y]$ 不可微 理论最优 无法梯度优化
Hinge $\max(0,1-y\hat f)$ 次梯度 SVM 对离群点敏感
平方误差 $(y-\hat y)^2$ 可微 回归 分类时边界处梯度饱和
交叉熵 $-\sum y\log\hat y$ 可微 分类 需数值稳定实现

二、分类损失谱系:从交叉熵到 Focal Loss

2.1 交叉熵的推导与数值稳定性

交叉熵源于最大似然估计。给定真实分布 $q$(one-hot)与预测分布 $p_\theta$,最小化 KL 散度 $D_{KL}(q\|p)=H(q)+\underbrace{-\sum q\log p}_{\text{交叉熵}}$。因为 $H(q)=0$(one-hot 熵为 0),最小化 KL 等价于最小化交叉熵。

实践中关键陷阱是数值稳定性:先算 softmax 再取 log 会触发 exp 上溢(当 logit 很大时)。正确做法是合并为 log_softmax(先做 x - max(x) 再 exp),把减最大值这一步吸收进去:


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

# ❌ 数值危险:先 softmax 再 log,大 logit 时 exp 溢出
probs = F.softmax(logits, dim=-1)
loss_bad = (-targets_onehot * torch.log(probs)).sum(-1).mean()

# ✅ 数值安全:log_softmax 内部做了 max 减治,等价于上式但稳定
logp = F.log_softmax(logits, dim=-1)
loss_good = (-targets_onehot * logp).sum(-1).mean()

# PyTorch 的 CrossEntropyLoss 内部正是 log_softmax + NLLLoss,切勿手动拼 softmax
loss = F.cross_entropy(logits, targets)  # targets 为类别索引(long),非 one-hot

CrossEntropyLoss 与 LogSoftmax + NLLLoss 数学等价,但前者在 C++ 内核里合并了 softmax 与 log,既省一次 exp 又避免中间溢出——这是生产环境必须直接调用 cross_entropy 的根本原因,而非"为了教学清晰"手写 softmax。

2.2 标签平滑(Label Smoothing)

Szegedy 等(2016)发现 one-hot 标签迫使模型对正确类输出概率 1.0、其余 0.0,造成过自信(over-confident)与过拟合。标签平滑用 $\varepsilon$ 把目标分布"软化":

$$q'(y|x)=\begin{cases}1-\varepsilon & y=y_{\text{true}} \\ \varepsilon/(K-1) & \text{otherwise}\end{cases}$$


class LabelSmoothingCrossEntropy(nn.Module):
    def __init__(self, eps=0.1, reduction='mean'):
        super().__init__()
        self.eps = eps
        self.reduction = reduction
    def forward(self, logits, targets):
        n_classes = logits.size(-1)
        logp = F.log_softmax(logits, dim=-1)
        # 平滑后的软标签:正确类 1-eps,其余 eps/(K-1)
        smooth = torch.full_like(logp, self.eps / (n_classes - 1))
        smooth.scatter_(1, targets.unsqueeze(1), 1.0 - self.eps)
        loss = (-smooth * logp).sum(dim=-1)
        return loss.mean() if self.reduction == 'mean' else loss.sum()

$\varepsilon=0.1$ 是图像分类的常用值;在 LLM 微调中也可用较小值(如 0.0~0.05)缓解过拟合。注意:标签平滑与知识蒸馏的软标签目标存在语义冲突——蒸馏要求模型去拟合教师的软分布,而标签平滑又去破坏硬标签的尖锐性,二者同用时需统一目标分布。

2.3 Focal Loss:聚焦难样本

Lin 等(2017,RetinaNet)针对类别极度不平衡 + 大量易分负样本的检测场景,提出对交叉熵降权的机制:

$$\text{FL}(p_t)=-\alpha_t(1-p_t)^\gamma\log(p_t),\quad p_t=\begin{cases}p & y=1\\1-p & y=0\end{cases}$$

$(1-p_t)^\gamma$ 让易分样本($p_t\to1$)的权重趋近 0,难分样本($p_t$ 小)保留大权重;$\gamma$ 控制聚焦强度(典型 2.0),$\alpha$ 做类别平衡。


def focal_loss(logits, targets, alpha=0.25, gamma=2.0, reduction='mean'):
    ce = F.cross_entropy(logits, targets, reduction='none')  # 逐样本 CE
    pt = torch.exp(-ce)                                      # pt = 模型对正确类的概率
    focal = alpha * (1 - pt) ** gamma * ce                   # 难样本(pt小)权重放大
    return focal.mean() if reduction == 'mean' else focal.sum()

Focal Loss 把优化注意力从"已被学会的易样本"转移到"边界难样本",在长尾分类、目标检测中收益显著,但在均衡数据集上常与常规 CE 持平甚至略差——不要无脑套用。

2.4 多标签与类别不平衡:BCE + pos_weight

文本多标签、医学多病种等"一个样本可属于多类"的场景,需用 BCEWithLogitsLoss(内部 sigmoid + BCE,数值稳定)而非 softmax 系损失。对正样本稀疏的类别,用 pos_weight 放大正类梯度:


# 假设正样本占比 1%,用 pos_weight 抵消不平衡
pos_weight = torch.tensor([99.0])  # = 负样本数 / 正样本数
loss = F.binary_cross_entropy_with_logits(logits, multi_hot_labels, pos_weight=pos_weight)

三、回归损失谱系:重尾噪声下的鲁棒性

损失 公式 对离群点 典型用途
MSE / L2 $(y-\hat y)^2$ 敏感(平方放大) 高斯噪声回归
MAE / L1 $\vert y-\hat y\vert$ 鲁棒 重尾噪声、中位数回归
Smooth L1 分段 较鲁棒 目标检测框回归
Huber 阈值 $\delta$ 分段 鲁棒 需兼顾连续梯度
Log-Cosh $\log(\cosh(y-\hat y))$ 鲁棒 平滑近似 MAE
Quantile $\rho_\tau$ 分位估计 区间预测

def smooth_l1_loss(pred, target, beta=1.0):
    diff = (pred - target).abs()
    return torch.where(diff < beta, 0.5 * diff ** 2 / beta, diff - 0.5 * beta).mean()

def huber_loss(pred, target, delta=1.0):
    diff = (pred - target).abs()
    return torch.where(diff <= delta, 0.5 * diff ** 2, delta * (diff - 0.5 * delta)).mean()

经验法则:噪声接近高斯用 MSE;存在离群点或重尾分布用 L1/Huber;L2 在异常值上梯度爆炸,可能拖垮整个 batch 的收敛。

四、蒸馏损失:带温度的 KL 散度

知识蒸馏(Hinton 2015)让学生对齐教师的软分布。核心是带温度 $T$ 的 KL 散度——温度越高,软标签越平滑、暗含的类别相似性信息越丰富:


def distillation_loss(student_logits, teacher_logits, labels, T=4.0, alpha=0.7):
    # 软目标:温度软化后的 KL(注意 kl_div 要求 log 概率在前)
    kd = F.kl_div(
        F.log_softmax(student_logits / T, dim=-1),
        F.softmax(teacher_logits / T, dim=-1),
        reduction='batchmean') * (T * T)        # 乘 T^2 抵消梯度幅度缩放
    # 硬目标:常规 CE
    ce = F.cross_entropy(student_logits, labels)
    return alpha * kd + (1 - alpha) * ce

(T*T) 缩放是必要的:软化的 logit 梯度幅度被除以 $T$,乘以 $T^2$ 还原到原尺度,否则蒸馏项会被"稀释"。本站《模型提取与盗取攻击》从安全视角讨论了相同数学,本文聚焦其作为训练损失的工程形态。

五、多任务损失:尺度失衡与自适应加权

多任务模型常把多个损失简单相加 $\mathcal{L}=\sum_i w_i\mathcal{L}_i$,但不同任务的梯度尺度、收敛速度差异巨大,固定权重会让主导任务"吞掉"其他任务。Kendall 等(2018)提出用同方差不确定性自动学习权重:

$$\mathcal{L}_{\text{total}}=\sum_i\frac{1}{2\sigma_i^2}\mathcal{L}_i+\log\sigma_i$$


log_sigma1 = nn.Parameter(torch.zeros(1))   # 可学习任务噪声 log σ
log_sigma2 = nn.Parameter(torch.zeros(1))
# 训练时与模型参数一同优化;σ 越大(越不确定)的任务权重自动变小
loss = (torch.exp(-2*log_sigma1) * L1 + log_sigma1
        + torch.exp(-2*log_sigma2) * L2 + log_sigma2)

这把"该给哪个任务多少权重"从超参搜索变成了可微分的学习问题,是生产级多任务训练的标准解法。

六、生产陷阱清单(必读)

# 陷阱 后果 修复
1 手写 log(softmax(x)) 大 logit 时 exp 溢出得 NaN 用 F.cross_entropy / log_softmax
2 reduction='mean' 与梯度累积步数不匹配 累积后梯度尺度偏小,等效 LR 减半 累积用 sum 并手动 /steps,或确认框架已除
3 ignore_index 未设置 padding/special token 计入损失污染梯度 对 padding 设 ignore_index=-100
4 软标签传 long 类型 类型错误崩溃 软标签用 float,硬标签用 long
5 混合精度下 CE 输入为 fp16 logit 溢出、loss 变 Inf 保持 logits 为 fp32 再算损失
6 类别不平衡仍用原生 CE 多数类主导,少数类学不到 class_weight / pos_weight / Focal
7 标签平滑与蒸馏软标签同用且目标不一致 训练目标自相矛盾 统一软目标分布来源
8 多任务损失直接求和 尺度大的任务压制其他 用 uncertainty 加权或显式归一
9 初期 loss 出现 Inf/NaN 不检测 权重被 NaN 污染不可逆 加 torch.isfinite(loss) 断言与梯度裁剪
10 仅看 loss 下降 loss 降但验证指标不升(标签噪声/过拟合) 同时监控业务指标与校准误差

七、可复现参考实现


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

class TrainingLosses:
    """损失函数工程工具箱:覆盖交叉熵、标签平滑、Focal、蒸馏、多任务。"""
    @staticmethod
    def ce(logits, targets):
        return F.cross_entropy(logits, targets)  # 始终首选内置,数值稳定

    @staticmethod
    def label_smoothing(logits, targets, eps=0.1):
        n = logits.size(-1)
        logp = F.log_softmax(logits, dim=-1)
        smooth = torch.full_like(logp, eps / (n - 1))
        smooth.scatter_(1, targets.unsqueeze(1), 1.0 - eps)
        return (-smooth * logp).sum(-1).mean()

    @staticmethod
    def focal(logits, targets, alpha=0.25, gamma=2.0):
        ce = F.cross_entropy(logits, targets, reduction='none')
        pt = torch.exp(-ce)
        return (alpha * (1 - pt) ** gamma * ce).mean()

    @staticmethod
    def distill(student, teacher, labels, T=4.0, alpha=0.7):
        kd = F.kl_div(F.log_softmax(student / T, -1),
                      F.softmax(teacher / T, -1), reduction='batchmean') * T * T
        return alpha * kd + (1 - alpha) * F.cross_entropy(student, labels)

# 多任务:用可学习噪声自动平衡权重
class MultiTaskHead(nn.Module):
    def __init__(self):
        super().__init__()
        self.log_s1 = nn.Parameter(torch.zeros(1))
        self.log_s2 = nn.Parameter(torch.zeros(1))
    def loss(self, l1, l2):
        return (torch.exp(-2*self.log_s1)*l1 + self.log_s1
                + torch.exp(-2*self.log_s2)*l2 + self.log_s2)

八、与系列其他文章的桥接

  • 梯度累积:reduction='mean' 的 loss 已对 batch 内样本数归一,多步累积时必须显式乘/除累积步数或改用 sum,否则等效学习率减半——详见《梯度累积深度实战》。
  • 学习率调度:loss 量级变化会改变梯度健康度,warmup 阶段若 loss 异常抖动,应先查陷阱清单第 9 条而非盲目调小 LR。
  • 激活函数 / 归一化:CE 对 logit 的尺度敏感,Pre-LN、RMSNorm 等稳定了前向尺度,间接让 loss 曲面更平滑——与《RMSNorm》《激活函数》互为因果。
  • 知识蒸馏 / 模型安全:蒸馏损失与《模型提取与盗取攻击》共享 KL 数学,理解本文有助于看穿"用查询 API 反演模型"的攻击面。

损失函数看似是训练脚本里一行 criterion(),实则承载了整个学习目标的设计哲学:选错代理、忽略数值稳定、无视类别失衡,都会让再深的网络也收敛到错误的解。把它当成一个需要工程化对待的组件,而非默认黑盒,是迈向生产级训练的第一课。

点赞(0) 打赏

评论列表 共有 0 条评论

暂无评论
立即
投稿
网站二维码

微信公众账号

微信扫一扫加关注

发表
评论
返回
顶部
/* 跳过导航链接 (无障碍) */ position: absolute; top: -100px; left: 15px; z-index: 99999; padding: 8px 16px; background: #007bff; color: #fff; font-size: 14px; border-radius: 0 0 4px 4px; text-decoration: none; transition: top 0.2s; } top: 0; outline: 3px solid #0056b3; }