交叉熵损失深度解析--大模型基础
主包是在边实习边学习大模型,会更一些这方面文章,一边记录方便平时翻看一边加深理解,主包是蒟蒻有不对的地方请大佬指正~
一、交叉熵损失解决什么问题?
交叉熵损失是深度学习中分类任务与生成任务的核心损失函数,也是 ChatGPT、Llama、Qwen 等大模型预训练阶段的基础损失。核心目标:衡量「模型预测分布」与「真实分布」的差异,差异越小,损失值越低,模型性能越好。
适用场景:
- 分类任务(二分类、多分类):如图像识别、文本分类;
- 生成任务(自回归生成):如 LLM 预训练(预测下一个 token)、机器翻译;
- 本质:基于「概率分布匹配」的损失,完美契合语言建模的概率本质(LLM 本质是建模
P(t_i | t_{<i})的条件概率)。
二、数学原理:从信息论到交叉熵公式
1. 前置概念(面试常考铺垫)
要理解交叉熵,必须先掌握「熵」和「KL 散度」,三者是递进关系:
(1)熵(Entropy):衡量真实分布的不确定性
熵是对「真实分布混乱程度」的度量,公式为:
- P(xi):真实分布中第i类的概率;
- 含义:真实分布越确定(如二分类中正负),熵越小;分布越均匀(正),熵越大;
- 例:硬币正反面均匀(P=0.5),熵H(P)=−log20.5−log20.5=1;若硬币作弊(正),熵H(P)≈0.47,不确定性降低。
(2)KL 散度(Kullback-Leibler Divergence):衡量两个分布的差异
KL 散度也叫「相对熵」,直接量化「模型预测分布Q」与「真实分布P」的差距,公式为:
- 性质:KL(P∣∣Q)≥0,当且仅当P=Q时,KL 散度 = 0;
- 问题:KL 散度不对称(KL(P∣∣Q)=KL(Q∣∣P)),但在损失函数中,我们只关心「最小化预测与真实的差异」,不对称性不影响优化目标。
(3)交叉熵(Cross-Entropy):KL 散度的简化形式
将 KL 散度展开:
其中:
这就是「交叉熵损失」的公式。
关键结论:在训练中,真实分布P是固定的(如标签已知),因此H(P)是常数。最小化 KL 散度(让Q逼近P)等价于最小化交叉熵损失CE(P,Q)。→ 这就是为什么交叉熵能作为损失函数:它是「分布差异」的有效代理。
2. 二分类与多分类的具体形式
(1)二分类场景(如文本情感分析:正 / 负)
真实标签y∈{0,1},模型预测正类概率为y^=σ(z)(σ是 sigmoid 函数,将输出映射到 [0,1])。交叉熵损失公式:
- 例:真实标签y=1,模型预测y^=0.9 → L=−[1⋅log0.9+0⋅log0.1]≈0.105;若y^=0.1 → L≈2.303,损失显著增大。
(2)多分类场景(如 LLM 预测下一个 token:词汇表大小V)
真实标签是「one-hot 向量」(如词汇表大小为 5,真实 token 是第 3 类 → y=[0,0,1,0,0]),模型预测分布Q(xi)=softmax(zi)(zi是模型输出的 logits,softmax 将其归一化为概率)。交叉熵损失公式:
- 由于y是 one-hot 向量,只有真实类别k的yk=1,其余为 0,因此公式可简化为:L=−logy^k→ 这就是 LLM 中「自回归语言建模损失」的核心!
(3)批量损失(Batch Loss)
实际训练中,需计算一批样本的平均损失:
- N:批量大小(batch size);
- kj:第j个样本的真实类别。
三、物理意义:为什么交叉熵能指导模型学习?
用通俗的语言解释:
-
惩罚「置信度低的正确预测」和「置信度高的错误预测」:
- 若模型对正确类别预测概率y^k=0.9(置信度高),log0.9≈−0.105,损失小;
- 若y^k=0.1(置信度低),log0.1≈−2.303,损失大;
- 若模型把错误类别预测为 0.9(置信度高),正确类别为 0.1,损失会极大(惩罚错误的自信)。
-
在 LLM 中的特殊意义:LLM 预训练是「自回归生成」,每个 token 的损失是−logP(ti∣t<i),批量平均后就是整个序列的交叉熵损失。→ 损失越小,说明模型基于前文预测下一个 token 的「概率置信度」越高,语言建模能力越强。
四、交叉熵损失 vs 均方误差(MSE):面试高频对比
面试常问:「为什么分类 / 生成任务用交叉熵,而不用 MSE?」核心差异在「梯度特性」和「适用场景」:
| 对比维度 | 交叉熵损失(Cross-Entropy) | 均方误差(MSE) |
|---|---|---|
| 适用场景 | 分类任务、生成任务(概率分布匹配) | 回归任务(连续值预测,如房价预测) |
| 梯度计算(二分类) | 梯度 = y^−y(与误差直接相关) | 梯度 = (y^−y)⋅σ′(z)⋅z′ |
| 梯度稳定性 | 梯度大小与预测误差成正比,训练稳定 | 当y^接近 0 或 1 时,σ′(z)趋近于 0(梯度消失),训练缓慢 |
| 优化目标 | 直接优化「概率分布差异」,契合分类本质 | 优化「预测值与真实值的平方差」,不适合概率建模 |
关键例子:二分类中,若模型预测y^=0.9(真实y=1),MSE 的梯度会因σ′(z)很小而趋近于 0,模型难以更新;而交叉熵的梯度是0.9−1=−0.1,梯度正常,能有效更新参数。
五、交叉熵在 LLM 中的具体应用:工程实现细节(面试加分)
以 Llama、Qwen 等 LLM 为例,交叉熵损失的工程实现有 3 个核心细节,面试中能说清这些说明你有实际经验:
1. 自回归语言建模损失的计算逻辑
LLM 的输入序列是t1,t2,...,tT,模型需要预测t2∣t1、t3∣t1t2、...、tT∣t1..T−1。
- 输入:将序列左移一位(输入t1..tT−1),标签:原始序列t2..tT;
- 损失计算:对每个位置i,计算
,然后求所有位置的平均;
- PyTorch 实现示例:
import torch import torch.nn as nn # 假设模型输出logits: [batch_size, seq_len-1, vocab_size] # 标签labels: [batch_size, seq_len-1](左移后的真实序列) criterion = nn.CrossEntropyLoss(ignore_index=-100) # 忽略padding token logits = model(input_ids) # input_ids是左移后的序列,shape: [B, T-1] loss = criterion(logits.reshape(-1, logits.size(-1)), labels.reshape(-1))
2. 处理 Padding Token:ignore_index 参数
LLM 训练时,不同长度的序列会被 padding 到相同长度(如用<pad> token),这些 token 的损失不应参与计算。
- 解决方案:设置
ignore_index=-100(PyTorch 默认),将 padding token 的标签设为 - 100,交叉熵损失会自动忽略这些位置; - 为什么用 - 100?因为 logits 的索引是 0~vocab_size-1,-100 是无效索引,不会被计入损失。
3. 正则化优化:标签平滑(Label Smoothing)
几乎所有主流 LLM(GPT、Llama、Qwen)的预训练都用了「标签平滑」,面试必问!
(1)什么是标签平滑?
将 one-hot 标签的「硬标签」软化,减少模型对正确标签的过度自信,避免过拟合。
- 原始硬标签:y=[0,0,1,0,0](真实类别是第 3 类);
- 标签平滑后的软标签:
;
- ϵ:平滑系数(主流 LLM 用ϵ=0.1);
- V:词汇表大小(多分类)或 2(二分类)。
(2)标签平滑的交叉熵公式
- 含义:不仅惩罚对正确类别的低置信度,还惩罚对错误类别的过度低置信度,迫使模型学习更泛化的特征。
(3)LLM 中使用标签平滑的原因
- 避免模型「过拟合到训练数据的噪声」:真实文本中可能有笔误或不规范表达,硬标签会让模型过度记住这些噪声;
- 提高模型的「鲁棒性」:面对未见过的文本时,不会因置信度过高而产生极端预测;
- 工程实现:PyTorch 的
CrossEntropyLoss直接支持label_smoothing参数,无需手动实现。
六、交叉熵损失的常见问题与解决方案(面试高频)
1. 梯度消失问题
问题描述:
在多分类任务中,若模型对正确类别预测概率y^k趋近于 1,logy^k趋近于 0,梯度会变小;若预测错误且置信度高,梯度也会趋近于 0。
解决方案:
- 用「Label Smoothing」:软化标签,避免y^k趋近于 0 或 1;
- 调整模型结构:用 ReLU 替代 Sigmoid/Tanh(减少梯度消失),添加 BatchNorm/LayerNorm;
- 优化器选择:用 AdamW 替代 SGD,自适应学习率缓解梯度消失。
2. 类别不平衡问题
问题描述:
若数据集中某些类别样本极少(如二分类中正样本占 1%,负样本占 99%),模型会偏向预测多数类,交叉熵损失会很低但泛化性差。
解决方案:
- 加权交叉熵(Weighted Cross-Entropy):给少数类样本分配更高的权重w,损失公式变为
- Focal Loss:在交叉熵基础上引入「难度因子」,降低易分样本的权重,聚焦难分样本:
- αt:类别权重(平衡样本数量);
- γ:聚焦参数(γ>0,难分样本损失被放大);
- 数据增强:对少数类样本进行扩充(如 LLM 中的回译、同义词替换)。
3. 标签噪声问题
问题描述:
训练数据中存在错误标签(如 LLM 预训练数据中的错别字、标注错误),交叉熵会惩罚模型的正确预测,导致性能下降。
解决方案:
- 标签平滑:软化标签,降低错误标签对模型的影响;
- 噪声鲁棒损失:如 Label Smoothing 的变体、MixUp(混合样本和标签);
- 数据清洗:预处理时过滤明显错误的样本(如 LLM 中的低质量文本过滤)。
七、面试高频问题与答案
1. 请简述交叉熵损失的公式和物理意义?
- 公式(多分类):
(单样本),批量平均为
;
- 物理意义:衡量模型预测分布与真实分布的差异,损失越小,模型对正确类别的置信度越高,分布匹配越好;
- 核心逻辑:基于 KL 散度推导,最小化交叉熵等价于最小化两个分布的差异。
2. 为什么 LLM 预训练用交叉熵损失?
- LLM 的核心任务是「自回归语言建模」,即预测下一个 token 的条件概率P(ti∣t<i);
- 交叉熵损失直接量化「模型预测概率与真实 token 分布的差距」,完美契合该任务;
- 工程上易实现(如 PyTorch 的
CrossEntropyLoss),梯度特性好,训练稳定; - 结合标签平滑后,能有效防止过拟合,提高模型泛化性。
3. 标签平滑的原理和作用是什么?
- 原理:将 one-hot 硬标签软化(如ϵ=0.1,真实类别概率变为1−ϵ,其他类别均分ϵ);
- 作用:
- 避免模型对正确标签过度自信,缓解过拟合;
- 提高模型鲁棒性,降低噪声标签的影响;
- 迫使模型学习更泛化的特征,而非死记硬背训练数据。
4. 交叉熵损失和 KL 散度的关系?
- KL 散度衡量两个分布的差异:KL(P∣∣Q)=CE(P,Q)−H(P);
- 由于真实分布P的熵H(P)是常数,最小化 KL 散度等价于最小化交叉熵损失;
- 交叉熵是 KL 散度的「可优化部分」,更适合作为损失函数(直接计算更简单,无需额外计算熵)。
5. 二分类问题中,交叉熵损失的梯度是什么?为什么比 MSE 好?
- 二分类交叉熵梯度:
,梯度大小与预测误差直接相关,训练稳定;
- MSE 梯度:
,当y^接近 0 或 1 时,σ′(z)趋近于 0,导致梯度消失,模型难以更新;
- 结论:交叉熵更适合分类任务,梯度特性更优,契合概率建模目标。
更多推荐


所有评论(0)