有次线上模型迭代,我在自定义模型头时图省事,手动把 logits 过了一遍 softmax 再取 log,结果验证集上 loss 全部变成 nan。排查了一个下午,最后发现是 float 精度的问题——softmax 之后概率已经接近 0 的位置,再取 log 直接给到 -inf,负号和另一个 inf 一碰就成 nan。那时候我才意识到,负对数似然这个每天在框架里自动执行的操作,我从没有认真把每一步的原理和工程细节想清楚。
这篇文章想把这个话题彻底展开。负对数似然函数(Negative Log-Likelihood,NLL)看着只是"取概率的负对数"这么一个简单动作,实际上它贯穿了统计推断、分类损失设计、信息论、数值稳定实现、训练状态诊断这些完全不同的层次。不管你是刚接触机器学习、搞不清损失函数来源的初学者,还是已经在调模型但想补理论短板的工程师,这篇文章都值得读一读。我会从最大似然估计出发,把它为什么长成这个样子讲明白,再落到分类任务里的具体形态、与交叉熵的关系、工程实现的坑,以及训练过程中通过 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的真实形态:从伯努利到Softmax
2.1 二分类:logistic loss 就是 NLL
二分类问题的真实标签 y 只有 0 和 1 两种情况。模型输出通常是一个经过 sigmoid 的值 p = P(y=1 | x),同时 p 也隐含了 P(y=0 | x) = 1-p。把这两个情况合并写成一个公式:
P(y | x) = p^y · (1-p)^(1-y)
当 y=1 时,只剩 p^1(1-p)^0 = p;当 y=0 时,只剩 p^0(1-p)^1 = 1-p。取负对数后,得到:
NLL = -[y log p + (1-y) log(1-p)]
这就是大家熟悉的 logistic loss,也是二元交叉熵。注意它的行为:如果真实标签是 1,模型给 p=0.99,那么损失大约是 0.01;模型给 p=0.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 类分类,模型随机初始化后跑第一个 batch,NLL 应该在 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-exp
4.1 前端最常踩的坑:先 softmax 再 log
回到开头那个 nan 的故事。很多人在自实现时按照数学公式的顺序:先算 softmax 得到概率 p,再计算 -log(p)。数学上没错,但在浮点数的世界里,softmax 出来的概率很多是极小值。一个 1000 类任务,某个类别的概率可能小于 1e-30,float32 还能勉强表示,但如果你用了 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 为 0,log(0) 就出现;二是因为经过了非线性乘法,梯度数值也会非常不稳定。
更稳妥的做法是在 log 空间里累加每个任务的 log p_j,等价但数值性质好很多。我在做多任务模型时,早期就是这么吃亏的,后来统一改成 log-sum-exp 和 log-prob 的方式,才彻底摆脱了概率下溢导致的偶发短路。
5. 训练中NLL曲线教会我的事:梯度、过拟合与label smoothing
5.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 求和。框架里对应的就是 reduction='mean' 和 reduction='sum'。这一点极其容易被忽略,但影响巨大——同样的学习率,在 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 方式、与交叉熵的等价边界这些细节,都是在实际调试中必须踩一遍才能理解的功课——希望这篇文章能帮你把这些弯路提前绕开。