动手学深度学习——RNN简洁实现
1. 前言
上一篇我们已经把 RNN 从零开始实现 了一遍。
那一篇的重点是“拆开看本质”:
-
token 如何做 one-hot
-
参数如何初始化
-
隐藏状态如何递推
-
每个时间步如何计算输出
-
如何根据前缀生成字符
这样做的好处是,RNN 的内部机制会看得非常清楚。
但在真实开发中,我们通常不会每次都手写这一整套,而是直接使用深度学习框架已经封装好的模块。
所以这一节就进入:
RNN 的简洁实现
这里的重点不是再重复讲 RNN 原理,而是看:
-
PyTorch 里的
nn.RNN怎么用 -
输入输出张量形状怎么对应
-
如何把现成 RNN 层接到语言模型里
-
它和“从零实现”到底对应在哪里
这篇本质上就是:
把手写版 RNN,换成框架版 RNN。
2. 为什么还要学“简洁实现”
这是一个很实际的问题。
既然上一节已经手写过了,为什么这一节还要再写一次?
原因很简单:
2.1 从零实现帮助理解原理
它让你真正知道 RNN 内部发生了什么。
2.2 简洁实现才更接近真实开发
在实际项目里,我们通常直接调用:
-
nn.RNN -
nn.GRU -
nn.LSTM
而不会从头手写矩阵递推。
2.3 两者对照,理解会更稳
如果你只会 API,不懂原理,很容易“会调不会讲”。
如果你只懂原理,不会框架实现,也不方便真正训练模型。
所以李沐这里先讲“从零实现”,再讲“简洁实现”,这个顺序非常合理。
3. PyTorch 里的 RNN 层是什么
PyTorch 已经给我们封装好了一个循环神经网络层:
nn.RNN
它本质上做的事情,和我们上一节手写的基础 RNN 是同一类运算:
H_t = tanh(X_t W_xh + H_{t-1} W_hh + b_h)
只不过这些参数初始化、时间步循环、状态传递,框架都帮我们做好了。
所以你可以把 nn.RNN 理解成:
官方封装好的基础循环神经网络单元
4. 先创建一个 RNN 层
最常见的写法类似这样:
import torch
from torch import nn
num_hiddens = 256
rnn_layer = nn.RNN(input_size=28, hidden_size=num_hiddens)
这里有两个最关键的参数。
input_size=28
表示每个时间步输入向量的维度。
如果这里做的是字符级语言模型,且词表大小是 28,
那每个字符经过 one-hot 后,输入维度就是 28。
hidden_size=num_hiddens
表示隐藏状态维度。
也就是每个时间步内部记忆向量的长度。
5. 为什么这里 input_size 常常等于词表大小
因为在这一阶段,李沐通常还是沿用 one-hot 表示。
假设词表大小是 28,那么:
-
每个 token 的 one-hot 向量长度就是 28
-
所以每个时间步输入特征维度就是 28
因此:
input_size = vocab_size
这是非常自然的。
后面如果换成词嵌入(embedding),那输入维度就不一定等于词表大小了。
6. 先看一个最小例子
通常会先构造一个输入试试看:
X = torch.rand(size=(35, 2, 28))
state = torch.zeros((1, 2, num_hiddens))
Y, state_new = rnn_layer(X, state)
这里三个张量的形状一定要看懂。
7. 输入 X 的形状为什么是 (35, 2, 28)
这三个维度分别表示:
-
35:时间步长度num_steps -
2:批量大小batch_size -
28:每个时间步输入向量维度input_size
也就是说:
这里一次性输入了 2 条序列,每条序列长度为 35,每个时间步是一个 28 维向量。
PyTorch 默认的 RNN 输入格式就是:
(num_steps, batch_size, input_size)
这一点和我们上一节手写版的输入组织方式是一致的。
8. 初始隐藏状态为什么是 (1, 2, num_hiddens)
这里的三个维度分别表示:
-
1:层数 × 方向数 -
2:batch_size -
num_hiddens:隐藏状态维度
因为这里:
-
只用了 1 层 RNN
-
不是双向 RNN
-
所以第一个维度就是 1
于是初始状态形状就是:
(1, batch_size, num_hiddens)
这和手写版的 (batch_size, num_hiddens) 相比,多了一个“层数维”。
9. 输出 Y 和新状态 state_new 的形状怎么看
继续看:
Y.shape, state_new.shape
通常会得到:
(torch.Size([35, 2, 256]), torch.Size([1, 2, 256]))
含义是:
Y.shape = (35, 2, 256)
表示每个时间步、每个样本,都输出了一个 256 维隐藏表示。
state_new.shape = (1, 2, 256)
表示最后一个时间步的隐藏状态。
你可以理解为:
-
Y保存了所有时间步的隐藏输出 -
state_new保存了序列结束时的最终状态
10. 为什么 Y 不是最终的词表预测
这是一个很容易混淆的点。
nn.RNN 输出的 Y 本质上还是:
每个时间步的隐藏表示
它的维度是:
hidden_size
而不是词表大小。
也就是说,RNN 层本身只负责:
-
吃序列
-
递推隐藏状态
-
产生隐藏特征
但它还没有完成最终的“字符分类”。
所以后面还需要再接一个线性层,把隐藏状态映射到词表大小。
11. 为什么还要再加一个输出层
因为语言模型的最终目标是:
对词表中每个 token 打分,预测下一个 token 是谁
而 RNN 层输出的是隐藏特征向量,
例如 256 维。
这 256 维不直接对应:
-
a -
b -
c -
...
所以还必须再接一个线性层:
hidden_size -> vocab_size
这样每个时间步才能变成对整个词表的分类结果。
12. 简洁实现里的完整模型怎么组织
李沐这里通常会再封装一个类,把:
-
RNN 层
-
输出线性层
-
one-hot 输入处理
整合到一起。
常见思路类似:
class RNNModel(nn.Module):
def __init__(self, rnn_layer, vocab_size):
super().__init__()
self.rnn = rnn_layer
self.vocab_size = vocab_size
self.num_hiddens = self.rnn.hidden_size
self.num_directions = 1
self.linear = nn.Linear(self.num_hiddens, self.vocab_size)
这个类的作用就是:
把 RNN 隐藏表示真正变成语言模型输出
13. 前向传播里要做什么
这个模型类里最关键的是 forward。
通常思路是:
第一步:把输入索引变成 one-hot
因为当前还没用 embedding。
第二步:送入 RNN 层
得到所有时间步的隐藏输出。
第三步:调整形状
把 (num_steps, batch_size, hidden_size) 变成二维张量,方便线性层处理。
第四步:用线性层映射到词表空间
得到最终分类分数。
14. 一个典型的前向传播写法
常见形式如下:
def forward(self, inputs, state):
X = nn.functional.one_hot(inputs.T.long(), self.vocab_size)
X = X.to(torch.float32)
Y, state = self.rnn(X, state)
output = self.linear(Y.reshape((-1, Y.shape[-1])))
return output, state
这段代码非常值得逐行理解。
15. 这一段 forward 怎么理解
15.1 inputs.T
因为原始输入通常是:
(batch_size, num_steps)
而 nn.RNN 默认需要:
(num_steps, batch_size, input_size)
所以先转置,把时间维放到前面。
15.2 one_hot
把每个 token 索引变成 one-hot 向量。
于是输入从:
(num_steps, batch_size)
变成:
(num_steps, batch_size, vocab_size)
15.3 Y, state = self.rnn(X, state)
这一步就是让 PyTorch 的 RNN 层完成所有时间步递推。
输出 Y 形状通常是:
(num_steps, batch_size, num_hiddens)
15.4 Y.reshape((-1, Y.shape[-1]))
把前两个维度合并。
例如:
(35, 2, 256) -> (70, 256)
这样就能直接送入线性层。
因为线性层通常是对最后一维特征做映射。
15.5 self.linear(...)
把隐藏状态映射成词表打分:
(70, 256) -> (70, vocab_size)
于是就得到了训练语言模型需要的输出形状。
16. 为什么这里要 reshape 成二维
因为我们的损失函数通常是交叉熵,
而它期待的输入往往是:
(样本数, 类别数)
在语言模型里:
-
每个时间步、每个样本,都是一个“分类样本”
-
所以
num_steps * batch_size个位置都可以合并起来统一算 loss
这就是为什么常常把时间步和 batch 两个维度合并。
17. 初始状态怎么写得更通用
对于简洁实现,状态初始化通常可以封装成:
def begin_state(self, device, batch_size=1):
if not isinstance(self.rnn, nn.LSTM):
return torch.zeros((self.num_directions * self.rnn.num_layers,
batch_size, self.num_hiddens),
device=device)
这里先只考虑普通 RNN 或 GRU。
如果是 LSTM,后面会变成两个状态。
这段代码的意义在于:
把状态初始化和模型结构绑定起来
这样无论 batch size 变化,还是后面层数变化,都更方便。
18. 简洁实现和从零实现一一对应在哪里
这是这一篇最该讲透的地方。
从零实现里的 W_xh, W_hh, b_h
在 nn.RNN 里被封装起来了。
从零实现里按时间步 for 循环
在 nn.RNN 内部自动完成了。
从零实现里的隐藏状态初始化
现在通过 begin_state() 来做。
从零实现里的输出层 W_hq, b_q
现在由 nn.Linear 负责。
所以你可以把简洁实现理解成:
手写版 RNN 的模块化封装版本
它们本质没有变,只是 API 更省事了。
19. 为什么简洁实现更适合真正训练
因为框架版实现通常有很多优势:
19.1 代码更短
更适合快速实验和工程开发。
19.2 计算更高效
框架内部通常做了优化。
19.3 更容易扩展
后面换成:
-
多层 RNN
-
GRU
-
LSTM
-
双向 RNN
接口都非常相似。
19.4 更不容易写错
特别是在状态维度、梯度传播这些地方,手写容易出 bug。
所以实际项目里,简洁实现是主流。
20. 训练时整体流程和上一节本质一样吗
本质一样。
无论是从零实现还是简洁实现,训练语言模型时的主线都没变:
输入
一批 token 索引序列。
模型
根据当前输入和历史状态,输出每个时间步的词表打分。
标签
右移一位后的真实 token。
损失
交叉熵。
优化
更新模型参数。
所以变的只是:
RNN 单元的实现方式
而不是整个语言模型训练逻辑。
21. 这一节代码最该掌握什么
如果从学习重点来看,最重要的是这几件事。
21.1 会创建 nn.RNN
知道:
-
input_size -
hidden_size
分别是什么意思。
21.2 看懂输入输出形状
特别是:
(num_steps, batch_size, input_size)
和
(num_steps, batch_size, hidden_size)
21.3 明白为什么还要接 nn.Linear
因为 RNN 输出的是隐藏表示,不是最终词表分类。
21.4 明白 reshape 的目的
为了统一计算每个时间步的分类损失。
21.5 能把它和“从零实现”对照起来
这样 API 才不会变成死记硬背。
22. 本节总结
这一节我们学习了 RNN 的简洁实现,核心内容可以总结为以下几点。
22.1 nn.RNN 封装了基础 RNN 的时间递推计算
它和上一节手写版本质一致。
22.2 输入通常先做 one-hot,再按时间维优先组织
这样才能送入 RNN 层。
22.3 RNN 层输出的是隐藏状态序列
不是最终词表预测。
22.4 需要再接一个线性层映射到词表空间
这样才能完成语言模型的多分类预测。
22.5 简洁实现更适合真实训练和后续扩展
这也是实际开发中更常用的方式。
23. 学习感悟
这一节最大的价值,在于它把“从零理解”和“工程实现”接上了。
如果只学手写版,你会觉得框架 API 太黑箱;
如果只学简洁版,你又可能根本不知道里面在算什么。
而当两者结合起来时,你就会真正明白:
nn.RNN并不神秘,它只是把我们上一节手写过的那套东西封装好了。
这种感觉非常重要。
因为后面学 GRU、LSTM 时,你会发现套路完全一样,只是内部单元更复杂。
更多推荐


所有评论(0)