Math-05.信息论-05.最大熵与Softmax

本页说明 Softmax 为何是分类的「自然」输出层:最大熵原理(maximum entropy principle)在约束下选最不确定的分布,得到指数族 → Softmax。

段末注释最大熵原理 在满足矩约束(如 $\mathbb{E}[T(X)]=\boldsymbol{\mu}$)的所有分布中,选取熵 $H(X)$ 最大的那个;离散情形下为 Gibbs 分布 $p(x) \propto \exp(\boldsymbol{\lambda}^\top T(x))$。

系列入口00.系列规划 | 前置:03 交叉熵与 KL


1. 最大熵问题(D1–D3)

图 1 约束下的最平分布

离散 $K$ 类,约束 $\mathbb{E}[x_k] = \mu_k$(或等价:给定 logits 的充分统计量)。求:

$$
\max_p H(p) = -\sum_k p_k \log p_k \quad \text{s.t.} \quad \sum_k p_k = 1,\ \sum_k p_k x_k = \mu
$$

拉格朗日乘子法 → 解形为:

$$
p_k = \frac{\exp(\lambda x_k)}{Z(\lambda)}, \quad Z(\lambda) = \sum_j \exp(\lambda x_j)
$$

在固定均值约束下,Gibbs/指数族熵最大。


2. Softmax 与多项 logistic(D3)

图 2 logit → Softmax → 概率

分类 logits $\mathbf{z} = f_\theta(\mathbf{x}) \in \mathbb{R}^K$:

$$
p_k = \mathrm{softmax}(\mathbf{z})k = \frac{\exp(z_k)}{\sum{j=1}^K \exp(z_j)}
$$

性质

  • $p_k > 0$,$\sum_k p_k = 1$
  • 对 $z$ 平移不变:$\mathrm{softmax}(\mathbf{z}+c) = \mathrm{softmax}(\mathbf{z})$
  • 15.多项分布 参数化一致
  • 配合 CE = 指数族 NLL(Math-00/50

数值稳定:$\mathrm{softmax}(\mathbf{z}) = \mathrm{softmax}(\mathbf{z} - \max_k z_k)$,防 $\exp$ 溢出。


3. 温度(D3–D7)

图 3 温度与采样

温度 $T > 0$:

$$
p_k^{(T)} = \frac{\exp(z_k / T)}{\sum_j \exp(z_j / T)}
$$

$T$ 效果
$T \to 0$ 趋 one-hot(argmax)
$T = 1$ 标准 Softmax
$T > 1$ 分布更平、熵更高、采样更随机
$T < 1$ 更尖锐、更确定

应用:LLM 解码(temperature sampling)、知识蒸馏软标签、模拟退火。与 top-p(nucleus)采样配合使用。


4. 与 LogSumExp(D6)

$$
\log Z = \log \sum_k \exp(z_k) = \mathrm{LSE}(\mathbf{z})
$$

$\log p_k = z_k - \mathrm{LSE}(\mathbf{z})$。log_softmax + nll_loss 比先 softmaxlog 更稳(Math-03/07 数值稳定)。


5. 局限(D8)

图 4 局限

问题 说明
Softmax 饱和 $z$ 极端 → 梯度消失;LayerNorm、合适初始化
类别极多 $K$ 大时 CE 贵;采样 softmax、层次 softmax
温度≠校准 高温不一定概率校准;需 Platt/temperature scaling
多标签 非互斥类用 sigmoid + BCE,非单一 Softmax

6. PyTorch 示例(D12)

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
import torch
import torch.nn.functional as F

z = torch.tensor([2.0, 1.0, 0.1])

# 标准 softmax
p = F.softmax(z, dim=0)
print("p:", p, "H:", -(p * p.log()).sum())

# 温度
T = 2.0
p_T = F.softmax(z / T, dim=0)
print("p_T (更平):", p_T)

# log_softmax 用于 CE
log_p = F.log_softmax(z.unsqueeze(0), dim=-1)
target = torch.tensor([0])
loss = F.nll_loss(log_p, target)

7. 小结

Softmax = 多项 logit 的规范指数族;温度 调节输出熵。语言模型见 10 Perplexity

系列导航03 CE/KL | 10 语言模型

-------------本文结束感谢您的阅读-------------