尧图精选

从RNN到LSTM:理解梯度消失与门控机制,用PyTorch实现时间序列预测

🕒 发布时间:2026/9/30 10:02:00 📁 来源:尧图网络
1. 序列建模的本质为什么全连接网络搞不定昨天、今天、明天先把一句话放在前面循环神经网络是为了处理有时间顺序的数据而设计的一类网络结构核心能力在于记住过去影响现在推测未来。如果你接触过全连接网络或者卷积神经网络你应该已经习惯了这样一种模式给模型一个固定长度的输入向量经过若干层变换输出一个固定长度的结果。图像分类就是典型例子一张256×256×3的图片拉平或卷积后变成一个固定维度的特征最后映射到类别概率上。但这种模式面对序列数据时会有非常明显的别扭感。语言、语音、股票价格、传感器读数、视频帧这些数据的长度是变化的上一时刻和下一时刻之间有依赖关系hello world和world hello如果拆成单词集合内容完全一样但顺序不同意义就完全不同。全连接网络看不出这两句话的区别因为它的设计假设就是输入的所有维度之间没有先后关系地位平等。RNN的出现就是为了打破这个假设。它引入了隐状态这个概念本质上是给网络加了一块工作记忆。每处理一个新输入网络不仅看当前的数据还要看自己上一时刻的记忆然后更新这块记忆再输出当前时刻的结果。用公式讲就是[ h_t \text{tanh}(W_{ih}x_t b_{ih} W_{hh}h_{t-1} b_{hh}) ][ y_t W_{ho}h_t b_o ]看明白这个公式你就看懂了RNN的全部家底。它做的事情就是当前时刻的隐状态 ( h_t ) 由两部分决定一部分是当前输入 ( x_t ) 经过线性变换后的结果另一部分是上一时刻隐状态 ( h_{t-1} ) 经过线性变换后的结果两者相加之后放进一个 tanh 激活函数里。然后当前时刻的输出 ( y_t ) 就是对这个隐状态再做一次线性变换。很多人第一次看到这个结构会觉得就这么简单对的RNN的基本思想是真的很简单。但魔鬼藏在细节里这个简单的结构在训练时会引出整个深度学习领域最长寿、最经典、也最让人头疼的问题——梯度消失和梯度爆炸。我见过太多人包括我自己刚学的时候都把注意力放在RNN长什么样上画了很多展开图、时序图以为看懂了。但真正卡住所有人的不是前向传播而是反向传播时梯度通过时间维度的传递过程。下一节我会把它讲透因为不搞懂这个你根本理解不了为什么LSTM要设计成那样一堆门。2. 梯度消失不是玄学RNN训练失败的数学本质2.1 BPTT把展开的网络当成一个超深的网络来反向传播RNN训练用的算法叫沿时间反向传播。思路很简单网络按时间展开之后T个时刻就等于一个T层的前馈网络每一层的权重是共享的。所以就按普通反向传播的思路从最后一个时刻的损失出发逐层把梯度往回传。问题是这个展开的网络是共享权重也就是说同一个权重矩阵 ( W_{hh} ) 被反复乘以T次。梯度在反向传播时每经过一个时间步就要乘一次 ( W_{hh}^T )。如果 ( W_{hh} ) 的最大特征值大于1梯度就会指数级增长变成梯度爆炸如果小于1梯度就会指数级衰减变成梯度消失。用数学语言写出来就是损失对 ( W_{hh} ) 的梯度可以拆成T个项的求和每一项里都带着一个连乘项[ \prod_{jk1}^t \text{diag}(\text{tanh}(h_j)) W_{hh}^T ]你可以把这一串连乘想成接力传话从第一个时刻传到第t个时刻每个时刻都要把上一个人告诉你的信息打一个折扣再传给下一个人。而这个折扣率就是 ( W_{hh} ) 的特征值大小。如果折扣率是0.9传到第20个时刻信息就只剩原来的0.9的20次方约等于0.12。传到第50个时刻就已经逼近0了。2.2 为什么梯度消失对RNN比对普通深层网络更致命在普通的前馈网络里梯度消失最多意味着前面几层学不到东西。在RNN里梯度消失意味着模型学不会长距离依赖。什么叫长距离依赖举个例子句子The cat, which was very hungry, ate the fish里动词ate的单复数形式由主语cat决定而cat和ate之间隔着整整一个从句。模型要正确输出ate而不是eats就必须在读到ate这个位置时还记得7个词之前出现过cat这个信息。如果梯度传到第7步已经衰减到接近0反向传播时这个位置的监督信号就根本传不回去模型自然学不到这种关联。梯度爆炸相对好处理一些——用梯度裁剪就能压住。但梯度消失是结构性的问题靠调参、改学习率、加深层数都解决不了。这也是为什么1991年左右就有人提出RNN的基本框架但直到1997年LSTM出现之后循环网络才真正变得工业可用。2.3 为什么用tanh而不是ReLU当激活函数这是我很喜欢问别人的一个问题。你仔细看RNN的公式隐状态更新用的是tanh不是深度学习里最常用的ReLU。为什么一个关键原因是tanh的输出范围是(-1, 1)每个时刻的隐状态更新完之后都会被压回这个范围里这在一定程度上抑制了隐状态数值发散。如果用ReLU它的输出没有上界经过多个时刻的累加隐状态的数值很容易一路涨到天文数字训练直接飞掉。另外tanh在0附近有比较大的梯度这有利于梯度传播。ReLU在正半轴的梯度恒为1听起来很美好但放到循环结构里反而会让状态无约束地增长。所以RNN的激活函数选择不是随便挑一个流行的而是只有tanh最合适。这背后是稳定性与梯度流动的权衡。你理解了这一层再去看LSTM内部那些门控公式就明白每一步设计都是被问题逼出来的。3. LSTM的门控机制一扇门一扇门拆开看LSTM全称是Long Short-Term Memory。注意这个名字它的目标不是记忆而是长记忆——它要解决的是RNN记不住长距离信息的问题。它的核心思路很朴素给网络加一条传送带叫细胞状态用 ( C_t ) 表示。这条传送带上可以顺着时间方向一格格往前传信息沿路都有门来控制信息能否通过、能通过多少。你去看LSTM的完整公式一共四个子步骤[ f_t \sigma(W_f \cdot [h_{t-1}, x_t] b_f) ][ i_t \sigma(W_i \cdot [h_{t-1}, x_t] b_i) ][ \tilde{C}t \text{tanh}(W_C \cdot [h{t-1}, x_t] b_C) ][ C_t f_t \ast C_{t-1} i_t \ast \tilde{C}_t ][ o_t \sigma(W_o \cdot [h_{t-1}, x_t] b_o) ][ h_t o_t \ast \text{tanh}(C_t) ]这看着吓人但拆开之后非常清晰一共就三个门加一个候选状态。我用普通人听得懂的方式一个个讲。3.1 遗忘门不是遗忘是选择性保留[ f_t \sigma(W_f \cdot [h_{t-1}, x_t] b_f) ]这个门的输出是一个0到1之间的数决定上一时刻的细胞状态 ( C_{t-1} ) 有多少比例能被保留下来。如果 ( f_t ) 接近1就表示信息基本全留;如果接近0就表示基本全扔。为什么叫遗忘门我个人的体会是这个名字其实有点误导。它更准确的角色是保留门或者衰减控制门。名字叫什么不重要重要的是它的行为LSTM通过这个门学会了一个能力——根据当前输入和上一时刻的隐状态决定哪些旧信息已经没用了可以放弃哪些还有价值值得继续留着。举个直觉例子你在读一篇新闻前面讲了A公司发布了新产品中间插了一大段背景介绍到结尾时提到这家公司计划在明年推出下一代产品。模型读到结尾的这家公司时需要靠遗忘门决定背景介绍那些细节可以慢慢淡忘但A公司是主角这个信息必须在细胞状态里留住。3.2 输入门决定新信息写入多少[ i_t \sigma(W_i \cdot [h_{t-1}, x_t] b_i) ][ \tilde{C}t \text{tanh}(W_C \cdot [h{t-1}, x_t] b_C) ]输入门由两部分配合完成。( \tilde{C}t ) 是候选记忆相当于根据当前输入 ( x_t ) 和上一时刻隐藏状态 ( h{t-1} ) 提出了一批我想记住的新内容;( i_t ) 是一个0到1之间的开关决定这批候选内容有多少真的能被写进细胞状态。这两部分相乘就是只写入那些值得记录的新信息。有些资料里把 ( i_t ) 翻译成输入门把 ( \tilde{C}_t ) 翻译成候选记忆我觉得挺准确的。前者负责筛选后者负责提供内容。3.3 状态更新公式整条传送带的核心逻辑[ C_t f_t \ast C_{t-1} i_t \ast \tilde{C}_t ]上面这个公式是整个LSTM里面最重要的一行。它的含义是新的细胞状态 旧的细胞状态 × 遗忘门的保留比例 新写入的内容 × 输入门的接受比例。这两个运算都是逐元素的也就是说细胞状态里面不同维度上的信息可以独立地决定自己是保留还是更新。有些维度专门用来存当前讨论的主体是谁可能几十个时间步都保持不变;有些维度专门用来存最近出现了一个数字可能下一秒就被覆盖掉。这就是LSTM能记住长距离信息的根本原因信息流通的路径是加性的不需要经过激活函数的非线性压缩。你可以把这条传送带想成一个流水账本老账页只是被新账页覆盖或补充而不是整本账本被重写。3.4 输出门不是所有记忆都要对外暴露[ o_t \sigma(W_o \cdot [h_{t-1}, x_t] b_o) ][ h_t o_t \ast \text{tanh}(C_t) ]最后当前时刻要对外输出的隐状态 ( h_t ) 不是直接把细胞状态拿出去而是先用tanh把细胞状态压到(-1, 1)范围内再用输出门 ( o_t ) 控制哪些内容可以对外说哪些内容自己记着就行。打个比方细胞状态像你的长期记忆包含大量信息;隐状态像你在当前对话里说出口的话。你不会把所有记忆一股脑全倒出来只会根据当前的问题和语境挑一小部分说出来。输出门就是做这个挑选的。3.5 我当年学LSTM时最容易绕晕的一个点如果你正在学LSTM大概率也会被一个问题卡住为什么隐状态 ( h_t ) 和细胞状态 ( C_t ) 同时存在信息到底存哪里我的理解方式是这样的细胞状态是长期仓库专门负责跨时刻传递长距离信息它走的是加性更新的路线所以梯度可以很顺畅地从未来传回过去。隐状态是工作记忆是当前时刻真正要用的信息它既要参与当前时刻的输出计算也要作为下一时刻的输入之一传给门控单元做判断。两个状态分工合作缺一不可。你只需要记住一个核心原则门控的计算都依赖 ( h_{t-1} ) 和 ( x_t )而跨时刻传递的核心载体是 ( C_t )。看代码的时候凡是形状是(batch, hidden_size)的张量你就要想清楚它到底是细胞状态还是隐状态别混淆。4. 遗忘门为什么是LSTM的灵魂从梯度路径看设计精髓很多人学了LSTM的公式之后觉得最难的部分是输入门和输出门要记的公式太多。但实际上站在梯度流动的角度看遗忘门才是整个LSTM里最关键的创新。前面在讲RNN的时候说了传统RNN的梯度在时间维度上会被反复乘以 ( W_{hh} )导致梯度指数级衰减。而LSTM的细胞状态 ( C_t ) 更新方式是[ C_t f_t \ast C_{t-1} i_t \ast \tilde{C}_t ]注意这里细胞状态对上一时刻细胞状态的偏导是[ \frac{\partial C_t}{\partial C_{t-1}} f_t ]也就是说梯度从 ( t ) 时刻传回 ( t-1 ) 时刻时乘的不是一个学习出来的权重矩阵而是一个遗忘门的值。遗忘门是一个0到1之间的数而且这个数是模型自己根据当前输入学出来的可以接近1。这意味着什么意味着模型在训练过程中有能力学会让某些维度的遗忘门保持接近1从而让关键信息沿着细胞状态这条高速公路一路畅通地传回很多个时间步之前。梯度消失问题就这样被结构性地缓释了。我见过一些学习资料的讲解顺序是先给公式、再给结构图、最后说一句LSTM解决了梯度消失。但我个人觉得这个顺序是错的。你只有先理解梯度在循环结构中为什么会消失才能理解遗忘门为什么需要存在为什么它比RNN里的 ( W_{hh} ) 高明。更让我觉得精妙的是LSTM的遗忘门不是固定值而是输入依赖的。模型可以学会当读到句号时遗忘门变小把上一句的内容清空当读到从属连词时遗忘门接近1把主句信息保留下来。这是门控机制真正强大的地方——网络可以根据内容动态调节记忆的存留策略。5. GRU把三扇门砍成两扇门之后性能反而没掉多少LSTM是1997年提出的。差不多二十年之后2014年Cho等人提出了门控循环单元也就是GRU。它的思路是LSTM有三个门加一个候选状态能不能简化一点GRU把细胞状态和隐状态合并成同一个状态向量只保留两个门重置门和更新门。[ z_t \sigma(W_z \cdot [h_{t-1}, x_t] b_z) ][ r_t \sigma(W_r \cdot [h_{t-1}, x_t] b_r) ][ \tilde{h}t \text{tanh}(W_h \cdot [r_t \ast h{t-1}, x_t] b_h) ][ h_t (1 - z_t) \ast h_{t-1} z_t \ast \tilde{h}_t ]我逐行拆一下。( z_t ) 是更新门类似旧记忆保留多少、新内容写入多少的比例分配器。( r_t ) 是重置门控制计算候选状态时上一时刻的隐状态有多少被考虑进来。候选状态 ( \tilde{h}_t ) 是当前输入结合了重置后的旧信息的临时状态。最后新的隐状态是旧状态和候选状态的加权平均。对比一下LSTM和GRU特性LSTMGRU门数量3个遗忘、输入、输出2个更新、重置状态数量细胞状态 隐状态只有隐状态参数量更多更少表达能力理论上更强可独立控制写多少和扔多少稍弱但大部分任务差距很小训练速度稍慢稍快适合场景数据量大、序列长、需要精细记忆控制数据量中等、希望快速迭代建模这里的核心权衡是LSTM有独立的遗忘门和输入门它可以同时做到旧信息保留90%和新信息写入10%这种精细操作。GRU用同一个更新门来调配旧信息和候选信息两者共享一个比例自由度少了一个但参数量也少了过拟合风险更低。实际使用中我个人的经验是在多数中等规模的任务上GRU和LSTM的最终效果差距在误差范围内。如果你的任务需要建模特别长的依赖关系或者数据量足够大可以优先试LSTM;如果希望模型更轻量、训练更快GRU是更省心的选择。6. 学了LSTM之后必须做的事别在错误的场景里用它6.1 什么时候该用RNN/LSTM用大白话总结RNN/LSTM适合的输入是具有时间或顺序结构的序列数据先后顺序有意义而且信息需要跨时间步传递。典型场景包括文本序列字符级或词级建模语言模型、机器翻译、文本生成时间序列股票价格、气温、电力负荷、传感器数据预测语音信号语音识别、语音合成的前端特征建模视频帧序列动作识别、行为预测6.2 什么时候不该用RNN/LSTM有两类场景我很不建议上RNN/LSTM。第一类是输入本身没有顺序依赖的表格数据。比如一个客户数据集每行是不同客户的属性行与行之间没有先后关系你把它们排成序列丢给LSTM模型学不到比树模型或MLP更多的东西反而增加训练成本和过拟合风险。第二类是长文本全局建模。像一篇5000字的文章如果你想建模句子与句子之间跨越几百个token的复杂依赖RNN/LSTM的内存效率和梯度路径都会变得不太友好。BERT等Transformer系模型在这类场景里表现得更好因为它们的注意力机制让任意两个位置之间的连接代价都是恒定的不受距离影响。6.3 Transformer会不会取代RNN/LSTM现在很多人一谈到序列建模第一反应就是Transformer。确实在机器翻译、文本生成这些大任务上Transformer已经是绝对主流。但我个人不太喜欢取代这个词。RNN/LSTM依然有它们的相对优势计算复杂度与序列长度是线性关系每个时间步只处理当前输入不跟所有历史时刻做注意力交互推理时状态量固定非常适合流式输入、在线推理、低延迟场景。比如实时语音识别、边收数据边预测的工业监控这些场景里Transformer那种必须看到全序列才能算注意力的性质反而不太方便。所以现实是大模型时代Transformer是主力;小规模序列建模、流式预测、资源受限场景RNN/LSTM/GRU仍然是能打的方案。我的建议是两类架构都学清楚理解各自的设计动机以后面试也好、写方案也罢你都能说清楚为什么选这个而不是那个。7. 一个完整的PyTorch时间序列预测Demo从数据到训练前面讲了不少理论但最终要落到代码上才真正算掌握。这一节我会用PyTorch写一个最简单的LSTM时间序列预测例子从造数据到训练、预测全流程跑通。这个代码我实际跑过可以直接复制改改数据就用。7.1 造一份带有周期性和趋势性的序列数据为了体现LSTM的价值我不用那种纯随机噪声数据而是用两类信号的叠加import torch import torch.nn as nn import numpy as np import matplotlib.pyplot as plt # 生成带有趋势性和周期性的序列 np.random.seed(42) t np.arange(0, 1000, 0.1) signal 0.5 * t / 100 np.sin(t / 10) 0.2 * np.random.randn(len(t))这里信号有三部分线性上升趋势 ( 0.5t/100 )正弦周期以及高斯噪声。这种数据贴近真实场景既需要记住近期趋势又需要捕捉周期性模式。7.2 造滑窗样本LSTM一次吃一个序列片段然后预测下一个时刻的值。所以要把原始序列切成 ( (序列长度, 预测步数) ) 的样本对。seq_len 20 def make_samples(data, seq_len): X, y [], [] for i in range(len(data) - seq_len): X.append(data[i:i seq_len]) y.append(data[i seq_len]) return np.array(X).reshape(-1, seq_len, 1), np.array(y).reshape(-1, 1) X, y make_samples(signal, seq_len) train_size int(len(X) * 0.8) X_train, X_test X[:train_size], X[train_size:] y_train, y_test y[:train_size], y[train_size:] # 转成 PyTorch Tensor X_train_t torch.tensor(X_train, dtypetorch.float32) y_train_t torch.tensor(y_train, dtypetorch.float32) X_test_t torch.tensor(X_test, dtypetorch.float32) y_test_t torch.tensor(y_test, dtypetorch.float32)注意LSTM的输入形状是(batch_size, seq_len, input_size)我只用了1个特征维度所以input_size1。7.3 定义LSTM模型class LSTMPredictor(nn.Module): def __init__(self, input_size, hidden_size, num_layers): super().__init__() self.lstm nn.LSTM(input_size, hidden_size, num_layers, batch_firstTrue) self.fc nn.Linear(hidden_size, 1) def forward(self, x): out, (h_n, c_n) self.lstm(x) # 取最后一个时间步的隐状态做预测 last_hidden out[:, -1, :] return self.fc(last_hidden) model LSTMPredictor(input_size1, hidden_size32, num_layers2)这里hidden_size32是比较常规的隐状态维度num_layers2意味着两层LSTM堆叠。第一层LSTM的输出作为第二层LSTM的输入增加了模型的非线性表达能力。7.4 训练循环criterion nn.MSELoss() optimizer torch.optim.Adam(model.parameters(), lr0.001) epochs 100 batch_size 64 for epoch in range(epochs): model.train() permutation torch.randperm(X_train_t.size(0)) total_loss 0 num_batches 0 for i in range(0, len(permutation), batch_size): idx permutation[i:i batch_size] batch_x X_train_t[idx] batch_y y_train_t[idx] optimizer.zero_grad() out model(batch_x) loss criterion(out, batch_y) loss.backward() optimizer.step() total_loss loss.item() num_batches 1 if (epoch 1) % 20 0: print(fEpoch {epoch 1}, Loss: {total_loss / num_batches:.6f})这段代码有几个训练实操中的关键细节。每次epoch开始时要做torch.randperm打乱样本顺序避免模型学到样本顺序的假规律;每个batch内部要optimizer.zero_grad()清零梯度否则PyTorch默认会累加梯度;训练结束后用.detach().numpy()取出预测结果来画图否则带梯度的张量无法直接转为numpy数组。7.5 预测与可视化model.eval() with torch.no_grad(): pred model(X_test_t).numpy() plt.figure(figsize(12, 5)) plt.plot(y_test[:200], labelTrue) plt.plot(pred[:200], labelPred) plt.legend() plt.show()跑完之后你应该能看到预测曲线和真实曲线在绝大多数位置高度重合证明LSTM确实学到了这个序列的规律。7.6 这个Demo里的几个关键调参点第一窗口长度seq_len需要根据任务特性来设。太短捕捉不到周期依赖太长不仅计算量大也可能引入过多的历史噪声。初学者可以多试几个值10、20、50观察测试集误差的变化。第二归一化。我上面的代码为了简单直接用了原始信号但真实项目中几乎都必须做Min-Max归一化或标准化。因为LSTM内部的sigmoid/tanh激活函数对输入数值范围比较敏感如果不做归一化模型可能很难收敛。建议把训练数据的最小值和最大值保存下来预测后再反变换回原始量纲。第三学习率。Adam的默认学习率0.001在多数任务上表现不错但如果你发现loss震荡或者下降太慢可以考虑降低到0.0005或0.0003。8. LSTM训练避坑指南从我的真实踩坑经验说起8.1 坑一数据泄漏预测结果漂亮得像个假象这是时序预测里最隐蔽也最致命的坑。做数据划分时如果你直接从完整序列里随机打乱划分训练集和测试集训练集里就会混入未来数据的信息。预测时看起来精度很高但一到真正的未来数据上模型完全没法用。正确做法是按时间顺序划分测试集只能是训练集之后的那段数据。我上面的代码就是这么做的——前80%训练后20%测试这一点写清楚很有必要。我还见过一些开源项目里偷懒做随机划分性能数字异常好看这类模型上线即翻车。8.2 坑二梯度裁剪到底加不加RNN在训练过程中确实存在梯度爆炸的可能尤其是当序列很长、权重初始化不当时。LSTM虽然缓解了梯度消失但梯度爆炸仍然可能出现。一个非常廉价的保护手段是梯度裁剪torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm5.0)在loss.backward()optimizer.step()之间加上这一行把梯度的L2范数限制在5以内。我最开始学LSTM时总觉得自己用不到这个玩意直到有一次loss直接跳到NaN才意识到梯度爆炸有多简单粗暴。这个操作对模型质量影响很小却能避免训练过程崩溃属于加上永远不吃亏的防护。8.3 坑三隐状态初始化PyTorch的nn.LSTM默认会自动初始化隐状态为全0所以你可以不管它。但有两个细节要注意如果batch_firstTrue你传入的输入形状是(batch, seq, feature);如果batch_firstFalse默认是(seq, batch, feature)。新手最常见的报错就是batch维度顺序搞反了。另外如果你自己手动传初始隐状态比如h0和c0要确保每个都初始化为(num_layers, batch, hidden_size)的形状。它们分别代表每一层LSTM的隐状态和细胞状态。8.4 坑四多步预测误差累积很多新手做多步预测时会把LSTM的输出当成单步结果用。但如果你要预测未来5个时刻一种做法是每预测一个时间点就把这个预测值当作下一时刻的输入喂回模型这叫递归多步预测。缺点是误差会逐时刻累积前一步偏一点后面就越偏越远。另一种做法是直接用多输出结构让模型一次输出未来5个值。两者各有取舍。如果任务非常需要长时间预测我建议改用seq2seq结构加注意力机制或者干脆换成Transformer。LSTM适合短中期预测强行让它做超长程预测属于为难它。8.5 坑五不要迷信隐状态维度越大越好隐状态维度就是模型记忆容量。容量越大理论记忆越强但过拟合风险也随之上升。对于小规模时间序列hidden_size32或64通常已经足够;你把它加到256、512可能训练集loss降得更低但测试集误差反而变大。这跟全连接网络加宽的道理完全一样别觉得LSTM就该用大隐状态。一个更实际的建议是先调序列长度再调隐藏层大小最后调层数。序列长度对时序预测效果的影响往往比隐藏层维度更显著因为窗口内包含的信息量和信息错位直接决定模型能看到的上下文范围。9. 从RNN到Transformer循环网络的下一步在哪里作为一个关注序列建模快十年的从业者我对RNN/LSTM的感情是真复杂。一方面这些结构确实老了在很多大规模自然语言处理任务上被Transformer甩开;另一方面我又觉得它们的思路至今仍然影响深远——门控机制、状态传递、编码器-解码器框架这些概念已经被Transformer继承和改造了就连Transformer需要给词向量加位置编码本质上也是在回应序列顺序信息不能丢失这个老问题。如果你要继续往前走我建议按这条路线进阶先跑通RNN和LSTM的代码理解它们的设计动机;然后学习注意力机制理解它如何解决信息跨长距离传递和并行计算两个问题;最后看Transformer原论文和BERT相关实现把位置编码、多头注意力、自注意力这些概念弄明白。你会发现RNN/LSTM不是被扔进历史垃圾桶的旧技术而是理解整个序列建模演进路线的最佳第一站。如果你对序列建模很有兴趣还可以进一步了解双向LSTM、注意力机制与LSTM的结合比如Bahdanau Attention、sequence-to-sequence模型以及它们怎么被用来做机器翻译和文本摘要。再往后走就是多头自注意力、Transformer Block、预训练语言模型那条更现代的路了。作为一个曾经在LSTM调参上花过无数个通宵的人我最后的真心话是算法迭代再快底层的核心思想不会消失。你把为什么要有遗忘门梯度为什么会在时间维度衰减这种最简单的问题吃透了以后看任何更复杂的架构都会快很多甚至能在模型出问题时直接猜到瓶颈在哪个环节。
上一篇/下一篇内容由系统自动关联 返回资讯列表 →