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 时,你会发现套路完全一样,只是内部单元更复杂。

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐