RNN 代码 - 基础版
1 课程概览
本课演示传统 RNN 的基础代码实现,讲解 nn.RNN API 的参数含义和输入输出形状。核心参数为 input_size=5、hidden_size=6、num_layers=1,输入数据形状为 (3, 1, 5),输出形状为 (1, 3, 6) 和 (1, 1, 6)。
2 核心概念与定义
- nn.RNN:PyTorch 中传统 RNN 的 API。
- input_size:输入特征的维度(词向量维度)。
- hidden_size:隐藏状态的维度(输出维度)。
- num_layers:隐藏层的层数。
- 序列数据(Sequence Data):后一个数据对前一个数据有依赖关系的数据。
3 模型与算法详解
RNN API 参数
nn.RNN(input_size=5, hidden_size=6, num_layers=1)
| 参数 | 值 | 含义 |
|---|---|---|
| input_size | 5 | 输入特征维度(词向量维度) |
| hidden_size | 6 | 隐藏状态维度(输出维度) |
| num_layers | 1 | 隐藏层层数 |
输入数据形状
# 输入数据形状:(seq_len, batch, input_size) = (3, 1, 5)
input = torch.randn(3, 1, 5)
| 维度 | 值 | 含义 |
|---|---|---|
| seq_len | 3 | 句子长度(3 个单词) |
| batch | 1 | 批次大小(1 个句子) |
| input_size | 5 | 每个单词的特征维度 |
输出形状
output, hn = rnn(input, h0)
# output 形状:(seq_len, batch, hidden_size) = (3, 1, 6)
# hn 形状:(num_layers, batch, hidden_size) = (1, 1, 6)
| 输出 | 形状 | 含义 |
|---|---|---|
| output | (3, 1, 6) | 所有时间步的输出 |
| hn | (1, 1, 6) | 最后一个时间步的隐藏状态 |
RNN 分类回顾
| 按输入输出 | 应用场景 |
|---|---|
| N vs N | 对联、诗词 |
| N vs 1 | 情感分析、文本分类、意图识别 |
| 1 vs N | 看图说话 |
| N vs M | 翻译、文本生成、摘要 |
4 数学原理与推导
RNN 前向传播公式
$$H_t = \tanh(W_{ih} X_t + W_{hh} H_{t-1} + b_h)$$
$$Y_t = W_{ho} H_t + b_o$$
输入输出关系
- 输入:$X_t$(本次输入)、$H_{t-1}$(上一时刻隐藏状态)
- 输出:$H_t$(本次隐藏状态)、$Y_t$(本次输出)
5 代码示例
import torch
import torch.nn as nn
# 1. 定义 RNN 模型
rnn = nn.RNN(input_size=5, hidden_size=6, num_layers=1)
# 2. 创建输入数据
# 形状:(seq_len, batch, input_size) = (3, 1, 5)
# 3 个时间步,1 个批次,每个单词 5 个特征
input = torch.randn(3, 1, 5)
# 3. 初始化隐藏状态
# 形状:(num_layers, batch, hidden_size) = (1, 1, 6)
h0 = torch.zeros(1, 1, 6)
# 4. 前向传播
output, hn = rnn(input, h0)
# 5. 查看输出形状
print("output shape:", output.shape) # torch.Size([3, 1, 6])
print("hn shape:", hn.shape) # torch.Size([1, 1, 6])
# 6. 查看输出内容
print("output:", output)
print("hn:", hn)
# hn 等于 output 的最后一个时间步
print("output[-1]:", output[-1]) # 与 hn 相同
输出示例
output shape: torch.Size([3, 1, 6])
hn shape: torch.Size([1, 1, 6])
6 重难点与易错提醒
- ❗重点:RNN API 的三个核心参数:
input_size、hidden_size、num_layers。 - ❗重点:输入数据形状为
(seq_len, batch, input_size)。 - ⚠️易错:
output是所有时间步的输出,hn是最后一个时间步的隐藏状态。 - ⚠️易错:
output[-1]与hn的内容相同。 - 💡深入理解:
seq_len是句子长度,batch是批次大小,input_size是词向量维度。
7 课堂问答精选
Q1:RNN API 的三个核心参数是什么?
A:input_size(输入特征维度)、hidden_size(隐藏状态维度)、num_layers(隐藏层层数)。
Q2:输入数据的形状是什么?
A:(seq_len, batch, input_size),即(句子长度,批次大小,词向量维度)。
Q3:output 和 hn 有什么区别?
A:output 是所有时间步的输出,形状为 (seq_len, batch, hidden_size)。hn 是最后一个时间步的隐藏状态,形状为 (num_layers, batch, hidden_size)。output[-1] 与 hn 内容相同。
Q4:RNN 按输入输出分为哪几种?
A:四种:N vs N(对联、诗词)、N vs 1(情感分析、文本分类)、1 vs N(看图说话)、N vs M(翻译、文本生成、摘要)。
8 本课小结
- RNN API:
nn.RNN(input_size, hidden_size, num_layers)。 - 输入形状:
(seq_len, batch, input_size)。 - 输出形状:
output为(seq_len, batch, hidden_size),hn为(num_layers, batch, hidden_size)。 output[-1]与hn内容相同。- RNN 分类:N vs N、N vs 1、1 vs N、N vs M。
9 延伸思考
- 如何调整 RNN 的参数以适应不同的任务?
- RNN 的批次大小如何影响训练效率?