# 背景及推导
需求:输入一句话,输出每个单词的褒贬性

假设对于上面 5 个单词,每个单词采用 300 维的 词向量,那么输入端就需要 1500 个词向量,这会带来几个缺点:
- 输入层的节点太多了,并且是变长的,会随着句子的长短发生变化
- 无法体现词语之间的先后顺序,仅仅只是把他们拉直展开成一个大向量,直接送入输入层(类似于 MLP)
为了解决上述缺点,可以让上一步的信息参与下一步的运算,具体可以这样做:


-
第一个词 经过非线性变换 g() 后得到一个中间结果(
隐藏状态) ,它再经过一次非线性变换得到第一个输出 -
接着,第二个词 与刚才的隐藏状态 一起参与非线性变换,得到隐藏状态 ,它再经过一次非线性变换得到
-
g():激活函数,一般取双曲正切 tanh
-
:专门针对词向量的矩阵
-
:专门针对隐藏状态的矩阵
-
:专门针对输出结果的矩阵
-
:偏置项 bias
上图流程可以简化为下面两种等效表达方式:


# RNN 的问题:梯度消失/爆炸

反向传播时,梯度需要从后往前传:
∂L∂W=∂L∂sn×∂sn∂sn−1×∂sn−1∂sn−2×...×∂s1∂W\frac{\partial L}{\partial W}=\frac{\partial L}{\partial s_n}\times\frac{\partial s_n}{\partial s_{n-1}}\times\frac{\partial s_{n-1}}{\partial s_{n-2}}\times ... \times\frac{\partial s_1}{\partial W}
- 如果 (如 0.5),连乘 100 次 →
0.5^100 ≈ 0→ 梯度消失 → 前面的权重几乎不更新 → 模型记不住长距离信息 - 如果 (如 1.5),连乘 100 次 →
1.5^100 ≈ 亿→ 梯度爆炸 → 权重变成NaN→ 模型训练崩溃
解决方案:
原始的 RNN 几乎已被淘汰
方法
说明
梯度裁剪
限制梯度最大值(如 1.0)
ReLU 激活
缓解梯度消失(但 RNN 常用 tanh)
LSTM/GRU
最有效的方案 ⭐ LSTM 和 GRU 中增加了一种名为“门”的结构
Gated RNN
在 RNN 的学习中,
梯度消失也是一个大问题。为了解决这个问题,需要从根本上改变 RNN 层的结构。人们已经提出了诸多 Gated RNN 框架,其中具有代表性的有 LSTM 和GRU
# PyTorch 代码示例
基础 RNN:
import torch
import torch.nn as nn
class SimpleRNN(nn.Module):
def __init__(self, input_size, hidden_size, num_classes):
super().__init__()
self.rnn = nn.RNN(input_size, hidden_size, batch_first=True)
self.fc = nn.Linear(hidden_size, num_classes)
def forward(self, x):
# x: (batch, seq\_len, input\_size)
out, hidden = self.rnn(x) # out: (batch, seq\_len, hidden)
# 取最后一个时间步
out = self.fc(hidden\[-1\])
return out
# 使用
model = SimpleRNN(input_size=100, hidden_size=128, num_classes=2)
用上面的 model,结合 IMDB 数据集进行情感分析:
# 超参数
batch_size = 64
seq_len = 100
learning_rate = 0.001
# 损失和优化器
criterion = nn.CrossEntropyLoss()
optimizer = torch.optim.Adam(model.parameters(), lr=learning_rate)
# 训练
for epoch in range(10):
model.train()
for batch_x, batch_y in train_loader:
optimizer.zero_grad()
output = model(batch_x)
loss = criterion(output, batch_y)
loss.backward()
# 梯度裁剪(防止 RNN 梯度爆炸)⭐
torch.nn.utils.clip\_grad\_norm\_(model.parameters(), 1.0)
optimizer.step()
# 验证
model.eval()
val\_acc = evaluate(model, val\_loader)
print(f"Epoch {epoch}: Val Acc = {val\_acc:.4f}")