d2l 学习笔记:softmax 回归
1. Softmax 回归解决的问题
Softmax 回归用于把线性回归的输出形式扩展到多分类问题。线性回归输出一个连续值,适合预测房价、温度、销量等数值;多分类任务需要输出每个类别的概率。
以图像分类为例,模型输入一张图片后,需要输出每个类别的预测概率:
T-shirt 0.80
Trouser 0.10
Pullover 0.05
Sneaker 0.05
概率输出需要满足两个条件:
Softmax 的作用就是把线性模型产生的未归一化分数转换为概率分布。
2. 数据与任务形式
D2L 章节中常用 Fashion-MNIST 作为 Softmax 回归示例数据集。
2.1. 输入数据
Fashion-MNIST 中每张图片大小为:
展开后变成一个长度为 784 的向量:
批量输入时,若 batch size 为 ,输入矩阵形状为:
2.2. 输出类别
Fashion-MNIST 共有 10 个类别,因此模型需要为每个样本输出 10 个类别分数。
0 T-shirt
1 Trouser
2 Pullover
3 Dress
4 Coat
5 Sandal
6 Shirt
7 Sneaker
8 Bag
9 Ankle boot
模型输入维度是 784,输出维度是 10。
3. 模型参数
Softmax 回归本质上仍然是一个线性模型,只是在线性输出后增加 Softmax 归一化。
num_inputs = 784
num_outputs = 10
W = torch.normal(0, 0.01, size=(num_inputs, num_outputs))
b = torch.zeros(num_outputs)
其中:
对于一个样本 ,线性部分计算为:
其中 是长度为 10 的向量:
中的每个元素表示一个类别的未归一化分数,也称为 logits。logits 还不是概率,可能是任意实数。
4. Softmax 函数
4.1. 数学定义
Softmax 将 logits 转换为概率分布:
其中:
- :第 个类别的 logit。
- :把 logit 转换为正数。
- :所有类别指数值之和。
- :第 个类别的预测概率。
4.2. 为什么使用指数
指数函数有两个作用:
- 保证输出为正数。
- 放大不同类别分数之间的差距。
例如 logits 为:
3.0
1.0
0.0
指数变换后约为:
20.09
2.72
1.00
再除以总和后得到概率分布。
4.3. 手动实现
def softmax(X):
X_exp = torch.exp(X)
partition = X_exp.sum(1, keepdim=True)
return X_exp / partition
keepdim=True 的作用是保持维度,便于后续按行广播相除。
5. 前向传播
Softmax 回归的前向传播流程包括图片展开、线性计算和 Softmax 归一化。
def net(X):
return softmax(torch.matmul(X.reshape((-1, W.shape[0])), W) + b)
5.1. 图片展开
原始图片形状为:
模型计算前需要展开为:
对应代码:
X.reshape((-1, W.shape[0]))
如果 batch size 为 100,则展开后的形状为:
5.2. 线性计算
矩阵乘法计算为:
若:
则输出 logits 形状为:
表示 100 个样本各自对应 10 个类别分数。
5.3. 概率输出
Softmax 作用在每一行 logits 上,使每个样本输出一组概率。
样本 1: [0.80, 0.10, 0.05, 0.05]
样本 2: [0.20, 0.70, 0.05, 0.05]
每一行概率之和为 1。
6. 交叉熵损失
6.1. 为什么需要交叉熵
模型输出的是预测概率 ,真实标签可以用 one-hot 向量 表示。
例如真实类别是第一个类别:
如果模型预测为:
说明模型把真实类别的概率只给到了 0.1,预测效果较差。
交叉熵用于衡量预测概率分布与真实标签分布之间的差距:
由于 one-hot 标签中只有真实类别的位置为 1,交叉熵可以简化为:
其中 是模型分配给真实类别的概率。
6.2. 交叉熵的直观含义
如果模型给真实类别的概率很高:
则损失很小:
如果模型给真实类别的概率很低:
则损失很大:
交叉熵本质上是在惩罚模型没有给真实类别足够高的概率。
6.3. 手动实现交叉熵
def cross_entropy(y_hat, y):
return -torch.log(y_hat[range(len(y_hat)), y])
示例:
y_hat = torch.tensor([
[0.1, 0.3, 0.6],
[0.3, 0.2, 0.5]
])
y = torch.tensor([0, 2])
索引表达式:
y_hat[range(len(y_hat)), y]
表示从每个样本的预测概率中取出真实类别对应的概率:
样本 1 取第 0 类概率:0.1
样本 2 取第 2 类概率:0.5
再取负对数得到每个样本的损失。
7. 分类准确率
准确率用于衡量模型预测类别是否与真实标签一致。
def accuracy(y_hat, y):
if len(y_hat.shape) > 1 and y_hat.shape[1] > 1:
y_hat = y_hat.argmax(axis=1)
cmp = y_hat.type(y.dtype) == y
return float(cmp.type(y.dtype).sum())
7.1. argmax 的作用
对于每个样本,argmax(axis=1) 取概率最大的类别作为预测类别。
[0.1, 0.7, 0.2] -> 1
如果真实标签也是 1,则该样本预测正确。
7.2. 准确率计算
准确率计算公式为:
8. 单轮训练流程
一轮训练通常包括前向传播、计算损失、反向传播和参数更新。
def train_epoch_ch3(net, train_iter, loss, updater):
for X, y in train_iter:
y_hat = net(X)
l = loss(y_hat, y)
l.sum().backward()
updater(X.shape[0])
8.1. 前向传播
y_hat = net(X)
这一步将输入图片转换为类别概率。
8.2. 损失计算
l = loss(y_hat, y)
这一步计算预测概率和真实标签之间的差距。
8.3. 自动求梯度
l.sum().backward()
这一步通过自动微分计算参数 和 的梯度。
8.4. 参数更新
updater(X.shape[0])
这一步根据梯度更新参数,使下一轮预测更接近真实标签。
9. 完整训练流程
Softmax 回归训练流程可以概括为:
读取图片
-> 展开为向量
-> 线性计算得到 logits
-> Softmax 转换为概率
-> 交叉熵计算损失
-> 反向传播计算梯度
-> 更新 W 和 b
-> 重复多轮训练
从数学角度看,Softmax 回归就是在优化参数 和 ,使真实类别对应的预测概率尽可能高。
10. 关键知识点
| 知识点 | 说明 |
|---|---|
| logits | 模型线性层输出的未归一化分数。 |
| Softmax | 将 logits 转换为概率分布。 |
| one-hot 标签 | 用向量表示真实类别,真实类别位置为 1。 |
| 交叉熵 | 惩罚真实类别预测概率过低。 |
| argmax | 取概率最大的位置作为预测类别。 |
| 反向传播 | 根据损失计算参数梯度。 |
Softmax 回归是理解神经网络分类任务的基础。后续的多层感知机、卷积神经网络和 Transformer 分类头,都会延续“logits -> Softmax 或 CrossEntropy -> 参数更新”的基本思路。
关联文档
- d2l 学习笔记:引言:D2L 学习入口。
- d2l 学习笔记:线性回归的损失函数:线性回归与损失函数基础。
- d2l 学习笔记:softmax 工程实现:Softmax 数值稳定性与框架实现。
- d2l 学习笔记:交叉熵:交叉熵的信息论解释与分类损失含义。
