翻书翻到《Deep Learning》循环神经网络这一部分,我第一感觉是:书里写得实在太省了。LSTM的公式就那么几句话,图表也不多,但“为什么要这么设计”“这几个门到底在解决什么问题”“手写代码时该注意什么”,全靠读者自己悟。这篇读书笔记是系列的第七篇,核心就一个:循环神经网络里的LSTM到底是怎么工作的,以及怎么把它用在真实任务里。
很多初学者卡在LSTM上,不是说公式看不懂,而是不知道这些公式分别对应什么直觉。更常见的情况是:论文读懂了一堆术语,自己拿PyTorch一写就报错;或者训练一个时间序列模型,loss怎么都不降。这篇文章不打算复述教材,而是把我从“看着公式发呆”到“能独立调通LSTM训练”整个过程中的思考、源码解读、实操代码和踩坑记录整理出来,适合三类人:
- 正在读深度学习教材、卡在循环神经网络章节的人。
- 想用LSTM做时间序列预测(比如销量预测、设备寿命预测)但不知道从哪下手的人。
- 想把PyTorch源码里那堆参数和维度搞清楚的人。
先说一个总体感受:LSTM其实是把“记忆”这件事拆成了几个可学习的开关,理解了这几个开关,后面的代码、调参、训练就都顺了。
1. 为什么我读LSTM时先放弃了数学推导
1.1 从普通RNN的“短期失忆”说起
先回顾一下循环神经网络的雏形。所谓的循环,就是同一个网络在每个时间步上反复运行,把上一时刻的隐藏状态传给下一时刻。不断把传入的信息压缩成一个固定维度的向量,再用这个向量做预测。这个思路本身没毛病,问题是当你反向传播时,误差信号要穿过时间维度一层一层往回传,每个时间步的梯度都要乘一次矩阵。
如果这个矩阵的特征值小于1,乘的次数多了以后梯度会指数级衰减。长序列里,前面几步的梯度基本上趋近于0,等于说网络根本学不到很久以前的信息。如果特征值大于1,梯度又会指数级爆炸,训练直接崩掉。这就是“梯度消失/爆炸”问题的来源,也是普通RNN只能记住短期依赖的根本原因。
打个比方:想象一排人传话,第一个人说A,第二个人传给第三个人时只能记住大概,第五个人再传时就丢了一半,传到第十个人那里只剩“确实说了点什么”的模糊印象。普通RNN就是这个传话游戏,传几步还行,传一百步就全散了。学语言的人都知道,一句话里关键信息常常出现在十几步之前,所以普通RNN在真实场景里几乎没法用。
1.2 门控机制的直觉:一条带闸门的传送带
那LSTM是怎么解决的?思路其实特别朴素:与其让信息在隐藏状态里被一次次非线性压缩、遗忘,不如专门修一条“高速公路”,让信息能在时间维度上相对完整地流动。这条高速公路就是细胞状态(cell state),用符号c表示。细胞状态上只有少量线性运算,比如加减乘除和逐元素乘法,没有非线性激活,所以梯度在反向传播时不会因为链式法则被反复压缩。
但光有高速公路也不行,不然什么信息都往里塞,记忆就成了一锅粥。所以LSTM加了三个门,相当于高速公路上的闸口:一个决定要不要丢掉旧的记忆,一个决定要不要把新信息写进去,一个决定拿当前状态对外输出什么。生活里可以类比成一个记账本:每天都有大量信息,但记什么、删什么、月底汇报什么,都需要取舍。LSTM就是让网络自己去学这个取舍规则的门控机制。
很多人一开始看公式觉得复杂,其实就是三个开关加一条传送带。后面我会把每一步拆开来看,保证看完能自己写出来。
2. LSTM内部机制拆解:三个门和一条细胞状态
2.1 三个门到底在做什么
LSTM在每个时间步t做的事情,可以拆成四个步骤。先定义几个变量:当前输入是x_t,上一时刻隐藏状态是h_{t-1},上一时刻细胞状态是c_{t-1},当前时刻输出是h_t,更新后的细胞状态是c_t。所有门都依赖h_{t-1}和x_t,所以门本质上是“根据过去的输出和当前的新输入,决定接下去怎么处理记忆”。
第一个是遗忘门:决定上一时刻的细胞状态c_{t-1}要保留多少。公式是f_t = σ(W_f · [h_{t-1}, x_t] + b_f)。别被矩阵吓到,它就是一个带sigmoid的全连接层。输出在0到1之间,0表示彻底忘掉,1表示全保留。举个例子,读句子“我在公园遇到一只猫,它很可爱”的时候,等到要预测“它”指代什么,网络需要保留“猫”这个信息;读到“遇到”这个动词时可能已经不需要“公园”这个地点了,遗忘门会让这部分记忆衰减。
第二个是输入门:决定当前这一时刻的新信息里,哪些值得写入细胞状态。公式里通常写成i_t = σ(W_i · [h_{t-1}, x_t] + b_i),它负责筛选强度;同时还要生成候选记忆c̃_t = tanh(W_c · [h_{t-1}, x_t] + b_c)。候选记忆提供了新的内容,它用tanh把数值范围压到-1到1,而输入门决定这个候选记忆有多少能进细胞状态。新旧相加后的细胞状态就是c_t = f_t ⊙ c_{t-1} + i_t ⊙ c̃_t。这里的⊙是逐元素乘,不是矩阵乘。
第三个是输出门:决定对外输出哪些信息。公式是o_t = σ(W_o · [h_{t-1}, x_t] + b_o),然后h_t = o_t ⊙ tanh(c_t)。细胞状态本身可以是一堆内部数值,但对外汇报要挑重点,所以先用tanh把细胞状态压到[-1,1],再用输出门过滤一遍。理解输出门还有个关键点:LSTM的输出h_t并不是记忆本身,而是从记忆里抽取的一个“对外版本”。以后你看源码时,看到return的永远是h而不是c,就是因为它只负责对外汇报。
这三段加一起你会发现,LSTM之所以能记住长期信息,核心在c_t那条通道。反传梯度要经过c_t时,路线是c_t = f_t ⊙ c_{t-1} + ...,如果遗忘门接近1,梯度几乎原样传回去,不会衰减。这就是为什么LSTM能在几十甚至几百步之后还能记住学过的东西。
2.2 用PyTorch源码解读维度与运算过程
看懂了公式,再看PyTorch源码就轻松多了。PyTorch里的nn.LSTM和nn.LSTMCell封装了同样的运算,区别在于LSTMCell只做一个时间步,LSTM帮你循环处理整个序列。我自己更喜欢先看LSTMCell的源码理解运算逻辑,再接LSTM来用。
LSTMCell的输入是x和上一时刻的(h, c)。定义的时候只需要指定input_size和hidden_size。它内部其实就做了一个大矩阵乘,把所有门的权重拼在一起一次算完:
import torch import torch.nn as nn # 单步 LSTM 的逻辑参考 def lstm_cell_forward(x, h_prev, c_prev, W_ih, W_hh, b_ih, b_hh): gates = x @ W_ih.t() + b_ih + h_prev @ W_hh.t() + b_hh # 把 gates 按 hidden_size 切成4份,顺序是 i, f, g, o 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_next = f * c_prev + i * g h_next = o * torch.tanh(c_next) return h_next, c_next这里的4份对应输入门i、遗忘门f、候选值g、输出门o。PyTorch底层把这个矩阵运算做成了一个大矩阵乘法,所以参数总量是4 * hidden_size * (input_size + hidden_size + 1),最后的1是bias。如果你检查nn.LSTM(10, 20)的参数,会发现weight_ih_l0的形状是(80, 10),weight_hh_l0是(80, 20),80正好是4 * 20,对应4个门。
再看nn.LSTM的常用参数,很多人一上来会被batch_first、num_layers、bidirectional这几个参数弄晕:
- batch_first=True时输入形状是(batch, seq_len, input_size),否则是(seq_len, batch, input_size)。建议一律用batch_first=True,不然数据预处理时容易出错。
- num_layers表示堆叠几层LSTM。第一层输出的隐藏状态作为第二层的输入,默认是1。堆叠适合复杂序列,但也会增加过拟合风险。
- bidirectional=True使用双向LSTM,前向和后向两个方向的隐藏状态会拼接。双向适合做分类、命名实体识别这类可以看全文的任务,但不适合纯时间序列预测,因为预测时你还没有未来数据。
- 输出有两部分:outputs是每个时间步的h,形状是(batch, seq_len, hidden_size * num_directions);最后返回的(h_n, c_n)是最后一个时间步的隐藏状态。很多教程里只拿outputs[:, -1, :]当最终特征,其实在多层LSTM里,最后一层最后一个时刻的h_n[:, -1, :]和outputs[:, -1, :]是同一个东西,但中间层的h不会被输出,这点阅读源码时要分清。
3. 用LSTM做一个时间序列预测:从数据到训练
理论说得再多,不如跑一个实际项目。我用LSTM做了一个单变量时间序列预测的任务:给定过去N天的数值,预测未来几天的数值。这个流程可以平移到销量预测、流量预测、设备剩余寿命预测等场景。下面直接给代码和关键步骤。
3.1 数据预处理:滑动窗口、归一化、时序切分
时间序列预测和图像分类不太一样,不能随便打乱数据。训练样本之间虽然可以随机抽取,但必须保证每个样本内部是连续的时间段。具体做法是用一个固定长度的滑动窗口把序列切成样本。假设原始数据有1000个点,窗口长度是24,预测目标是未来3个点,那么第1个样本是0到23,预测24到26;第2个样本是1到24,预测25到27,以此类推。
我通常会把切好的样本转成PyTorch的Dataset:
import torch from torch.utils.data import Dataset class SequenceDataset(Dataset): def __init__(self, data, seq_len, pred_len): self.data = torch.FloatTensor(data) self.seq_len = seq_len self.pred_len = pred_len def __len__(self): return len(self.data) - self.seq_len - self.pred_len + 1 def __getitem__(self, idx): x_start = idx x_end = idx + self.seq_len y_start = x_end y_end = y_start + self.pred_len x = self.data[x_start:x_end] y = self.data[y_start:y_end] return x, y有几个细节要特别强调。第一,归一化时绝对不能用整个数据集的均值和标准差,那会造成数据泄漏,因为测试集的信息提前泄露进了训练过程。正确做法是只用训练集的统计量归一化,然后把同样的scale和shift应用到验证集和测试集。第二,如果特征量纲差异大,LSTM很难训练,数值小的特征会被高维特征盖过去。第三,切分时按时间顺序留出最后20%作为测试集,不要去随机抽样,否则时间关系就错了。
3.2 模型定义与训练循环代码实战
模型本身很简单,就是个LSTM加一层全连接:
import torch.nn as nn class LSTMPredictor(nn.Module): def __init__(self, input_size=1, hidden_size=64, num_layers=2, output_size=3, dropout=0.2): super().__init__() self.lstm = nn.LSTM( input_size=input_size, hidden_size=hidden_size, num_layers=num_layers, batch_first=True, dropout=dropout ) self.fc = nn.Linear(hidden_size, output_size) def forward(self, x): out, _ = self.lstm(x) # out: (batch, seq_len, hidden_size) out = out[:, -1, :] # 只用最后一个时间步的隐藏状态 return self.fc(out)训练循环我习惯套一个固定模板,里面包含三件容易被忽略的事:梯度裁剪、学习率调节、每个epoch记录验证集loss。
import torch.optim as optim from torch.utils.data import DataLoader model = LSTMPredictor(input_size=1, hidden_size=64, num_layers=2, output_size=3) optimizer = optim.Adam(model.parameters(), lr=3e-4) criterion = nn.MSELoss() train_loader = DataLoader(train_dataset, batch_size=64, shuffle=True) val_loader = DataLoader(val_dataset, batch_size=64, shuffle=False) best_val_loss = float("inf") for epoch in range(100): model.train() train_loss = 0.0 for x, y in train_loader: optimizer.zero_grad() pred = model(x.unsqueeze(-1)) # 输入需要增加特征维度 loss = criterion(pred, y) loss.backward() nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.step() train_loss += loss.item() * x.size(0) train_loss /= len(train_dataset) model.eval() val_loss = 0.0 with torch.no_grad(): for x, y in val_loader: pred = model(x.unsqueeze(-1)) loss = criterion(pred, y) val_loss += loss.item() * x.size(0) val_loss /= len(val_dataset) if val_loss < best_val_loss: best_val_loss = val_loss torch.save(model.state_dict(), "best_model.pt") print(f"epoch {epoch:3d} | train loss {train_loss:.5f} | val loss {val_loss:.5f}")这里有个细节我想多说一句:x.unsqueeze(-1)是为了把输入从(batch, seq_len)变成(batch, seq_len, 1),LSTM要求的输入最后必须是特征维度。单变量序列只有一列特征,所以补上1。如果你有多变量,就把多个特征堆叠在最后一维上,比如流量、温度、湿度三个字段就是input_size=3。
3.3 参数选择的工程经验
很多新手问我hidden_size、num_layers、seq_len到底怎么选。说实话没有万能公式,但有几个经验值可以给大家参考。
- hidden_size:一般取32到128之间。太小记不住信息,太大容易过拟合,而且显存消耗直线上升。先取64,看验证集loss的趋势再调整。
- num_layers:大多数单变量时间序列任务2层就够,再多容易出现梯度不稳定和训练变慢。3层以上通常只在数据量很大、任务很复杂时划算。
- seq_len(窗口长度):这个取决于任务的记忆长度。如果预测每天的销量,序列可能和“上一周”“去年同期”有关,窗口可以取7、14、30。但窗口也不是越长越好——太长的窗口会把无关信息也塞进去,而且训练更慢。可以先从经验值开始,比如用seq_len=24小时或7天,再反复尝试。
- learning rate:我建议从3e-4开始,这是Adam比较稳的基线。如果loss不降,试着减小到1e-4;如果训练很慢,试试1e-3。
- batch_size:时间序列预测里我一般不取太大,32或64比较稳。太大容易收敛慢,也会有内存压力。
预测多步时,输出层的维度决定你要预测几个未来时间点。想预测3天就把output_size设成3,想预测7天就设成7。直接把输出层设成多个点是“直接多步预测”,实现简单,但长期预测误差会累积的现象更容易出现。另一个方案是递归预测,即每次预测一步,把预测结果当作下一步的输入;这种方案在长期预测时误差会慢慢漂移。选择哪个没有绝对好,通常是试着来,以验证集为准。
4. 训练LSTM最常见的坑和排查办法
LSTM训练过程比普通前馈网络更容易出幺蛾子。我这几年的实操里,几乎把常见问题都踩过一遍,这里挑几个最典型的分享。
4.1 我踩过的五个坑
第一个坑:归一化用了全量数据统计量。这是数据泄漏的重灾区。我在早期做销量预测时,用整个序列的均值和方差做了MinMaxScaler,把训练和测试数据混在一起归一化,结果测试集上的指标好得离谱。后来我仔细一查才发现,模型其实偷偷“看”到了测试集的统计信息。线下评测和线上真实表现的差距,几乎都来自这种数据泄漏。后面所有时间序列项目里,我都先只拟合训练集的scaler,再transform验证和测试集。
第二个坑:loss变成了NaN。训练中途突然出现NaN,我的排查顺序是:先看数据里有没有NaN或Inf,再看学习率是不是太大,最后看是否梯度爆炸。LSTM的梯度路径长,爆炸概率比普通网络高得多,所以不管任务大小,我都会加上梯度裁剪。我的经验是max_norm取1.0或5.0都能起到保护作用。
第三个坑:loss卡住不降。最常见的原因是输出层缺少合适的激活,或者输出数值范围与目标差太多。比如目标值在0到1之间,模型输出却是任意实数,这个时候网络要自己学一个缩放,学得慢很正常。更常见的原因是数据没归一化,或归一化时不彻底。还有一个容易被忽略的原因:LSTM的初始h和c没给好。PyTorch默认初始化为0其实通常问题不大,但有些人喜欢自己初始化成很大的数,结果训练半天梯度消失。老老实实用默认初始状态。
第四个坑:预测结果总是“上一个值的平移”。这是时间序列预测里典型的惰性预测问题。模型发现把最近一个观测值复制一下,loss就很小,于是懒得学真实规律。这种情况常见于高度随机的数据。缓解办法有三个:一是对数据进行差分,让模型预测增量而不是原始值;二是增加seq_len,让模型看到更多上下文,削弱“抄近路”的动机;三是用多步预测目标,让模型必须学到趋势才能降低loss。
第五个坑:训练集loss狂降,验证集loss上升。这就是过拟合。LSTM的过拟合和别人没两样,但有一个特殊的点:序列数据内部高度自相关,如果验证集的划分时间太接近训练集,看起来验证集很好,实际泛化很差。所以在划分时,验证集和测试集都要跟训练集留出足够的时间间隔,比如训练集最后一天和验证集第一天之间隔一两周,模拟真实场景。
4.2 LSTM调参速查表
为了方便排查,我把常见情况整理成一个速查表:
| 症状 | 可能原因 | 建议处理方式 |
|---|---|---|
| loss=NaN | 学习率过高、梯度爆炸、数据含Inf | 降低学习率,加梯度裁剪,检查数据清洗 |
| loss不降 | 数据未归一化、输出层不合适、初始状态异常 | 用训练集统计量归一化,检查输出层尺度,重置模型 |
| 预测是滞后值 | 模型“抄近路”,数据随机性大 | 差分序列、加窗口长度、换多步预测目标 |
| 训练好测试差 | 过拟合、数据泄漏 | 减少hidden_size/num_layers,加dropout,检查划分布局 |
| 训练慢 | 序列太长、模型太深、batch太大 | 缩短seq_len可以快很多,减小深度,降batch_size |
| 长序列记不住 | 遗忘门学成关闭状态、网络容量不足 | 增加hidden_size,检查是否多层需要更多数据,尝试双向(若允许) |
4.3 关于设备寿命预测:从单步到多步的迁移
热词里经常出现“lstm设备寿命预测实战”。这类任务本质上还是时间序列预测,把传感器信号(振动、温度、压力等)作为输入,去预测剩余使用寿命(RUL)。但直接从销量预测切过去,还是有几个坑。
第一是工况变化。设备的工作环境不是恒定的,转速、负载、环境温度都会变,模型如果只学了单一工况下的规律,换工况就失效。解决办法是把工况变量也作为输入特征,或者按工况分段做数据增强。第二是多变量问题。设备寿命预测很少只有一个传感器,通常有十几个甚至上百个通道。LSTM的input_size直接等于通道数,但如果特征太多、样本太少,过拟合会非常严重。可以先用PCA或重要特征筛选压缩维度,别一股脑全塞进去。第三是多步预测策略。设备寿命预测的目标往往是一串剩余寿命曲线,而不是单点预测,递归预测容易累积误差,直接多步输出又可能让模型只盯着最近几步。比较稳的做法是输出未来一段时间的健康指标,再根据阈值反推寿命区间。
顺便说一句,看到网上有人问“能不能直接拿一个LSTM代码改吧改吧做寿命预测”,我的建议是先把流程跑通,不要急着换模型。拿一段公开的传感器数据,把3.2节的代码改改input_size和output_size,跑通之后再去研究特征工程。原理扎实比模型花哨重要得多。
5. LSTM之后:变体、GRU,以及什么时候该放弃它
5.1 LSTM与GRU的取舍
读者多半会遇到GRU。GRU是LSTM的简化版,它把遗忘门和输入门合并成更新门,又移除了细胞状态,只保留了隐藏状态。它的参数比LSTM少,训练速度更快,在不少任务上效果和LSTM差不多。
如果项目还没有跑起来,我通常会建议先用GRU快速验证数据是否有可学习信号,等确定需要更强表达力再换成LSTM。但如果任务需要精细的记忆控制,或者你已经验证了LSTM效果更好,不要为了参数少而强行换GRU。具体可以用一个小表格对比:
| 对比项 | LSTM | GRU |
|---|---|---|
| 参数数量 | 多(三个门+候选值) | 少(两个门) |
| 训练速度 | 较慢 | 较快 |
| 长期依赖建模 | 细胞状态通道更显式 | 依赖隐藏状态循环,相对隐式 |
| 适用场景 | 数据量大、需要精细记忆控制 | 数据量小、快速原型验证 |
5.2 双向、多层和堆叠的适用场景
关于双向LSTM,我想多说一点。双向的优势是每个位置既能看左边也能看右边,很适合文本分类、命名实体识别这类对全局信息敏感的任务。但代价是推理时必须有完整输入序列,任何不能预知未来的在线预测任务都不能用。多层LSTM则适合自动提取层次化特征,底层学习短时局部结构,高层学习长时依赖。堆叠时要留意dropout位置:PyTorch里只在层间的输入上做dropout,最上层输出不做,这个行为和CNN的spatial dropout类似,能缓解过拟合但也不能完全依赖,数据不够时不建议堆太深。
5.3 从LSTM到注意力机制:LSTM的瓶颈在哪
我每次读LSTM的源码都会感叹:它把“记忆”做到了极致,但它的训练是串行的。每一个时间步都要依赖前一步的结果,没法并行,所以处理长序列时训练效率很低。也就是因为这个瓶颈,后来大家才想到让序列中任意两个位置都能直接跳到对方,不需要逐步传递。这个思路后来变成了注意力机制,再后来演变成Transformer,如今大语言模型的基础架构。有人说LSTM是不是过时了,对这种说法我持保留态度。大语言模型时代的很多任务确实用不到LSTM,但在时间序列、工业预测、小数据场景里,LSTM依然是一个极其可靠、容易训练、方便部署的模型。你理解了LSTM,再去看注意力机制,你会更容易理解它到底想解决什么。
最后分享一个我个人的体会:读LSTM这类模型,真正有用的动作不是盯着公式背,而是亲手用LSTMCell把公式敲一遍,再跟nn.LSTM的结果做对比。你一旦发现两个结果一模一样,门、细胞状态、输出这些概念就彻底落地了。后续你去看任何基于LSTM的论文、改任何代码,心里都会稳得多。