沐神-动手学深度学习 课后习题3.4. softmax回归
·
这篇文章是《动手学深度学习》第 3.4 节,它将线性回归的概念扩展到了分类问题(Classification)。
- 网页内容核心总结
线性回归预测的是“多少”(连续值),而 Softmax 回归 预测的是“是什么”(离散类别)。
A. 分类问题的表示
- 独热编码 (One-hot Encoding):由于类别之间通常没有自然的数值大小关系(例如“猫”不等于“狗”的一半),我们使用向量来表示类别。如猫表示为 (1, 0, 0),狗表示为 (0, 1, 0)。
- 网络结构:Softmax 回归是一个单层神经网络。每个类别对应一个输出,输出的个数等于类别的总数。
B. Softmax 运算
为了将网络的原始输出(logits)转化为合法的概率分布(即所有项非负且和为 1),引入了 Softmax 函数:
\hat{y}_j = \frac{\exp(o_j)}{\sum_k \exp(o_k)}
它不仅保证了概率属性,还通过指数运算拉大了最大值与其他值之间的差距。
C. 损失函数:交叉熵 (Cross-Entropy)
对于分类问题,我们不关心预测值和真实值之间的“距离”(平方误差),而关心预测概率分布与真实分布的匹配程度。
- 公式:l(\mathbf{y}, \hat{\mathbf{y}}) = -\sum_j y_j \log \hat{y}_j。
- 直觉:由于 y 是独热向量,损失函数实际上只关注模型对“正确类别”所分配的概率。正确类别的概率越接近 1,损失越低。
- 习题详细答案
Q1: 我们可以通过将 Softmax 运算应用到输出层来计算导数。证明交叉熵损失关于偏置 b 的导数是 \hat{y} - y。
解析:
这是一个非常经典的结论。在线性回归中,误差是 \hat{y} - y;在 Softmax 回归中,梯度依然是预测概率分布与真实分布的差。
- 推导思路:
- 写出损失函数 L 关于输出 o 的表达式。
- 利用链式法则 \frac{\partial L}{\partial b} = \frac{\partial L}{\partial o} \cdot \frac{\partial o}{\partial b}。
- 由于 \frac{\partial o}{\partial b} = 1,最终结果简化为预测概率减去标签值 \text{softmax}(o) - y。
Q2: 既然我们有交叉熵损失,为什么还要用 Softmax?如果我们直接用 o 计算损失会怎样?
答案:
- 概率约束:o 的取值范围是 (-\infty, \infty),直接使用无法直接解释为概率。
- 指数放大:Softmax 倾向于让正确的类别脱颖而出。
- 数值稳定性:如果没有 Softmax 的指数处理,对数损失在处理非正数时会出错(\log 无法处理负数)。
Q3: 为什么分类问题通常不使用平方损失(MSE)?
答案: - 梯度消失:当预测非常错误时,平方损失的梯度往往变得很小,导致模型学习缓慢。而交叉熵在错误严重时梯度依然很大。
- 非凸优化:平方损失配合 Softmax 往往会导致非凸的优化表面,增加训练难度。
- 不符合统计假设:平方损失假设误差服从高斯分布,而分类问题的输出通常服从多项分布。
Q4: 在实际应用中,如果类别数量非常大(例如几十万个),Softmax 会有什么问题?
答案: - 计算开销:分母上的求和操作 \sum \exp(o_k) 需要遍历所有类别,当类别极多时,计算非常缓慢。
- 解决方案:在自然语言处理等领域,通常会使用 层次 Softmax (Hierarchical Softmax) 或 负采样 (Negative Sampling) 来近似计算。
Q5: 假设我们有三个类别,对应的权重向量分别为 w_1, w_2, w_3。如果我们将所有 w_k 都加上一个常数向量 c,预测结果会变吗?
答案:
不会改变。 - 证明:在 Softmax 公式中,分子分母同时乘以 \exp(c^\top x),这个项会相互抵消。这说明 Softmax 回归存在参数冗余,真正起作用的是权重之间的差值。
下一步建议:
理论已经铺垫完成,下一步是 3.5 节“图像分类数据集”。你会接触到著名的 Fashion-MNIST 数据集,并学习如何加载真实的图像数据进行训练。
你想让我帮你梳理一下 Fashion-MNIST 图像加载 的代码逻辑,还是直接看 3.6 节的从零实现?
针对你上传的图片中《动手学深度学习》softmax回归章节末尾的进阶习题,我为你整理了详细的逻辑分析与解答:
- 指数族与softmax的联系
1.1 计算softmax交叉熵损失 l(\mathbf{y}, \hat{\mathbf{y}}) 的二阶导数
- 一阶导数(回顾):在前文中我们已知,对于输出 o_i 的导数为 \frac{\partial l}{\partial o_i} = \hat{y}_i - y_i。
- 二阶导数(Hessian矩阵):我们需要对一阶导数再次求导。
- 当 i=j 时:\frac{\partial^2 l}{\partial o_i^2} = \frac{\partial \hat{y}_i}{\partial o_i} = \hat{y}_i(1 - \hat{y}_i)。
- 当 i \neq j 时:\frac{\partial^2 l}{\partial o_i \partial o_j} = \frac{\partial \hat{y}_i}{\partial o_j} = -\hat{y}_i \hat{y}_j。
- 结论:二阶导数矩阵(Hessian)可以表示为 \text{diag}(\hat{\mathbf{y}}) - \hat{\mathbf{y}}\hat{\mathbf{y}}^\top。
1.2 计算softmax(o)给出的分布方差,并与二阶导数匹配 - 背景:在概率论中,多项分布的协方差矩阵正好是 \text{diag}(\mathbf{p}) - \mathbf{p}\mathbf{p}^\top。
- 匹配:你会发现,损失函数的二阶导数矩阵(Hessian)恰好等于模型预测概率分布的协方差矩阵。这在统计学上意味着交叉熵损失的曲率完全由预测的不确定性(方差)决定。
- 概率向量为 (1/3, 1/3, 1/3) 的编码问题
2.1 如果我们尝试为它设计二进制代码,有什么问题?
- 信息量计算:根据信息论,每个类别的自信息量为 -\log_2(1/3) \approx 1.58 比特。
- 问题:二进制编码只能是整数(如 1 位或 2 位)。如果用 1 位不够表示 3 类,用 2 位则会浪费空间(平均码长 2 > 1.58),导致编码效率不高。
2.2 设计更好的代码(联合编码) - 两个独立观察结果:如果联合编码两个样本,共有 3^2=9 种可能。-\log_2(1/9) \approx 3.17 比特。用 3 位二进制无法表示,用 4 位则平均每个样本用 2 位,依然有浪费。
- n个观测值:随着 n \to \infty,根据 香农第一信源编码定理,我们可以通过对超长序列进行联合编码,使平均每个样本消耗的比特数无限接近于 1.58。
- 关于 RealSoftMax 的证明
3.1 证明 \text{RealSoftMax}(a, b) > \max(a, b)
- 假设 a \geq b,则 \exp(a) + \exp(b) > \exp(a)。
- 两边取对数:\log(\exp(a) + \exp(b)) > \log(\exp(a)) = a。
- 既然结果大于 a 且 a = \max(a, b),则命题得证。
3.2 证明 \lambda^{-1}\text{RealSoftMax}(\lambda a, \lambda b) > \max(a, b) - 同理,\log(\exp(\lambda a) + \exp(\lambda b)) > \lambda \max(a, b)。
- 由于 \lambda > 0,两边除以 \lambda 保持不等号方向,结论成立。
3.3 证明对于 \lambda \to \infty,该式趋近于 \max(a, b) - \lambda^{-1} \log(\exp(\lambda a) + \exp(\lambda b)) = \lambda^{-1} \log(\exp(\lambda a)(1 + \exp(\lambda(b-a))))(假设 a > b)。
- 展开得:a + \frac{\log(1 + \exp(\lambda(b-a)))}{\lambda}。
- 当 \lambda \to \infty 时,b-a < 0,\exp(\lambda(b-a)) \to 0,第二项趋于 0。最终结果收敛于 a(即 \max(a, b))。
3.4 soft-min 会是什么样子? - \text{soft-min}(a, b) = -\log(\exp(-a) + \exp(-b))。它通过对输入取负,利用 softmax 的特性寻找最小值。
3.5 扩展到两个以上的数字 - \text{RealSoftMax}(x_1, \dots, x_n) = \log\left(\sum_{i=1}^n \exp(x_i)\right)。这在数学上被称为 LogSumExp 函数,是深度学习中计算 log-probability 的标准做法。
下一步建议:
这些题目涉及到信息论和数值稳定性。在实际编程实现中,直接计算 exp 容易溢出,所以 PyTorch 内部通常使用 LogSoftmax。
你想看看 3.7 节 Softmax 的简洁实现 中是如何处理这些数值稳定性的吗?
更多推荐

所有评论(0)