首页 / 资讯中心 / 文章详情

负对数似然与交叉熵:原理、等价关系与数值稳定实现

负对数似然与交叉熵:原理、等价关系与数值稳定实现 ★ FEATURED ARTICLE
有次线上模型迭代我在自定义模型头时图省事手动把 logits 过了一遍 softmax 再取 log结果验证集上 loss 全部变成 nan。排查了一个下午最后发现是 float 精度的问题——softmax 之后概率已经接近 0 的位置再取 log 直接给到 -inf负号和另一个 inf 一碰就成 nan。那时候我才意识到负对数似然这个每天在框架里自动执行的操作我从没有认真把每一步的原理和工程细节想清楚。这篇文章想把这个话题彻底展开。负对数似然函数Negative Log-LikelihoodNLL看着只是取概率的负对数这么一个简单动作实际上它贯穿了统计推断、分类损失设计、信息论、数值稳定实现、训练状态诊断这些完全不同的层次。不管你是刚接触机器学习、搞不清损失函数来源的初学者还是已经在调模型但想补理论短板的工程师这篇文章都值得读一读。我会从最大似然估计出发把它为什么长成这个样子讲明白再落到分类任务里的具体形态、与交叉熵的关系、工程实现的坑以及训练过程中通过 NLL 曲线能观察到的细节。1. 为什么优化目标必须变成负的对数似然1.1 最大似然估计的直觉从抛硬币说起假设有一枚不均匀的硬币正面朝上的概率是 θ连续抛 10 次得到 6 次正面、4 次反面。用数学公式表达这次实验结果发生的概率就是 L(θ) θ^6 · (1-θ)^4这个 L(θ) 就是似然函数。你在直观上会觉得这次实验最支持 θ 0.6因为当 θ 0.6 时 L(θ) 最大。把求使 L(θ) 最大的 θ这个过程叫最大似然估计基本就是把直觉翻译成最优化问题。放到机器学习场景里训练集是 (x_i, y_i) 的配对模型带着参数 θ要估计的是条件概率 P(y_i | x_i; θ)。此时整个训练集的似然函数是每个样本似然的连乘L(θ) ∏_i P(y_i | x_i; θ)训练的目标就是找到一个 θ*让这整个连乘最大。说得直白点我们希望模型赋予真实数据的概率最高也就是模型看到这些数据时不觉得惊讶。1.2 对数变换不是偏好是被迫的选择直接最大化这个连乘在数学上完全正确但在实际计算中会遇到两个致命问题。第一个是数值下溢。假设每个样本的概率在 0.9 到 0.99 之间一万个样本连乘下来结果在 float32 精度下早就变成了 0。log 把一个微小到无法表示的数映射成 -23000 这样的有限值至少还能在计算图里正常传播。第二个是求导和幂运算的麻烦。连乘求导要做一长串乘积法则而连加求导只需要简单地求和。对数函数有一个核心性质它是单调递增的所以最大化 log L(θ) 和最大化 L(θ) 得到的最优点完全一致。取对数后连乘变成连加log L(θ) ∑_i log P(y_i | x_i; θ)很多数学技巧其实是被迫的——不是因为它优雅而是因为它让原本不现实的计算变得现实。取对数这个操作就是典型例子。1.3 负号从哪里来优化器的统一约定为什么前面还要加个负号变成负对数似然主要原因在于深度学习框架里的优化器默认只做最小化。梯度下降更新参数时执行的是参数 - 梯度 × 学习率这个公式假设目标函数越小越好。最大化的对数似然不能直接喂给这种优化器所以取个相反数从最大化对数似然变成最小化负对数似然形式上是Loss(θ) -∑_i log P(y_i | x_i; θ)深层逻辑上把一切目标都统一成最小化后续加正则项、加约束、多任务加权都方便了。L2 正则项是越小越好、模型的复杂度惩罚是越小越好你要参与统一优化负对数似然自然也得变成越小越好。有意思的是这背后还藏着信息论的影子-log p 可以被理解为观察到概率为 p 的事件所带来的惊讶程度。概率越小的样本这条公式给出的惩罚越大。最大似然从让数据尽可能可能变成让数据尽可能不令人惊讶本质上是一致的视角。2. 分类任务里NLL的真实形态从伯努利到Softmax2.1 二分类logistic loss 就是 NLL二分类问题的真实标签 y 只有 0 和 1 两种情况。模型输出通常是一个经过 sigmoid 的值 p P(y1 | x)同时 p 也隐含了 P(y0 | x) 1-p。把这两个情况合并写成一个公式P(y | x) p^y · (1-p)^(1-y)当 y1 时只剩 p^1(1-p)^0 p当 y0 时只剩 p^0(1-p)^1 1-p。取负对数后得到NLL -[y log p (1-y) log(1-p)]这就是大家熟悉的 logistic loss也是二元交叉熵。注意它的行为如果真实标签是 1模型给 p0.99那么损失大约是 0.01模型给 p0.4损失大约是 0.92。这个损失不是简单判对错而是对错误程度做连续惩罚。这也解释了为什么 NLL 训练出来的模型能给出有概率意义的置信分数而不是只输出一个硬分类结果。2.2 多分类Softmax 回归中的下标记号K 分类问题里模型输出的是一个 K 维 logits 向量 z经过 softmax 后变成概率分布P(k | x) exp(z_k) / ∑_j exp(z_j)真实标签 y 是一个整数。把这个概率代入 NLL 公式单样本损失就是NLL -log P(y | x) -log(softmax_y(z))这个式子看起来简单到没什么可说的但请记住 softmax 的概率是归一化的所有类别概率加起来等于 1。所以模型要让正确类别的概率提高就必然会压低其他类别的概率这是一种天然的竞争机制。这也是为什么多分类 NLL 在训练中经常表现为自信——它鼓励模型把概率质量往真实类别集中。有时你会看到有人把 NLL 写成多行带求和号的式子-∑_k y_k log p_k其中 y_k 是 one-hot 编码。这个写法在数学和解代码上更通用但因为 one-hot 中除了真实类别以外全是 0最终结果和我上面写的一行式完全一致。2.3 具体数值算一遍感受距离如何被放大假设一个三分类问题模型对某个样本输出的 logits 是 z [2.2, 1.0, 0.1]。经过 softmax 后概率大约是 [0.67, 0.23, 0.10]。如果真实标签是第 0 类那么 NLL -log(0.67) ≈ 0.40模型对这个样本已经比较有信心了。但如果真实标签是第 2 类那么 NLL -log(0.10) ≈ 2.30模型犯了大错。同样是分类错误这个损失从 0.40 跳到 2.30相差接近六倍。原因是 log 在靠近 0 的地方下降速度极快也就意味着模型每把概率压到接近 0 的位置一次一旦该位置是真实标签就会被扣很大的分数。这个特性是 NLL 在分类任务中有效性的来源——它让错误预测的惩罚呈超线性增长避免模型装傻来逃避惩罚。2.4 一个实用锚点均匀分布时的 NLL log K如果你完全不知道任何信息在 K 个类别上给出均匀概率 1/K那么单样本 NLL 就是 -log(1/K) log K。这个值可以作为训练早期的重要参照。我记得有个项目是 125 类分类模型随机初始化后跑第一个 batchNLL 应该在 4.8 附近log 125 约等于 4.83。如果你训练的初始 loss 远低于 log K反而是个危险信号——可能数据泄漏了也可能模型初始化出了问题。这个锚点还有一个作用很多人习惯用 loss 下降了多少来衡量训练效果但 loss 的绝对值在不同类别数量下不可比。10 分类任务从 2.30 降到 0.50和 1000 分类任务从 6.91 降到 1.50体验是完全不同的。3. NLL和交叉熵是同一件事吗两种视角的对齐3.1 信息论定义下的交叉熵交叉熵来自信息论描述的是用模型分布 q 去编码真实分布 p 中样本时的平均编码长度H(p, q) -∑_k p(k) · log q(k)如果真实分布 p 是一个 one-hot 标签即 p(y)1、其他都是 0那么上式里只有一个非零项H(p, q) -log q(y)在分类任务中q(y) 就是模型给真实类别预测的概率。所以 NLL 和交叉熵在这个场景下恰好是同一个数。NLL 是从统计学角度喊的名字交叉熵是从信息论角度喊的名字两者在平地上碰到了一起。3.2 为什么说完全等价需要小心严格来讲NLL 定义是经验分布下的负对数似然。当标签是 hard one-hot 标签时NLL 就等于交叉熵。但在标签不是 one-hot 的情况下两者有微妙差别。比如知识蒸馏里用 teacher 模型的 soft label t_k 作为监督损失函数通常写成Loss -∑_k t_k log p_k这个式子没法写成 -log p_y 的形式因为每个类别都贡献了一项。你用加权交叉熵来描述它比用NLL更精确。不过很多框架里依然把它归为soft target 的 NLL或者带温度交叉熵本质上都是同族的公式。这也是为什么我建议你学 NLL 时一定要理解这个等价边界。如果你只知道NLL 就是取正确类别的负 log 概率碰到 soft label 场景就会发懵。3.3 和 KL 散度的关系训练就是在做分布逼近KL 散度描述两个分布之间的距离KL(p || q) H(p, q) - H(p) -∑_k p(k) log q(k) ∑_k p(k) log p(k)训练过程中 p 是数据集上的经验分布它固定不变所以 H(p) 是一个常数。最小化 KL(p || q) 就等于最小化交叉熵 H(p,q)也就等于最小化 NLL。换句话说分类模型的训练过程是在寻找一个模型分布 q让它尽可能接近数据分布 p。这个视角对理解生成模型、对比学习里的 InfoNCE 损失很有帮助。很多进阶损失函数设计到最后都能看到最小化负对数似然 最小化 KL 散度这个统一的骨架。你掌握了这一层推导就能看懂各种新型损失函数从哪来。3.4 回归任务为什么用的是 MSE高斯 NLL 的展开前面说的都是分类实际上回归里最常见的目标函数也是 NLL 的一种特殊形态。假设 y|x 服从高斯分布即均值为模型输出 μ_θ(x)方差 σ² 固定。那么单样本的负对数似然是NLL 0.5·log(2πσ²) (y - μ_θ(x))² / (2σ²)第一项和 θ 无关是常数第二项去掉系数 0.5/σ² 之后就变成了 (y - μ)² 也就是均方误差 MSE。也可以换一种理解MSE 假设误差服从高斯分布而分类的 NLL 假设标签服从伯努利或类别分布。这两种损失的目标一致只是对数据生成过程的假设不同。这意味着不需要为分类用交叉熵、回归用 MSE找两套理由它们背后是同一个统计框架。以后你在设计新任务时只需要思考你的目标变量服从什么分布就能直接写出合适的 NLL 损失。4. 工程实现里的数值陷阱log(0)、下溢与log-sum-exp4.1 前端最常踩的坑先 softmax 再 log回到开头那个 nan 的故事。很多人在自实现时按照数学公式的顺序先算 softmax 得到概率 p再计算 -log(p)。数学上没错但在浮点数的世界里softmax 出来的概率很多是极小值。一个 1000 类任务某个类别的概率可能小于 1e-30float32 还能勉强表示但如果你用了 float16早就下溢成 0 了。log(0) 是什么是 -inf。在 loss 后面再做什么运算-inf 一旦出现整个梯度都有可能变成 nan。正确的做法是把 softmax 和 log 合并成 log_softmax 一步完成。数学上 log(softmax(z)_k) z_k - log(∑_j exp(z_j))这样在 log 内部先处理掉极小的概率项不会直接出现 log(0)。名称上这只是一次运算合并工程上却能避开一个让新人数小时的 bug。4.2 log-sum-exp 和 max-shift 的原理光是合并还不够log_softmax 里面的 ∑ exp(z_j) 同样可能上溢。如果 logits 里某个值特别大比如 exp(1000)在 float 里直接变 inf。业界标准的解法是 log-sum-exp trick先找到 logits 的最大值 m然后所有 logits 先减去 m 再算 exp。log_softmax(z)_k z_k - m - log(∑_j exp(z_j - m))这样做的原因是 exp(z_j - m) 的最大值必然不超过 1因为 z_j - m ≤ 0所以 exp 的结果都在 [0, 1] 范围内不可能上溢。数学上减掉 m 再恢复值是不变的因为(m 会出现在 log 内和 z_k 中互相抵消)。所有主流框架的 Categorical 分布、softmax、交叉熵内部都做了这一步。自实现时如果不做哪怕你在 CPU 上测没问题换 GPU 半精度训练就会崩。4.3 PyTorch 里的推荐姿势与维度细节实际写代码时推荐做法是直接把 logits 传给 F.cross_entropy。它内部做了 log_softmax 和 nll_loss 的合并数值稳定性最好。如果你需要对 log_probs 做日志或者自定义逻辑可以分开写。import torch import torch.nn.functional as F logits torch.randn(16, 10) # (batch_size, num_classes) labels torch.randint(0, 10, (16,)) # (batch_size, )必须是 long 型 # 推荐一步到位 loss1 F.cross_entropy(logits, labels) # 拆开来用先自己算 log_softmax再做 NLL log_probs F.log_softmax(logits, dim-1) loss2 F.nll_loss(log_probs, labels) print(loss1.item() loss2.item()) # 两种写法数值上完全一致自实现时要注意两个细节。第一个是 labels 必须传类别索引而不是 one-hot如果你非要用 one-hot得自己把 target 转成 (batch, class_num)然后按元素相乘再求和第二个是注意 log_probs 的维度顺序默认假设第一维是 batch类别在最后一维如果你的张量布局不同务必指定正确的 dim。4.4 自己算多标签的概率时也要用 log 一族函数还有一个隐蔽的坑多标签分类或多输出结构的概率计算。有些人实现时会先算各个任务的概率 p_j再用 P ∏ p_j^{y_j}(1-p_j)^{1-y_j} 去算联合概率。某种程度这是合理的建模但如果你在损失函数里直接对这个 P 取 log然后让梯度反向传播会遇到两个问题一是 p_j 一旦为 1模型极端自信1-p_j 为 0log(0) 就出现二是因为经过了非线性乘法梯度数值也会非常不稳定。更稳妥的做法是在 log 空间里累加每个任务的 log p_j等价但数值性质好很多。我在做多任务模型时早期就是这么吃亏的后来统一改成 log-sum-exp 和 log-prob 的方式才彻底摆脱了概率下溢导致的偶发短路。5. 训练中NLL曲线教会我的事梯度、过拟合与label smoothing5.1 最优雅的一阶导数结论∂L/∂z_k p_k - y_k如果拿掉中间层的复杂结构单看 softmax NLL 这个组合对 logits z_k 求梯度会得到一个惊人的简洁结果∂L / ∂z_k p_k - y_k这里 y_k 是 one-hot 的真实标签p_k 是 softmax 输出的概率。这个式子的意思是模型对第 k 类的梯度就是模型预测概率与真实标签之间的差值。如果真实类是第 y梯度会告诉模型你给这个类的概率还差 1-p_y 那么多请继续提高它对其他类梯度是你多给了 p_k 的概率请压下去。这也是为什么分类任务里 NLL 比 MSE sigmoid 组合更受欢迎。MSE 加 sigmoid 的梯度中会包含 sigmoid 的导数项在概率接近 0 或 1 时梯度趋近于 0直接导致梯度消失而 NLL 的梯度是 p - y模型越自信时 p 越接近 y梯度越小但它不会因为 sigmoid 饱和而人为地消失。这个差异在实际训练中的体验非常明显。5.2 如何通过 NLL 的数值判断模型状态看训练曲线时不要只看loss 降了没有NLL 的数值本身带着语义。我总结过几个自己经常用的参照初始 NLL 是否接近 log K。类别数是 K随机初始化的模型应当在 log K 附近。如果初始 loss 太低检查数据标签是否泄漏如果初始 loss 比 log K 高很多可能初始化不当或已有正则过强。训练集 NLL 降到接近 0验证集 NLL 开始反弹这是过拟合的明确信号。因为训练集模型已经对每个样本给出接近 1 的概率但这种自信无法泛化到验证集。NLL 为 inf 或 nan 时先别赖在函数实现上检查学习率和 logits 范围。通常学习率过高前几步更新后 logits 就溢出后续全乱套。训练 NLL 出现断崖式下降然后回升要怀疑某些样本的标签噪声过大因为 NLL 对噪声样本极其敏感错误的标签会导致损失异常大拉动模型走向奇怪的方向。5.3 label smoothing 的本质给 NLL 加一个均匀先验既然 NLL 鼓励模型变得极端自信而过度自信在过拟合时危害不小业界常用 label smoothing 来抑制这个问题。它的做法是把 one-hot 标签替换成y_k (1 - ε) · y_k ε / K损失从 -log p_y 变成了 -(1-ε)log p_y - (ε/K)∑_k log p_k。后面那一项其实就是模型输出与均匀分布之间的交叉熵它的作用是阻止模型把所有概率推到 1因为一旦这么做了均匀分布那一项就会带来惩罚。用 NLL 的框架来理解 label smoothing会很清楚它是在告诉模型我允许你保持一定的不确定性因为这通常更符合真实世界中的标签噪声。我在做大规模图像分类时ε 取 0.1 经常能带来 0.5% 到 1% 的精度提升而关键是验证集 NLL 更平滑曲线更稳。5.4 平均还是求和一个容易忽略但影响学习率的点计算 NLL 时理论上可以按 batch 平均也可以按 batch 求和。框架里对应的就是 reductionmean 和 reductionsum。这一点极其容易被忽略但影响巨大——同样的学习率在 average 模式下有效在 sum 模式下可能直接发散因为批量越大总损失越大梯度也就越大。我自己就在从 PyTorch 迁移到 JAX 时踩过这个坑用 jnp.sum 写 NLL忘了除以 batch size收敛速度变得完全不可控。所以当你看到别人的代码、或者其他框架的默认配置时第一件事就是确认他们用的是平均还是求和。这也顺带解释了为什么很多库的 CrossEntropyLoss 默认 reduction 是 mean——为了让损失值不随 batch 大小缩放方便跨实验比较。5.5 类别不平衡时 NLL 的隐患与加权处理NLL 天然对少数类不友好。如果某个类别在训练集里只出现 1% 的次数它的期望损失贡献自然就很小模型会倾向于忽略它。常见做法有两种一是直接按类别频率的倒数给样本加权相当于把损失函数改成加权 NLL二是采用 Focal Loss在 NLL 前面乘一个 (1-p_t)^γ 的调制因子让模型把注意力集中在难分样本上。这两种方法本质上都在调整 NLL 对不同样本的关注度了解 NLL 本身的构造后理解它们的动机就容易多了。我自己现在设计任何监督损失时第一件事永远是问这个问题的数据生成过程是什么目标变量服从什么分布对应的 NLL 是哪种形式想清楚这一步80% 的损失函数设计问题都能找到答案。至于数值稳定性、reduction 方式、与交叉熵的等价边界这些细节都是在实际调试中必须踩一遍才能理解的功课——希望这篇文章能帮你把这些弯路提前绕开。
阅读完成 · 觉得有帮助?
咨询建站