AI

d2l 学习笔记:softmax 回归

·11 分钟阅读·4130 字

整理 Softmax 回归从线性模型扩展到多分类任务的数学原理、手动实现与训练流程

📋 目录

d2l 学习笔记:softmax 回归

1. Softmax 回归解决的问题

Softmax 回归用于把线性回归的输出形式扩展到多分类问题。线性回归输出一个连续值,适合预测房价、温度、销量等数值;多分类任务需要输出每个类别的概率。

以图像分类为例,模型输入一张图片后,需要输出每个类别的预测概率:

T-shirt     0.80
Trouser     0.10
Pullover    0.05
Sneaker     0.05

概率输出需要满足两个条件:

0≤pi≤10 \leq p_i \leq 1 ∑ipi=1\sum_i p_i = 1

Softmax 的作用就是把线性模型产生的未归一化分数转换为概率分布。

2. 数据与任务形式

D2L 章节中常用 Fashion-MNIST 作为 Softmax 回归示例数据集。

2.1. 输入数据

Fashion-MNIST 中每张图片大小为:

28×2828 \times 28

展开后变成一个长度为 784 的向量:

x=(x1,x2,…,x784)x=(x_1,x_2,\dots,x_{784})

批量输入时,若 batch size 为 nn,输入矩阵形状为:

X∈Rn×784X \in \mathbb{R}^{n \times 784}

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)

其中:

W∈R784×10W \in \mathbb{R}^{784 \times 10} b∈R10b \in \mathbb{R}^{10}

对于一个样本 xx,线性部分计算为:

o=xW+bo=xW+b

其中 oo 是长度为 10 的向量:

o=(o1,o2,…,o10)o=(o_1,o_2,\dots,o_{10})

oo 中的每个元素表示一个类别的未归一化分数,也称为 logits。logits 还不是概率,可能是任意实数。

4. Softmax 函数

4.1. 数学定义

Softmax 将 logits 转换为概率分布:

y^i=eoi∑jeoj\hat{y}_i=\frac{e^{o_i}}{\sum_j e^{o_j}}

其中:

  • oio_i:第 ii 个类别的 logit。
  • eoie^{o_i}:把 logit 转换为正数。
  • ∑jeoj\sum_j e^{o_j}:所有类别指数值之和。
  • y^i\hat{y}_i:第 ii 个类别的预测概率。

4.2. 为什么使用指数

指数函数有两个作用:

  1. 保证输出为正数。
  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. 图片展开

原始图片形状为:

28×2828 \times 28

模型计算前需要展开为:

784784

对应代码:

X.reshape((-1, W.shape[0]))

如果 batch size 为 100,则展开后的形状为:

100×784100 \times 784

5.2. 线性计算

矩阵乘法计算为:

XW+bXW+b

若:

X∈R100×784X \in \mathbb{R}^{100 \times 784} W∈R784×10W \in \mathbb{R}^{784 \times 10}

则输出 logits 形状为:

100×10100 \times 10

表示 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. 为什么需要交叉熵

模型输出的是预测概率 y^\hat{y},真实标签可以用 one-hot 向量 yy 表示。

例如真实类别是第一个类别:

y=[1,0,0]y=[1,0,0]

如果模型预测为:

y^=[0.1,0.7,0.2]\hat{y}=[0.1,0.7,0.2]

说明模型把真实类别的概率只给到了 0.1,预测效果较差。

交叉熵用于衡量预测概率分布与真实标签分布之间的差距:

l(y,y^)=−∑iyilog⁡y^il(y,\hat{y})=-\sum_i y_i \log \hat{y}_i

由于 one-hot 标签中只有真实类别的位置为 1,交叉熵可以简化为:

l=−log⁡y^truel=-\log \hat{y}_{\text{true}}

其中 y^true\hat{y}_{\text{true}} 是模型分配给真实类别的概率。

6.2. 交叉熵的直观含义

如果模型给真实类别的概率很高:

y^true=0.99\hat{y}_{\text{true}}=0.99

则损失很小:

−log⁡(0.99)-\log(0.99)

如果模型给真实类别的概率很低:

y^true=0.01\hat{y}_{\text{true}}=0.01

则损失很大:

−log⁡(0.01)-\log(0.01)

交叉熵本质上是在惩罚模型没有给真实类别足够高的概率。

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. 准确率计算

准确率计算公式为:

accuracy=预测正确样本数总样本数\text{accuracy}=\frac{\text{预测正确样本数}}{\text{总样本数}}

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()

这一步通过自动微分计算参数 WW 和 bb 的梯度。

8.4. 参数更新

updater(X.shape[0])

这一步根据梯度更新参数,使下一轮预测更接近真实标签。

9. 完整训练流程

Softmax 回归训练流程可以概括为:

读取图片
  -> 展开为向量
  -> 线性计算得到 logits
  -> Softmax 转换为概率
  -> 交叉熵计算损失
  -> 反向传播计算梯度
  -> 更新 W 和 b
  -> 重复多轮训练

从数学角度看,Softmax 回归就是在优化参数 WW 和 bb,使真实类别对应的预测概率尽可能高。

10. 关键知识点

知识点说明
logits模型线性层输出的未归一化分数。
Softmax将 logits 转换为概率分布。
one-hot 标签用向量表示真实类别,真实类别位置为 1。
交叉熵惩罚真实类别预测概率过低。
argmax取概率最大的位置作为预测类别。
反向传播根据损失计算参数梯度。

Softmax 回归是理解神经网络分类任务的基础。后续的多层感知机、卷积神经网络和 Transformer 分类头,都会延续“logits -> Softmax 或 CrossEntropy -> 参数更新”的基本思路。


关联文档

Yanche Blog

记录云原生、Linux、数据库等技术领域的学习心得,以及日常生活的思考与感悟。

© 2026 Yanche Blog. All rights reserved.

Powered by Astro