拓十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
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,核心参数只有32x32+32x1这一百多个数字,就算序列再长也不会增加。其次它强迫模型学到"每一步都适用的通用变换规则"——不管你看到的是序列的第3个点还是第80个点,处理逻辑是一致的。这非常符合序列数据的特性:规律是稳定的,变的只是内容。

我当时写了一个最简RNN的forward,代码只有几行,但每一步都在体现这个思想:

for t in range(seq_len): h = torch.tanh(x[:, t:t+1] @ 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_y

x_t是t时刻的输入,h_t是t时刻的隐藏状态,y_t是该时刻的输出。很多文章把h_t类比成"记忆",但这个类比不够精确。更准确地说,h_t是网络从序列起点到t时刻全部信息的有损压缩编码。因为tanh的非线性压缩,h_t的每个维度通常表示某种特征的强度,维度数量决定了信息容量上限。

第一行公式里有两个矩阵:W_xh负责把当前输入"投影"到状态空间,W_hh负责把上一时刻的状态"搬运"到当前时刻。这个搬运过程就是记忆在时间维度上的传播。注意,如果忽略tanh,h_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]。

t=1时:

h1 = tanh(W_xh * 2 + W_hh * h0) = tanh([[2.0], [-2.0]] + [0, 0]) = tanh([2.0, -2.0]) = [0.964, -0.964]

t=2时:

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的反向传播叫BPTT(Backpropagation Through Time),全称是"随时间反向传播"。名字听着玄,本质就是把RNN在时间维度上展开成一个深层的全连接网络,然后用标准的链式法则求梯度。

比如一个长度为T的序列,把前向传播展开,就是T个"隐层"堆叠起来的网络。第t层的输入是x_t和h_{t-1},输出h_t又作为第t+1层的输入。唯一的特殊之处在于:这T层共享同一套权重W_xh和W_hh。所以误差对W_hh的梯度,是所有时间步贡献的总和。

PyTorch里你不需要手动实现BPTT,loss.backward()会自动把时间维度的计算图展开并求梯度。但如果你不理解梯度是怎么一路传回去的,遇到loss不收敛或者NaN就无从下手。

3.2 梯度连乘与消失/爆炸的数学

从t时刻到t-k时刻传播的梯度,链式法则里会出现一组连乘项:

∂h_t / ∂h_{t-k} = ∏(从i=t-k+1到t) diag(f'(h_i)) * W_hh

其中f'是tanh的导数,是一个对角矩阵;W_hh是隐藏层间的权重矩阵。整个连乘项的范数大概受 |λ_max(W_hh)| 的k次方控制,λ_max是W_hh的最大奇异值。

这里有两个隐藏的杀手:

  1. 如果λ_max(W_hh) < 1,k越大,连乘项越小,梯度呈指数级衰减,最终消失。
  2. 如果λ_max(W_hh) > 1,k越大,连乘项越大,梯度呈指数级增长,最终爆炸。

为什么都说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_samples=2000, seq_len=40, pred_len=10): 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, dtype=np.float32).reshape(-1, seq_len, 1), np.array(ys, dtype=np.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, device=x.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_size=1, hidden_size=32, pred_len=10) criterion = nn.MSELoss() optimizer = optim.Adam(model.parameters(), lr=0.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_size=batch_size, shuffle=True) 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_norm=5.0) optimizer.step() total_loss += loss.item() * batch_x.size(0) if (epoch + 1) % 50 == 0: avg_loss = total_loss / len(train_dataset) print(f"epoch {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]), label="true wave", linewidth=2) plt.plot(range(40, 50), sample_y_pred, marker='o', linestyle='--', label="predicted", linewidth=2) plt.legend() plt.show()

如果预测点能顺着真实波形的趋势平滑延伸,说明模型确实把相位和频率都学到了;如果预测出来的是一段直线或者朝着错误方向走,那基本可以判定模型没有真正利用历史信息。从我多次实验的经验看,隐层小于16时,预测的后段容易出现向右偏移,这是因为信息容量不够,模型只记住了大致的周期,记不住精确相位;隐层32以上就稳定很多,这也说明隐层规模对这个任务是有实际影响的。

5. 训练RNN时我踩过的坑和调试建议

5.1 loss震荡不收敛的几种原因

RNN的loss曲线比全连接网络更容易出现震荡,很多新手一看到loss上下乱跳就开始怀疑模型写错了。实际上最常见的几种原因,按出现频率排:

  • 学习率太大。RNN的损失曲面非常陡峭,学习率0.01在普通网络上可能还好,在RNN上就会导致loss剧烈震荡。我用同样的代码,lr从0.01降到0.005,loss曲线就从"疯狗式乱跳"变成了"稳步下降"。
  • 没有做梯度裁剪。这是我入坑时犯的最严重的错误。第一次训练RNN,loss在第30轮附近突然变成nan,找了一晚上原因,最后发现是梯度爆炸。加了一行clip_grad_norm_之后问题直接消失。
  • 数据没有归一化。如果序列值范围特别大,比如几万到几百万,tanh的输入会被推到饱和区,梯度消失成零,模型直接就死掉了。RNN的输入最好归一到[-1, 1]或者零均值单位方差附近,这和激活函数的特性强相关。

5.2 梯度裁剪:一个必须养成的习惯

具体说下梯度裁剪的实操。PyTorch里最常用的就是按范数裁剪:

nn.utils.clip_grad_norm_(model.parameters(), max_norm=5.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,或者用截断BPTT(Truncated 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持平,训练还更快。就我个人经验,如果项目里不要求必须用LSTM,GRU是个很好的默认选择。

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里的位置编码,都会有一种"原来都是老朋友"的熟悉感。

返回列表