Softmax 回归是为 K 个互斥类别分配概率的线性模型,常用于单标签分类。它对每个类别输出一个未归一化的分数,即 Logit:
o=W⊤x+b.
用一个三分类例子贯穿计算:logits 为 [log2,0,0],取指数得到 [2,1,1],总和是 4。每项除以 4,就得到概率 [1/2,1/4,1/4]。
Softmax 函数将这些 Logit 转换为非负且总和为 1 的概率值:
p(y=k∣x)=∑j=1Kexp(oj)exp(ok).
在工程实现中,为了数值稳定,通常会在计算指数前减去最大的 Logit。由于 Softmax 对整体平移不变(即给所有 Logit 加上同一常数结果不变),这一操作不会改变最终结果。Softmax 保持分数的大小顺序,因此推理时可直接对 logits 取 argmax。
交叉熵
对于独热编码(One-hot)的目标 y,单个样本的交叉熵损失如下,其中 py 指模型给观测类别的概率:
ℓ(y,p)=−k=1∑Kyklogpk=−logpy.
假设第二类才是正确答案,独热目标就是 y=[0,1,0]。只有这一类的概率进入损失:ℓ=−log(1/4)=log4≈1.3863。模型给第一类的概率最高,但提高第二类的概率才能降低这个样本的损失。
最小化该损失,等价于在模型假设下最大化观测类别标签的条件似然。
信息论中的交叉熵比较目标分布与预测分布。这里独热向量 y 表示单个已观察样本的经验目标,p 是模型预测;使用自然对数时,损失单位是 nat。这不意味着总体中给定输入后的真实标签分布也一定是独热分布。
稳定损失与梯度
若输入有 d 个特征,W 的形状为 d×K,b 长度为 K;B 行输入产生 B×K 个 logits。用 c 表示正确类别的索引,与独热向量 y 区分。Softmax 推导给出
ℓ=m+logj∑eoj−m−oc,m=jmaxoj,∂oj∂ℓ=pj−yj.
应直接计算这个 log-sum-exp 表达式,或使用接收 logits 的融合交叉熵损失。即使 Softmax 已做稳定处理,先把概率舍入再取对数,仍可能得到无穷大。
回到同一个例子,m=log2、oc=0,稳定损失为 log2+log(1+1/2+1/2)−0=log4。所有 logits 同加 1000 时,这个常数会在表达式中抵消,损失和概率都不变。
用预测概率减去目标,就能看出分数该往哪个方向调整:
p−y=[1/2,1/4,1/4]−[0,1,0]=[1/2,−3/4,1/4].
第二项为负,说明应提高正确类的 logit;其余项为正,说明应降低对应分数。若直接对这些 logits 做一步梯度下降,就要减去这个向量的一个正倍数。
从 logit 接到可训练参数,只需对 ℓ=log∑keok−oc 求导,再代入 oj=∑iWijxi+bj:
∂oj∂ℓ=pj−1j=c,∂Wij∂ℓ=xi(pj−yj),∂bj∂ℓ=pj−yj.
这就是穿过仿射 logit 层的反向自动微分:损失梯度往回传递,并乘以对应输入。矩阵式 ∇Wℓ=x(p−y)T 只是把这些逐坐标导数放在一起。批次平均损失对应各样本梯度的平均值。
适用边界
- 互斥多分类: Softmax 回归常用于每个样本对应一个类别标签的任务。Softmax 函数还有其他用途,例如计算注意力权重。
- 多标签分类: 若多个标签可能同时成立,应使用独立输出,通常配合 Sigmoid 激活函数和二元交叉熵损失。
- 决策规则: 默认使用
argmax 选取类别,但在误报与漏报成本不对称的场景下,可能需要调整决策阈值。
- 概率校准: 分类准确率高不代表概率预测是校准良好的(Calibrated)。
- 线性边界: 若特征间存在非线性关系,需进行特征变换或采用表达能力更强的模型。
评估时,除了对比类别频率和简单基线,务必检查混淆矩阵及各类别的错误分布。整体准确率可能会掩盖模型在稀有或关键类别上的失效。
接下来阅读 多层感知机,引入非线性隐藏层表示。推导过程及维护良好的实现代码可参考 Dive into Deep Learning: Linear Neural Networks for Classification。