拓冰建站拓冰建站
首页 / 资讯中心 / 正文

从零手写RNN到LSTM/GRU:循环神经网络代码实现与梯度消失实战

这一节终于要动真格的了。前面几讲咱们把深度学习的基础概念、反向传播、全连接网络和卷积网络都过了一遍现在轮到循环神经网络RNN的代码实现了。很多朋友对RNN的理解停留在“能处理序列数据”这句话上真正打开编辑器写的时候会发现一堆问题维度怎么排、隐藏状态怎么传、LSTM和GRU到底改了什么、训练时为什么loss死活不降。这讲就是把这些问题一个一个解决掉。这次的代码实现走的是“从零手写再对比验证”的路线。我会先手写一个标准循环神经网络Vanilla RNN把核心公式落到代码里然后在这个基础上扩展到LSTM和GRU最后给出训练过程中的踩坑实录。适合已经会Python和PyTorch基础操作、但对RNN内部机制还比较模糊的读者。看完之后你不仅能跑通代码还能在面试或者做项目时把你写的网络结构讲清楚。1. 复杂循环神经网络到底在解决什么问题1.1 从标准RNN的核心公式聊起标准循环神经网络Vanilla RNN是整个序列模型的基石。它的核心思想用一句话概括每一时刻的输出都依赖当前的输入和上一时刻的隐藏状态。这个机制做了一件在普通全连接网络里做不到的事——让网络拥有“记忆”。就像一个人逐字阅读文本时读到第10个字的时候脑子里还保留着前9个字的信息从而能理解整个句子的语境。RNN的核心公式极其简洁只有两行隐藏状态更新$h_t \tanh(W_{xh} x_t W_{hh} h_{t-1} b_h)$输出计算$y_t W_{hy} h_t b_y$其中 $x_t$ 是t时刻的输入$h_t$ 是t时刻的隐藏状态$y_t$ 是t时刻的输出。$W_{xh}$ 是输入到隐藏层的权重$W_{hh}$ 是隐藏层到隐藏层的循环权重——这个权重就是“记忆”的来源。它把上一时刻的状态 $h_{t-1}$ 跟当前输入 $x_t$ 融合在一起。我们这讲标题里说的“复杂循环神经网络代码实现”其实包含两层含义一是从零写出上面这个标准RNN的过程二是在标准RNN基础上扩展到LSTM、GRU、双向结构、多层堆叠等更复杂的形式。很多人学RNN直接上LSTM结果连标准RNN都写不熟出了问题根本不知道是哪一层在作怪。所以我的建议是先把标准RNN手写一遍理解清楚了再往上加结构。举个例子帮助理解循环机制假设网络在预测句子“我今天吃了苹果”的下一个词当处理到“苹果”时隐藏状态里已经编码了“我”“今天”“吃了”等信息网络就是靠这个状态来决定下一个词最可能是“。”还是“吗”。把整个时间步展开算一遍输入序列有多少个词网络就有多少个时间步每一步都共享同一套权重参数。1.2 为什么标准RNN不够用梯度消失的真相既然标准RNN已经有了记忆能力为什么还要折腾LSTM和GRU核心原因就是梯度消失。而且这个问题在RNN里比在深层前馈网络里严重得多。原因从数学上看非常直接。反向传播时损失对 $h_0$ 的梯度需要沿着时间步一路乘回去这是一个连乘的过程。每乘一步都会被激活函数的导数“缩小”一次。tanh的导数最大值是1sigmoid的导数最大值是0.25。如果网络有50个时间步哪怕每一步都取最大导数0.25乘50次之后就是 $0.25^{50}$约等于 $7.9 \times 10^{-31}$这已经不是一个“小数字”而是彻底的数值归零。梯度消失带来的表现非常典型网络只能学到短距离的依赖关系距离超过5到10个时间步的信息基本就传不过去了。比如“我在北京长大后来去上海工作现在很想念__”这种句子要预测“__”需要记住很远之前出现的“北京”。如果梯度传不回来网络就只能靠局部语境瞎猜。那我们怎么解决两条路一是使用LSTM、GRU这类带门控机制的结构让信息可以通过“加法路径”跨时间步传播避免连乘带来的指数衰减二是使用梯度裁剪、更好的参数初始化、ReLU激活函数等技巧来缓解。这讲后面的内容都会覆盖到。2. 五步到位从零实现标准RNN的完整过程2.1 第一步构造时间序列数据动手写网络之前先得有数据。我们用一个经典任务来验证RNN的正确性——正弦波预测。给定过去若干个时间步的取值预测下一个时间步的值。这个任务足够简单能快速验证网络结构是否正确同时又能直观看到“记忆”的作用。数据生成的代码很直接import numpy as np import torch import torch.nn as nn def make_sin_data(seq_len50, num_samples1000): x np.linspace(0, 20 * np.pi, num_samples) data np.sin(x) 0.1 * np.random.randn(num_samples) # 用滑窗构造样本每 seq_len 个连续点作为输入预测下一个点 xs, ys [], [] for i in range(len(data) - seq_len): xs.append(data[i:i seq_len]) ys.append(data[i seq_len]) xs np.array(xs, dtypenp.float32).reshape(-1, seq_len, 1) ys np.array(ys, dtypenp.float32).reshape(-1, 1) return torch.from_numpy(xs), torch.from_numpy(ys) X, y make_sin_data() print(X.shape, y.shape) # torch.Size([950, 50, 1]) torch.Size([950, 1])样本的维度是[batch_size, seq_len, input_size]。这个顺序是整个RNN实现里最容易搞混的地方我见过不止一个新手把维度排成[seq_len, batch, input]传给网络之后各种报错。PyTorch官方的nn.RNN默认接受[seq_len, batch, input]的格式但我们手写实现的时候用[batch, seq_len, input]这样在循环里取x[:, t, :]更直观后面算损失也不容易出错。还有一个细节值得注意滑窗构造样本时相邻样本之间有49个点是重叠的。这会造成训练样本之间高度相关模型可能对噪声过拟合。不过对于演示性质的正弦预测任务来说问题不大如果做更严肃的时间序列预测可以考虑间隔采样或者使用更多样的数据来源。2.2 第二步手写一个RNN单元标准RNN的核心就是一个全连接层加一个tanh激活只不过这个全连接层同时接收当前输入和上一时刻的隐藏状态。先写一个最底层的Cellclass VanillaRNNCell(nn.Module): def __init__(self, input_size, hidden_size): super().__init__() # 直接声明 Parameter而不是用 nn.Linear好处是权重形状看得一清二楚 self.W_xh nn.Parameter(torch.randn(input_size, hidden_size) * 0.01) self.W_hh nn.Parameter(torch.randn(hidden_size, hidden_size) * 0.01) self.b_h nn.Parameter(torch.zeros(hidden_size)) def forward(self, x_t, h_prev): # x_t: [batch, input_size] # h_prev: [batch, hidden_size] h_t torch.tanh(x_t self.W_xh h_prev self.W_hh self.b_h) return h_t这里有个问题需要解释为什么初始化权重时要乘以0.01因为如果权重初始值太大tanh的输入会落入饱和区导数接近0模型一开始就处于梯度消失状态。乘以0.01可以让初始输出集中在0附近tanh在这个区域导数最大梯度能顺利回传。这是一个简单但极其有效的初始化技巧。为什么这里用tanh而不是ReLU如果你的网络层数不深用ReLU在数学上没问题但tanh输出范围是(-1, 1)有中心对称性在RNN里更容易保持隐藏状态的数值稳定。如果要换成ReLU就必须配合梯度裁剪和精细的学习率调整否则炸的概率很高。2.3 第三步把前向传播完整跑通有了Cell下一步就是把整个时间序列循环起来class VanillaRNN(nn.Module): def __init__(self, input_size, hidden_size): super().__init__() self.cell VanillaRNNCell(input_size, hidden_size) self.hidden_size hidden_size def forward(self, x): # x: [batch, seq_len, input_size] batch, seq_len, _ x.shape h torch.zeros(batch, self.hidden_size, devicex.device) hs [] for t in range(seq_len): h self.cell(x[:, t, :], h) hs.append(h) return torch.stack(hs, dim1)循环里的逻辑非常清晰初始化隐藏状态为全零向量然后从t0到tseq_len-1每一步把当前输入和上一步隐藏状态喂进Cell得到的输出作为下一步的隐藏状态。hs列表把每一步的隐藏状态都存了下来最后用torch.stack拼成[batch, seq_len, hidden_size]的张量。这里有一个设计取舍为什么返回所有时间步的隐藏状态而不是只返回最后一个因为在实际任务中两种需求都存在。文本分类通常只需要最后一个时间步的状态因为在处理完整个序列之后最后的状态浓缩了所有信息而序列标注比如词性标注或seq2seq的编码器则需要每个时间步的输出。返回全部状态灵活性更高用的时候想取哪个就取哪个。2.4 第四步加训练逻辑和梯度裁剪网络结构搭好了接下来是最容易出错的部分——训练循环def train_rnn(model, X, y, epochs200, lr0.01): criterion nn.MSELoss() optimizer torch.optim.Adam(model.parameters(), lrlr) dataset torch.utils.data.TensorDataset(X, y) loader torch.utils.data.DataLoader(dataset, batch_size32, shuffleTrue) for epoch in range(epochs): total_loss 0 for xb, yb in loader: optimizer.zero_grad() output model(xb) # [batch, seq_len, hidden_size] pred output[:, -1, :] # 取最后一个时间步 # 为了简化这里让 hidden_size 1 loss criterion(pred, yb) loss.backward() nn.utils.clip_grad_norm_(model.parameters(), max_norm5.0) optimizer.step() total_loss loss.item() if (epoch 1) % 50 0: print(fEpoch {epoch1}, Loss: {total_loss / len(loader):.6f})这里的nn.utils.clip_grad_norm_就是梯度裁剪。它的作用是把所有参数的梯度按L2范数缩放如果范数超过设定的阈值就成比例缩小。为什么必须加因为RNN的反向传播是沿时间展开的梯度除了可能消失也可能爆炸——尤其是任务里存在长序列时。反向传播的连乘过程里只要权重矩阵的最大奇异值大于1梯度就会指数级增长几个epoch之后loss就会变成NaN。梯度裁剪不能解决网络结构带来的梯度消失但能非常有效地防止训练过程被梯度爆炸打断。输出层这里做了简化直接让hidden_size设为1这样隐藏状态本身就是预测值。实际项目中通常会接一个线性层把hidden_size映射到输出维度这个在后面LSTM部分我会展开说。2.5 第五步与PyTorch官方实现对齐验证自己手写实现最大的风险就是“写完了但不知道自己写得对不对”。所以最后一步用PyTorch官方的nn.RNN做对比验证torch.manual_seed(42) model_custom VanillaRNN(input_size1, hidden_size1) model_official nn.RNN(input_size1, hidden_size1, batch_firstTrue) # 把官方模型的权重同步过来 with torch.no_grad(): model_official.weight_ih_l0.copy_(model_custom.W_xh.T) model_official.weight_hh_l0.copy_(model_custom.W_hh.T) model_official.bias_ih_l0.copy_(model_custom.b_h) model_official.bias_hh_l0.copy_(torch.zeros_like(model_custom.b_h)) x_test torch.randn(4, 50, 1) with torch.no_grad(): out_custom model_custom(x_test) out_official, _ model_official(x_test) print(最大误差:, (out_custom - out_official).abs().max().item())这一步非常关键。如果你手写的实现有问题在这里就能立刻发现。我实际测试时最大误差通常在1e-6量级这说明手写逻辑和官方实现完全一致。为什么不是完全为0因为浮点数运算的顺序不同会带来极小的数值误差这是正常的。这里顺带说明一下为什么官方实现的参数命名里有weight_ih和weight_hh。ih表示input-to-hiddenhh表示hidden-to-hidden对应我们手写代码里的W_xh和W_hh。搞懂官方命名规则调试的时候查文档会快很多。3. 从标准RNN到复杂RNNLSTM和GRU的实现要点3.1 LSTM的门控机制代码解读标准RNN的梯度消失问题促使研究者提出了长短期记忆网络LSTM。LSTM的核心改动是引入了一个独立的记忆单元 $c_t$以及三个门控遗忘门、输入门、输出门。这三个门都用sigmoid激活输出范围是(0,1)用来控制信息保留的比例。关键的代码实现如下class LSTMCell(nn.Module): def __init__(self, input_size, hidden_size): super().__init__() self.input_size input_size self.hidden_size hidden_size # 一次性初始化4组权重顺序对应 [i, f, g, o] self.W_ih nn.Parameter(torch.randn(input_size, 4 * hidden_size) * 0.01) self.W_hh nn.Parameter(torch.randn(hidden_size, 4 * hidden_size) * 0.01) self.b_ih nn.Parameter(torch.zeros(4 * hidden_size)) self.b_hh nn.Parameter(torch.zeros(4 * hidden_size)) def forward(self, x_t, h_prev, c_prev): gates (x_t self.W_ih self.b_ih) (h_prev self.W_hh self.b_hh) i, f, g, o gates.chunk(4, dim-1) i torch.sigmoid(i) # 输入门 f torch.sigmoid(f) # 遗忘门 g torch.tanh(g) # 候选记忆 o torch.sigmoid(o) # 输出门 c_t f * c_prev i * g h_t o * torch.tanh(c_t) return h_t, c_t讲一下四个分量的作用。输入门i决定当前输入有多少信息写入记忆单元遗忘门f决定上一时刻的记忆要保留多少g是候选记忆提供真正要写入的新内容输出门o决定最终输出多少记忆内容给隐藏状态。对比标准RNN可以发现LSTM的梯度传播路径发生了质的变化。记忆单元 $c_t$ 的更新公式是 $c_t f \cdot c_{t-1} i \cdot g$这是一条加法路径梯度的反传不再是一连串矩阵乘法的连乘而是通过遗忘门的系数逐时间步传播。遗忘门的输出是0到1之间的数一般初始化偏向1也就是默认保留大部分历史信息。这从根源上避免了梯度指数级消失的问题。实际操作时LSTM的参数数量大约是标准RNN的四倍训练更慢、更容易过拟合所以在大数据集上效果显著但小数据量任务上未必能赢过标准RNN。这个需要根据具体场景权衡。3.2 GRU的简化设计思路GRU门控循环单元是LSTM的一种精简变体把LSTM的三个门简化成两个更新门和重置门。更新门合并了输入门和遗忘门的功能重置门用来控制历史信息的遗忘程度。它只有两个门的参数参数数量约为LSTM的四分之三训练速度更快在小数据集上通常表现更好。GRU的核心更新公式可以简单描述为$z_t \sigma(W_{xz} x_t W_{hz} h_{t-1} b_z)$$r_t \sigma(W_{xr} x_t W_{hr} h_{t-1} b_r)$$\tilde{h}_t \tanh(W_{xh} x_t r_t \odot (W_{hh} h_{t-1}) b_h)$$h_t (1 - z_t) \odot h_{t-1} z_t \odot \tilde{h}_t$从公式可以看到GRU不像LSTM那样维护单独的记忆单元而是把全部信息都存在隐藏状态 $h_t$ 里。更新门 $z_t$ 控制“新的候选状态”和“旧状态”的融合比例当 $z_t$ 接近0时$h_t$ 几乎完全保留旧状态重置门 $r_t$ 控制计算候选状态时参考多少旧状态信息。在实际编码时GRU的代码结构和LSTM高度相似只是gate数量从4个变成了3个或更少有的实现把矩阵合并优化得更紧凑。如果你能手写LSTM改造出GRU只需要半小时。3.3 双向RNN与多层堆叠的实现要点除了门控机制“复杂RNN”还有一个重要方向是结构层面的扩展。最常见的就是双向RNNBiRNN和多层堆叠RNN。单向RNN只能利用当前时刻之前的信息这在很多任务里不够用。比如命名实体识别时判断“苹果”是水果还是公司既需要看它前面的词“吃了一个”也需要看它后面的词“发布了新手机”。双向RNN的思路很简单跑一个正向RNN再跑一个反向RNN把每个时间步的两个隐藏状态拼接起来作为最终输出。在PyTorch里只要给nn.RNN或nn.LSTM加上bidirectionalTrue参数就能实现底层会自动处理反向序列的计算。多层堆叠RNN也很常见就是把多个RNN层串联起来上一层的输出序列作为下一层的输入。这样做的好处是不同层能捕捉不同粒度的特征第一层可能学到词的局部结构第二层学习短语级别的模式更高层学习句子级别的语义。在PyTorch里设置num_layers2或num_layers3就能直接实现。多层RNN的梯度传播路径更长所以LSTM或GRU在多层结构里几乎是标配。4. 实操过程中我踩过的坑与排查实录4.1 预测结果是一条直线的元凶用标准RNN做正弦波预测时我第一次跑完训练打开预测图看结果模型输出的是一条水平直线完全没有任何波动。loss虽然在下滑但下降得非常慢最后停在一个平台期不再变化。这个现象很典型原因大概率是梯度消失。模型学到的最优策略就是“输出训练集标签的均值”因为这样MSE损失是次优的比随机猜测好。由于梯度传不回去网络根本没有能力学习到输入序列前后的依赖关系。我的解决办法是先用LSTM复现同一任务确认loss能大幅下降然后再回头优化标准RNN调整学习率、初始化方式、减少序列长度。想用标准RNN做长序列任务必须清清楚楚看到它的天花板在哪。4.2 loss突然变成NaN怎么办训练过程中loss在正常下降某一步突然变成NaN这是另一个高频问题。我用正弦波数据测试时如果学习率设成0.1Adam优化器也会翻车。梯度爆炸在RNN训练里太常见了尤其是网络层数较多或者序列较长时。解决手段按优先级排列第一加梯度裁剪这能解决90%的NaN问题第二降低学习率第三检查数据里有没有NaN或无穷大值第四把参数初始化改成正交初始化——RNN里这个技巧比普通随机初始化更有效。Python里可以用torch.nn.init.orthogonal_对隐藏层权重做正交初始化官方在很多RNN模型里也推荐这种方式。4.3 手写实现与官方实现对不上的常见原因很多人在我给的第五步对比验证时发现手写实现和nn.RNN输出差异巨大。我排查下来最常见的有三个原因。第一个是权重同步时维度转置错了。官方nn.RNN的weight_ih_l0形状是[hidden_size, input_size]而我们手写的W_xh是[input_size, hidden_size]同步时必须转置。第二个是忘记batch_firstTrue参数。默认情况下nn.RNN期望输入是[seq_len, batch, input]如果你按[batch, seq_len, input]喂进去模型不会报错但会把batch_size当成seq_len来循环输出结果完全错误。第三个是bias的处理官方模型的bias是分成bias_ih和bias_hh两份而我们手写时很可能合并到一个b_h里同步时需要都赋值。4.4 常见问题速查表问题现象可能原因排查方法解决方案loss不下降学习率过低、梯度消失观察loss数值变化曲线调大lr、改用LSTM/GRU预测为均值直线梯度消失检查不同时间步的梯度范数用门控结构、缩短序列长度loss突然变NaN梯度爆炸查看loss变化拐点梯度裁剪、降低lr、正交初始化官方/手写实现结果对不上维度顺序、初始化不一致逐层对比输出张量形状转置同步权重、加batch_first训练慢参数过多、收敛缓慢监控时间步循环耗时用GRU、减少hidden_size5. 模型效果验证与后续扩展5.1 用序列预测任务验证实现的正确性完成上述所有代码后最终结果应该是标准RNN在正弦预测任务上能学到基本趋势但预测曲线有明显滞后LSTM和GRU的预测曲线与真实曲线高度重合尤其在有噪声的数据上GRU的训练速度和收敛稳定性通常比LSTM更好。我在实践中用同样的数据和超参数跑过对比GRU在50个epoch时就能达到LSTM在100个epoch的loss水平。这个实验本身不大但验证价值很高。它能同时验证你网络的正确性、训练循环的正确性以及框架API的使用是否正确。接下来你可以把任务换成语料库上的文本生成输入前几个单词预测下一个词或者换成情感分类方法都完全一样只是换数据和输出层。5.2 从循环神经网络到注意力机制的扩展路线最后说说后续还能往哪个方向走。RNN在实际工业场景里的地位正在被Transformer架构大幅替代但RNN的训练思路、序列建模思想、梯度传播机制是所有序列模型的基础。掌握了从零实现RNN的能力再去理解Transformer里的自注意力机制会顺畅很多。比如你现在能理解“为什么Transformer可以并行处理序列”因为它的计算不像RNN这样强依赖上一个时间步的输出。如果继续深入建议按这个路线走先看注意力机制如何解决RNN长距离依赖的瓶颈然后看Transformer的整体结构最后了解Transformer在NLP和CV里的应用。这中间会遇到很多新概念但有RNN代码实现的底子在你至少不会被“位置编码”“多头注意力”这些术语吓住。还有一个小技巧想分享手写实现一遍 lassLSTM 和 GRU 之后再看官方文档里的公式说明你会发现每一行都无比清晰再遇到模型效果不好时你也能从网络内部结构的角度分析问题而不是只会在网上搜调参攻略。这种“能洞察网络内部到底在算什么”的能力才是从入门到进阶的分水岭。
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门