ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

从零手写RNN:理解循环神经网络、梯度消失与PyTorch实现

从零手写RNN:理解循环神经网络、梯度消失与PyTorch实现 很多人学循环神经网络的时候第一反应都是直接上LSTM觉得RNN太弱、太基础。我之前也这么想直到在一个时间序列预测项目里被一个简单的RNN狠狠上了一课。当时我用全连接网络去预测一段带随机相位的正弦波结果模型输出几乎是一条水平线怎么调都学不到序列的变化规律。后来换成循环神经网络几十个epoch就能拟合得很漂亮。那一刻我才意识到不是全连接不够深而是它的结构里压根没有记忆这个概念。这篇文章我就把这套东西从头到尾讲透包括RNN到底在算什么、BPTT为什么容易梯度消失、怎么用PyTorch从零手写一个可训练的RNN以及我在训练过程中踩过的那些坑。1. 为什么全连接和CNN搞不定序列数据1.1 序列数据的本质顺序本身就是信息先说一个最根本的问题什么叫序列数据语音是一帧一帧按时间排出来的文本是一个字一个字按顺序写出来的股票价格是每天收盘后按日期串起来的。这类数据有一个共同特征当前时刻的语义依赖前面若干时刻的信息。你听到我喜欢你和你我喜欢用的是同样的三个字但含义完全不同因为顺序变了。全连接网络处理输入时把所有特征拼成一个固定长度的向量每个位置是独立的。模型内部没有跨位置的信息传递机制。CNN稍微好一点卷积核可以覆盖一个局部窗口但窗口之外的远距离依赖依然无能为力。你让CNN去预测一段正弦波的未来走势它只能看到窗口内的几十个点对于这个波形的相位从哪开始、振幅多大这种需要在整个时间轴上建立认知的问题局部窗口是不够的。RNN的解决方案非常直白给网络增加一个隐藏状态hidden state这个状态在每个时间步都会被更新更新时既看当前输入也看上一个时刻的状态。相当于网络带了一个小本子每走一步都往本子上记点东西下个时刻做决策时先翻一翻本子。这个小本子就是RNN的记忆。1.2 参数共享是RNN最核心的思维方式RNN和全连接网络还有一个本质区别全连接网络每一层有自己独立的权重层数越多参数越多而RNN在所有时间步上共享同一套参数。也就是说不管序列是10步还是1000步用来处理每一步的Wxh和Whh都是同一份。参数共享的现实意义很大。首先它极大减少了参数量一个隐层32维的RNN核心参数只有32x3232x1这一百多个数字就算序列再长也不会增加。其次它强迫模型学到每一步都适用的通用变换规则——不管你看到的是序列的第3个点还是第80个点处理逻辑是一致的。这非常符合序列数据的特性规律是稳定的变的只是内容。我当时写了一个最简RNN的forward代码只有几行但每一步都在体现这个思想for t in range(seq_len): h torch.tanh(x[:, t:t1] Wxh.T h Whh.T bh)不同时刻的x_t进来了但Wxh和Whh始终是同一个h被一遍一遍地更新然后传给下一个时刻。1.3 一个直觉例子词性标注为了把这个记忆机制讲得更具象我拿词性标注举例。假设有句话小明用手机看电影网络逐个词读入。读到用的时候隐藏状态里已经编码了小明这个主语以及动词的感觉读到手机时状态里同时有用手机的关联所以这个手机更容易被识别为工具宾语再读到看的时候整个状态携带了前文的完整语义环境判断电影是看的宾语就非常自然了。每一步的判断都不只是靠当前词本身的嵌入向量而是靠当前词整个历史状态的压缩摘要。这种摘要当然会损失细节但常用信息会被保留下来。这也是为什么RNN后来会被注意力机制部分替代的重要原因——记忆容量是有限的但那是后话。2. RNN的前向传播从公式到数值计算2.1 隐藏状态是网络的记忆RNN前向传播的标准公式就是下面这两行h_t tanh(W_xh * x_t W_hh * h_{t-1} b_h) y_t W_hy * h_t b_yx_t是t时刻的输入h_t是t时刻的隐藏状态y_t是该时刻的输出。很多文章把h_t类比成记忆但这个类比不够精确。更准确地说h_t是网络从序列起点到t时刻全部信息的有损压缩编码。因为tanh的非线性压缩h_t的每个维度通常表示某种特征的强度维度数量决定了信息容量上限。第一行公式里有两个矩阵W_xh负责把当前输入投影到状态空间W_hh负责把上一时刻的状态搬运到当前时刻。这个搬运过程就是记忆在时间维度上的传播。注意如果忽略tanhh_t就是x_t和h_{t-1}的线性组合tanh则加了非线性让状态更新不再是简单叠加而是能模拟更复杂的依赖关系。第二行公式把隐藏状态映射成输出分类任务后面通常还会接softmax回归任务直接就用y_t。2.2 一个具体数字例子手算前向光看公式还是抽象我手算一个超小例子。假设隐藏层维度是2输入维度是1初始状态h0 [0, 0]W_xh [[1.0], [-1.0]]两行一列W_hh [[0.5, -0.3], [0.2, 0.4]]b_h [0, 0]。输入序列是x [2.0, 1.0]。t1时h1 tanh(W_xh * 2 W_hh * h0) tanh([[2.0], [-2.0]] [0, 0]) tanh([2.0, -2.0]) [0.964, -0.964]t2时h2 tanh(W_xh * 1 W_hh * h1) tanh([1.0, -1.0] W_hh * [0.964, -0.964])先算矩阵乘法0.5*0.964 (-0.3)*(-0.964) 0.482 0.289 0.771 0.2*0.964 0.4*(-0.964) 0.193 - 0.386 -0.193所以h2 tanh([1.0 0.771, -1.0 - 0.193]) tanh([1.771, -1.193]) [0.943, -0.831]看到没有h2同时受x2和h1的影响。就算现在x2很小h1里携带的x1信息依然在起作用。如果把序列拉长到100步这种影响会一路传递下去但如果中间经过的tanh导数太小影响也会逐级衰减这就是后面要说的梯度消失。2.3 激活函数的作用与选择RNN里最常用的激活函数是tanh其次才是ReLU。为什么不用ReLU当默认因为ReLU在正区间导数为常数1在循环结构中容易让隐藏状态的值不断累积放大导致训练不稳定。tanh的输出被限制在[-1, 1]之间每步更新后状态不会无限膨胀这一点在长序列上非常重要。但tanh的问题也很明显它的导数最大也只有1而且只有在输入接近0时才接近1输入稍微大一点导数就迅速衰减到接近0。这意味着误差信号经过一个时间步的传播最多保持不变通常会被压缩到原来的零点几倍。多传几步梯度就趋近于0。理解了这一点后面BPTT时梯度消失的推导就是顺理成章的事。3. BPTT反向传播梯度消失不是玄学是数学3.1 BPTT的计算图展开RNN的反向传播叫BPTTBackpropagation Through Time全称是随时间反向传播。名字听着玄本质就是把RNN在时间维度上展开成一个深层的全连接网络然后用标准的链式法则求梯度。比如一个长度为T的序列把前向传播展开就是T个隐层堆叠起来的网络。第t层的输入是x_t和h_{t-1}输出h_t又作为第t1层的输入。唯一的特殊之处在于这T层共享同一套权重W_xh和W_hh。所以误差对W_hh的梯度是所有时间步贡献的总和。PyTorch里你不需要手动实现BPTTloss.backward()会自动把时间维度的计算图展开并求梯度。但如果你不理解梯度是怎么一路传回去的遇到loss不收敛或者NaN就无从下手。3.2 梯度连乘与消失/爆炸的数学从t时刻到t-k时刻传播的梯度链式法则里会出现一组连乘项∂h_t / ∂h_{t-k} ∏(从it-k1到t) diag(f(h_i)) * W_hh其中f是tanh的导数是一个对角矩阵W_hh是隐藏层间的权重矩阵。整个连乘项的范数大概受 |λ_max(W_hh)| 的k次方控制λ_max是W_hh的最大奇异值。这里有两个隐藏的杀手如果λ_max(W_hh) 1k越大连乘项越小梯度呈指数级衰减最终消失。如果λ_max(W_hh) 1k越大连乘项越大梯度呈指数级增长最终爆炸。为什么都说RNN难训练就是因为这个连乘项。即使W_hh的谱半径接近但不超过1只要序列稍长梯度还是会衰减到几乎为0。梯度消失导致的结果是网络无法学习长距离依赖序列前部的信息对后部的预测完全没有贡献。梯度爆炸则更直接更新一步权重就直接变成NaN。从实践角度看梯度爆炸相对好解决——梯度裁剪梯度消失则棘手得多要么换结构LSTM/GRU要么做残差连接要么用更好的初始化。3.3 缓解手段的适用范围几类常见做法我按实际效果排个序梯度裁剪gradient clipping解决爆炸最直接设一个阈值比如5.0梯度范数超过就整体缩放。权重初始化把W_hh初始化为单位矩阵附近是很多RNN任务里让训练稳定的关键。换门控结构LSTM和GRU通过门控机制让梯度有一个高速公路这是解决消失最彻底的办法。双向RNN解决的是未来信息看不到的问题和梯度消失无关可别混为一谈。我在实际项目中检验过大多数时候挂梯度裁剪后loss就不再乱跳了但想真正学到长序列依赖还是得靠门控。4. 手写RNN并用PyTorch训练正弦波预测4.1 为什么选正弦波预测当入门任务不少教程喜欢拿文本生成当RNN例子但对新手来说有几个坎文本预处理复杂、词表很大、训练动辄几十分钟。正弦波预测是干净得多的选择数据是自己生成的多少都行标签是连续的Loss用MSE训练结果直接画图就能看出好坏。更重要的是正弦波带有明确的相位和频率属性模型想知道下一个点是什么必须对过去整段波形有全局感知能清楚反映出RNN的序列建模能力。数据生成我设置了随机相位、随机振幅、随机频率这样模型必须真的学会记住并外推波形规律而不是简简单单背下某个固定形状。训练集生成2000条样本每条取50个连续点前40点当输入后10点当要预测的目标。import numpy as np import torch import torch.nn as nn import torch.optim as optim def generate_sine_data(num_samples2000, seq_len40, pred_len10): xs, ys [], [] for _ in range(num_samples): phase np.random.uniform(0, 2 * np.pi) amp np.random.uniform(0.5, 1.5) freq np.random.uniform(0.8, 1.2) start np.random.uniform(0, 10) t np.linspace(start, start (seq_len pred_len) * 0.1, seq_len pred_len) wave amp * np.sin(freq * t phase) xs.append(wave[:seq_len]) ys.append(wave[seq_len:]) return ( np.array(xs, dtypenp.float32).reshape(-1, seq_len, 1), np.array(ys, dtypenp.float32).reshape(-1, pred_len) ) x_data, y_data generate_sine_data() x_train, y_train x_data[:1600], y_data[:1600] x_val, y_val x_data[1600:], y_data[1600:]为什么输入要reshape成[batch, seq_len, 1]而不是[batch, seq_len]因为RNN在时间步上处理的是特征向量哪怕每个特征只有1维也要保持三维结构方便后面在处理每个时间步时取x[:, t, :]。4.2 手写RNNCell而不是直接调nn.RNN我故意不用nn.RNN而是自己写一个循环体目的就是让你看清每一步在干什么。nn.RNN封装得太干净了初学者很容易把它当黑盒出了问题完全不知道从哪排查。核心模型代码class SimpleRNN(nn.Module): def __init__(self, input_size, hidden_size, pred_len): super().__init__() self.hidden_size hidden_size # 输入到隐状态 self.Wxh nn.Linear(input_size, hidden_size) # 隐状态到隐状态 self.Whh nn.Linear(hidden_size, hidden_size) # 最后的输出层 self.fc nn.Linear(hidden_size, pred_len) def forward(self, x): # x: [batch, seq_len, input_size] batch_size x.size(0) h torch.zeros(batch_size, self.hidden_size, devicex.device) seq_len x.size(1) for t in range(seq_len): x_t x[:, t, :] # 当前步输入 h torch.tanh(self.Wxh(x_t) self.Whh(h)) out self.fc(h) return out这个循环的本质就是第2章那个公式的代码化每来一个新的x_t先和h通过两个Linear变换组合在一起再经过tanh得到新的h。循环结束后我们把最终的h当作整段序列的压缩摘要丢给全连接层去预测未来10个点。有些同学可能会问为什么每个时间步不直接输出非要取最后一个h因为这个任务的设定是看完40个点预测未来10个点这是sequence-to-one模式只需要最后的汇总状态。如果是做逐词预测那就是sequence-to-sequence模式每个时间步都得输出。理解这个区别你就能根据任务改结构而不只是抄代码。4.3 训练循环和关键参数模型定义好之后训练循环看着和普通全连接网络几乎一样但有一个关键动作梯度裁剪。加上这一行之后整个训练过程会稳非常多。model SimpleRNN(input_size1, hidden_size32, pred_len10) criterion nn.MSELoss() optimizer optim.Adam(model.parameters(), lr0.005) batch_size 128 epochs 300 train_dataset torch.utils.data.TensorDataset( torch.from_numpy(x_train), torch.from_numpy(y_train)) train_loader torch.utils.data.DataLoader(train_dataset, batch_sizebatch_size, shuffleTrue) for epoch in range(epochs): model.train() total_loss 0.0 for batch_x, batch_y in train_loader: optimizer.zero_grad() pred model(batch_x) loss criterion(pred, batch_y) loss.backward() # 关键防止梯度爆炸 nn.utils.clip_grad_norm_(model.parameters(), max_norm5.0) optimizer.step() total_loss loss.item() * batch_x.size(0) if (epoch 1) % 50 0: avg_loss total_loss / len(train_dataset) print(fepoch {epoch 1}, loss {avg_loss:.6f})我实测下来用Adam优化器、学习率0.005、隐层32维的经验值搭配训练到第50轮loss就在0.05以下第300轮能到0.002左右。这里有个细节sin值范围在[-1.5, 1.5]之间MSE降到0.002意味着平均误差约0.045预测曲线肉眼看基本和真实值重合。4.4 用预测结果反推模型是否学会Loss数值毕竟抽象我习惯把预测结果画出来看。用验证集里随机抽一条样本把前40个点作为输入把预测的10个点和真实的10个点叠加在一起对比。import matplotlib.pyplot as plt model.eval() with torch.no_grad(): idx 0 sample_x torch.from_numpy(x_val[idx]).unsqueeze(0) sample_y_true y_val[idx] sample_y_pred model(sample_x).squeeze().numpy() plt.figure(figsize(10, 4)) plt.plot(range(50), np.concatenate([x_val[idx].flatten(), sample_y_true]), labeltrue wave, linewidth2) plt.plot(range(40, 50), sample_y_pred, markero, linestyle--, labelpredicted, linewidth2) plt.legend() plt.show()如果预测点能顺着真实波形的趋势平滑延伸说明模型确实把相位和频率都学到了如果预测出来的是一段直线或者朝着错误方向走那基本可以判定模型没有真正利用历史信息。从我多次实验的经验看隐层小于16时预测的后段容易出现向右偏移这是因为信息容量不够模型只记住了大致的周期记不住精确相位隐层32以上就稳定很多这也说明隐层规模对这个任务是有实际影响的。5. 训练RNN时我踩过的坑和调试建议5.1 loss震荡不收敛的几种原因RNN的loss曲线比全连接网络更容易出现震荡很多新手一看到loss上下乱跳就开始怀疑模型写错了。实际上最常见的几种原因按出现频率排学习率太大。RNN的损失曲面非常陡峭学习率0.01在普通网络上可能还好在RNN上就会导致loss剧烈震荡。我用同样的代码lr从0.01降到0.005loss曲线就从疯狗式乱跳变成了稳步下降。没有做梯度裁剪。这是我入坑时犯的最严重的错误。第一次训练RNNloss在第30轮附近突然变成nan找了一晚上原因最后发现是梯度爆炸。加了一行clip_grad_norm_之后问题直接消失。数据没有归一化。如果序列值范围特别大比如几万到几百万tanh的输入会被推到饱和区梯度消失成零模型直接就死掉了。RNN的输入最好归一到[-1, 1]或者零均值单位方差附近这和激活函数的特性强相关。5.2 梯度裁剪一个必须养成的习惯具体说下梯度裁剪的实操。PyTorch里最常用的就是按范数裁剪nn.utils.clip_grad_norm_(model.parameters(), max_norm5.0)一行代码在optimizer.step()之前调用即可。它的原理是计算所有参数的梯度总范数如果超过max_norm就按比例整体缩放。这样能保证更新步长的上限被控制住即便某个样本产生了特别极端的梯度也不会一次把权重打到不可恢复的状态。max_norm选多少我常用的范围是1.0到10.0对大多数RNN任务5.0是个保险的默认值。如果你发现loss需要更激进的下降可以把阈值调大如果训练不稳定就先往小调。注意裁剪要放在backward()之后step()之前顺序别搞反。我之前看有人把裁剪放在zero_grad()之前等于白做因为backward的梯度还没计算出来。5.3 隐状态初始化、序列长度的取舍RNN的h0习惯上初始化为全零这是合理的默认选择因为序列开始时没有任何历史信息。但如果你的任务有明确先验比如预测波形初始相位已知那也可以把h0初始化为对应的编码向量。不过绝大多数情况下别自己乱设全零最稳。序列长度是另一个容易被忽略的超参数。理论上RNN支持任意长度序列但实际训练中序列越长BPTT的连乘路径越长梯度消失越严重训练越难。我的建议是训练时先用较短的序列长度比如20到50个点把模型跑通确认loss能下降再拉长序列做正式训练如果必须处理长序列优先考虑LSTM/GRU或者用截断BPTTTruncated BPTT即把长序列切成多段每段只回传固定步数的梯度。在我做正弦波任务时有一个体会序列长度从40增加到100普通RNN的loss明显变差而GRU几乎不受影响。你可以在自己的实验里对比一下这个对比本身对理解RNN的梯度问题非常有帮助。5.4 一个容易被忽视的坑预测滞后期在做时间序列预测时RNN经常会出现一种看起来拟合很好但实际是滞后预测的情况——预测曲线比真实曲线晚了若干个时间步但两者的形状非常接近。这种现象在金融时序预测里尤其典型原因是模型发现把上一个观测值直接搬过来当预测值能让loss很低成本比预测精确转折点低得多。怎么判断你的模型有没有这种问题把预测值和真实值画在同一张图上看预测曲线是否整体右移。如果滞后明显说明模型没有真正学到序列的动态规律需要调整Loss函数比如加入一阶差分惩罚项或者预测多步时用noise scheduling提高模型对扰动的鲁棒性。6. 从RNN到LSTM/GRU门控机制为什么能救场6.1 LSTM的门控直觉既然RNN的梯度消失问题这么严重那LSTM到底做了什么事来救场一句话它给信息传递加了两条专用通道。一条是候选记忆通道负责写入新信息一条是遗忘门控制的记忆通道决定上一时刻的哪些记忆要被保留、哪些要被丢弃。关键是第二条通道的传递路径非常干净——它是一条从c_{t-1}到c_t的直连线性路径没有经过tanh的非线性压缩。误差信号可以从这条路上高速穿越很多时间步而不衰减这就从根上缓解了梯度消失。用生活化类比来说普通RNN像个每次考试前都要把所有书重新背一遍的学生时间一长前面的内容全忘光了LSTM则像个做笔记的人每隔一段时间审一遍笔记本重要的留下不重要的划掉关键信息不用从头再背一遍。GRU是LSTM的简化版把遗忘门和输入门合并成更新门参数更少在数据量不是特别大的情况下效果往往和LSTM持平训练还更快。就我个人经验如果项目里不要求必须用LSTMGRU是个很好的默认选择。6.2 什么时候应该直接用LSTM/GRU什么时候RNN够用这个问题没有标准答案但你可以根据这三条来判断序列长度序列短小于30步普通RNN完全够用序列长直接上GRU/LSTM。依赖距离如果关键信息离预测位置很远比如文本里前30个词决定最后一个词的时态普通RNN很难学到这种远距离依赖门控结构更靠谱。任务精度要求精度要求高的任务比如语音识别、机器翻译现在的主流方案甚至已经是Transformer了但如果你只是想理解序列建模的基本原理或者做一个快速的序列预测Demo普通RNN的训练周期短、结构简单、易于调试反而更合适。我的一个实用建议是在你自己的项目里先把普通RNN跑通拿到一个baseline loss然后无损替换成GRU把self.Whh换成GRU的cell逻辑观察loss能降多少。这个对比实验做一次你对门控机制的理解会比看十篇教程都深刻。说到底循环神经网络最核心的价值不在于它今天还是不是SOTA而在于它第一次让信息在时间维度上流动这件事变得可训练、可理解。你把这个结构吃透了再去看LSTM、GRU、甚至Transformer里的位置编码都会有一种原来都是老朋友的熟悉感。
返回列表