人人都会AI编程

4.3 损失函数:交叉熵损失在大模型训练中的作用

更新时间:2026-07-09

在 4.1 节中,我们讨论了因果语言建模(CLM)如何通过“预测下一个 Token”让模型掌握语言的生成规律;在 4.2 节中,我们也看到掩码语言建模(MLM)通过“完形填空”让编码器习得双向语义表示。无论采用哪种训练目标,最终都需要一个量化的“打分器”来回答一个核心问题:模型预测的概率分布,与真实答案之间的差距到底有多大?

在大模型领域,这个打分器几乎无一例外地选择了交叉熵损失(Cross-Entropy Loss)。它不仅是理论上的最优选择,更是工程实践中的“默认配置”。理解它的工作原理、数学本质和工程细节,是你读懂训练日志、诊断模型异常和调优训练流程的必备基础。


一、直观理解:衡量两个概率分布的“距离”

交叉熵的核心思想非常朴素:对比“模型相信的”与“真实发生的”

对于语言模型来说,真实答案始终是一个确定性的 Token(比如词表中的第 42 号词)。我们可以把这个真实标签表示为一个 One-Hot 向量——只有正确答案对应的位置是 1,其余全是 0。而模型经过 Softmax 层输出的是一个概率分布向量,它对词表中每一个词都赋予了一个 0 到 1 之间的置信度。

交叉熵的作用,就是度量这两个分布之间的差异。当模型把绝大多数概率“押注”在正确答案上时,交叉熵很小;当模型对正确答案只给了很低的概率,甚至把高概率错押给其他词时,交叉熵会急剧增大。

用一个不严谨的直觉类比:交叉熵就像是模型在每场考试后交的“悔过书长度”——押对题时只写一句“我知道了”,押错题时要写长篇检讨来解释自己为什么错得离谱。


二、数学形式:从单 Token 到全序列

对于单个位置(一个 Token)的预测,交叉熵损失的计算公式为:

$$
\mathcal{L}_{\text{CE}} = -\log P_{\theta}(x_t \mid x_{<t})
$$

其中:

  • $x_t$ 是第 $t$ 个位置的真实 Token;
  • $P_{\theta}(x_t \mid x_{<t})$ 是模型参数 $\theta$ 下,预测该位置为真实 Token 的条件概率(由 Softmax 输出);
  • 负对数概率意味着:概率越接近 1,损失越接近 0;概率越接近 0,损失趋向正无穷。

在 CLM(自回归训练)中的展开:4.1 节提到,CLM 要求模型基于前文 $x_{<t}$ 预测下一个 Token。因此,对于一个长度为 $T$ 的序列,总损失是所有有效预测位置的平均值:

$$
\mathcal{L}_{\text{CLM}} = -\frac{1}{T} \sum_{t=1}^{T} \log P_{\theta}(x_t \mid x_{<t})
$$

注意,这里每个位置 $t$ 的预测都共享同一套模型参数,且依赖因果掩码(Causal Masking,见 3.3 节)保证模型只能“往前看”。

在 MLM(掩码训练)中的适配:4.2 节提到的 BERT 类模型并非预测每个位置,而是只预测被掩码(Mask)的少数位置(通常 15%)。此时损失函数变为:

$$
\mathcal{L}_{\text{MLM}} = -\frac{1}{|\mathcal{M}|} \sum_{t \in \mathcal{M}} \log P_{\theta}(x_t \mid x_{\backslash \mathcal{M}})
$$

其中 $\mathcal{M}$ 是被掩码位置的集合,$|\mathcal{M}|$ 是被掩码的 Token 总数。未被掩码的位置不参与损失计算——这通过 PyTorch/TensorFlow 中的 ignore_index 机制实现,避免了模型学到“复制输入”这种平凡解。


三、为什么是大模型的“天作之合”?

交叉熵之所以成为语言模型训练的默认损失函数,并非偶然,而是因为它在数学、计算和信息论三个层面都与自回归任务完美契合。

1. 与 Softmax 的梯度友好性

Softmax + 交叉熵的组合在反向传播时具有极其简洁的梯度形式。对于 Logits(Softmax 之前的原始分数)$z_i$,当真实标签为 $k$ 时,梯度为:

$$
\frac{\partial \mathcal{L}_{\text{CE}}}{\partial z_i} = P_{\theta}(i) - \mathbb{1}(i=k)
$$

这意味着梯度大小直接等于预测概率与真实标签的偏差。模型越自信且越错误,梯度信号越强;当预测完全正确时,梯度趋于零,参数自然停止剧烈更新。这种“自适应步长”特性让优化器(如 4.4 节将详述的 AdamW)能高效工作。

2. 极大似然估计的等价性

最小化交叉熵,在数学上等价于最大化训练数据的对数似然(Maximum Likelihood Estimation, MLE)。这让大模型的预训练有了清晰的统计解释:我们在寻找一组参数,使得训练语料中观测到的所有文本序列的联合概率最大。

3. 信息论视角:压缩即智能

从信息论角度看,交叉熵衡量的是用模型分布对真实数据进行编码所需的平均比特数。训练过程本质上是在学习一个更好的数据压缩器——模型对语言规律理解越深,生成下一个 Token 的预测越准,压缩率就越高。这也是评估大模型时常用 Perplexity(困惑度) 的原因:

$$
\text{Perplexity} = \exp(\mathcal{L}_{\text{CE}})
$$

困惑度可以被理解为模型在做选择题时,面对的平均“有效选项数”。Perplexity 越低,说明模型越不困惑,对文本的把握越精准。


四、工程实践中的关键细节

在千卡集群上训练百亿参数模型时,交叉熵损失的实现远不止是调用一行 CrossEntropyLoss() 那么简单。以下几个工程细节直接影响训练稳定性和效率。

1. Padding 与 ignore_index

实际训练采用批次(Batch)处理,序列长度不一,需要填充(Padding)到固定长度。这些填充符(如 <pad>)不应参与损失计算。实现上会把 Pad Token 的 Label 设为 -100(PyTorch 默认值),让损失函数自动忽略这些位置,避免模型浪费算力去“预测无意义的填充符”。

2. 数值稳定性:log_softmax 融合

在混合精度训练(见 4.5 节)中,直接先算 Softmax 再取 Log 容易导致数值下溢(Underflow)。现代框架普遍采用 Log-Softmax 与 NLLLoss 融合(即 Fused Cross-Entropy),在 Kernel 层面一次性完成计算,既省显存又防溢出。

3. Label Smoothing:谨慎使用

Label Smoothing 将硬标签(Hard Target)从 [0, 0, 1, 0] 变成 [0.01, 0.01, 0.97, 0.01],目的是防止模型过度自信(Overconfident)。在预训练阶段,大模型通常不使用或仅使用极小的 Smoothing 值(如 0.0–0.1),因为预训练的目标是尽可能拟合数据分布;而在监督微调(SFT,见 5.1 节)或蒸馏场景中,适度的 Label Smoothing 有时能改善泛化。

4. Z-Loss:稳定 MoE 与深层训练

对于 6.5 节将介绍的 MoE(混合专家模型)或极深 Transformer,Logits 的绝对值可能因层数累积变得极大,导致 Softmax 极端尖锐、梯度爆炸。Google 在 PaLM 等模型中引入了 Z-Loss

$$
\mathcal{L}_{\text{Z}} = \frac{1}{B} \sum \log^2 Z, \quad Z = \sum_{i} \exp(z_i)
$$

将其作为辅助损失加到主交叉熵损失上,惩罚过大的 Logits 幅度,显著提升训练稳定性。

5. 损失曲线怎么看

  • 预训练初期:Loss 应从高位(接近词表大小的对数,如 $\ln(32000) \approx 10.37$)快速下降;
  • 中期:下降斜率变缓,可能出现平台期;
  • 末期:Loss 持续缓慢下降,但需配合下游任务指标(Perplexity、Benchmark 分数)判断,而非只看绝对值。

一个值得注意的实用指标是 Token-Level Loss 的分布不均:如果模型在代码段上 Loss 明显低于散文段,说明数据配比或分词器(Tokenizer)可能存在偏置。


五、交叉熵作为“指挥棒”:连接前后向传播

损失函数位于前向传播的最末端,却是反向传播的起点。理解这一点对排查训练故障至关重要:

  1. 梯度信号的来源:4.5 节将讨论的梯度爆炸/消失问题,其根源在于损失反向传播经过 LayerNorm、残差连接和 Attention 时的数值行为。交叉熵提供的初始梯度如果因为 Logits 过大而异常,会在深层网络中被放大。
  2. 优化器的“食物”:4.4 节的 AdamW、Cosine Decay 等策略,本质上都是在消费交叉熵损失回传的梯度。没有清晰的梯度信号,再精妙的优化器也无济于事。
  3. 与对齐技术的衔接:进入第 5 章后,RLHF 中的 Reward Model(奖励模型)和 DPO 中的偏好损失,虽然不再是单纯的交叉熵,但其底层往往仍保留对策略模型输出概率的某种对数概率计算——可以说,交叉熵的思想贯穿了从预训练到对齐的全生命周期。

六、小结

交叉熵损失是大模型训练的“静默基石”:它看似简单,却同时承担了概率建模、梯度生成和信息度量三重角色。

关键 takeaway:

  • 本质:比较模型预测分布与真实 One-Hot 标签的差异,等价于极大似然估计;
  • CLM 与 MLM 的适配:CLM 对所有位置求平均,MLM 仅对掩码位置求平均;
  • 工程必知:注意 Padding 掩码、数值稳定性(Log-Softmax 融合)、MoE 场景下的 Z-Loss;
  • 观测指标:Perplexity = exp(Loss),是衡量模型“困惑程度”的核心指标;
  • 衔接后续:它为 AdamW 提供梯度原料,其数值稳定性直接关联 4.5 节的梯度控制与混合精度训练。

在 4.4 节中,我们将顺着这些梯度信号,深入探讨优化器与学习率调度策略——当交叉熵告诉我们“往哪走”时,AdamW 和 Cosine Decay 决定“走多快、何时减速”。