ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

Informer模型复现详解:长序列时间序列预测与稀疏自注意力实战

Informer模型复现详解:长序列时间序列预测与稀疏自注意力实战 Informer这篇论文刚出来那阵子做时间序列预测的基本人手一份源码跑实验。我前后用它在ETT、电力负荷、天气这几个数据集上折腾了大半个月把整个模型从数据处理到训练评估完整重写了一遍。今天这篇就把Informer复现的全过程拆开讲清楚从算法原理到核心代码再到那些论文里不会写但测试中必定会踩的坑一次性梳理完。这篇内容适合谁看如果你正在做长序列时间序列预测LSTF或者想搞懂稀疏自注意力在Transformer里到底怎么落地又或者你只是想把Informer源码跑通但卡在某个地方这篇都能给你一个完整的参考路径。我会把每个模块的输入输出、尺寸变化和参数选择逻辑都讲明白做到能直接照着复现。1. 复现前必须先想清楚的事Informer到底改了什么1.1 长序列预测的三个痛点Informer提出之前用Transformer做长序列时间序列预测有几个绕不开的老大难问题。第一个是自注意力的二次复杂度。标准Transformer里任意两个位置都要计算注意力权重序列长度L下就是O(L²)的内存和计算开销。序列一旦上探到几千甚至上万显存直接爆炸。我当时用标准Transformer跑过一个长度为3000多的负荷序列单卡显存直接吃到爆这还是在batch size设得很小的情况下。第二个问题是长序列的注意力分布高度稀疏。论文里有个很直观的观察对于大部分query真正有意义的key其实就那么几个但标准注意力机制仍然对每个query做了全量计算大量算力浪费在低价值的位置对上。第三个问题是Decoder自回归推理导致误差累积。传统Transformer解码是一个接一个生成的预测长度一长前一步的误差会一路传导放大。自回归推理本身速度也慢生成1000步就要循环1000次。Informer整篇论文就是在解决这三个问题对应产出了三个核心设计ProbSparse自注意力、自注意力蒸馏、生成式Decoder。复现Informer的完整模型本质上是把这三块在代码层面各就各位再把它们拼成一个整体训练闭环。1.2 三个关键创新每一块都在解决什么ProbSparse自注意力的核心思路是先评估每个query的注意力稀疏程度只让稀疏性最高的Top-u个query参与完整的注意力计算其余query直接用注意力分布的均值近似。这里有个很关键的概念叫稀疏性度量用query和key的概率分布KL散度来定义。计算的时候不需要真正算出所有注意力得分而是用query的均值、最大值和key的均值做一个近似计算得到一个衡量每个query“是否需要特殊关注”的得分公式为# 稀疏性得分近似计算 M max(q * k^T / sqrt(d)) - mean(q * k^T / sqrt(d))得分越高说明这个query的注意力分布越不均匀信息量越大越值得参与全量计算。这样之后每个head只需要选前u c * ln(L)个query参与计算。c是采样因子一般取5L是序列长度。长序列下ln的增长非常慢所以这一步几乎能把复杂度从O(L²)压到O(L ln L)。自注意力蒸馏解决的问题是即使有了稀疏注意力多层Encoder堆叠后特征图尺寸依然很大计算量会随着层数累积。Informer的做法是在每一层Encoder之后对特征做裁剪逐步缩短序列长度有点像CNN里的池化。实现上就是每层后面接一个Conv1d加MaxPool维度减半、通道数加一。生成式Decoder和标准Transformer的Decoder差异最大。Informer的Decoder只需要前向一次就能一次性输出所有预测步的值不需要循环生成。做法是把Decoder输入拼成两部分一部分是已知的真实序列start token一部分是用0填充的占位序列。Decoder通过带掩码的稀疏注意力一次性把整段预测值映射出来训练和推理速度都快得多。1.3 复现前的心理预期不是调包是理解这件事我觉得有必要说在前面。复现Informer不代表pip install informer然后调用一下而是要把模型代码一行行写出来、训练起来、指标跑出来。Informer官方开源了PyTorch实现代码很紧凑模型部分大概900行左右。我建议复现过程中至少做到不改动核心结构的前提下把代码重写一遍搞清楚每个维度变换而不是照抄。我自己复现时遵循的路径是先跑通官方代码再不看源码自己实现一遍最后用官方代码验证自己的结果。这样三轮下来对模型的理解会扎实很多。做完这些你才算真正拥有了这个模型而不是“用过”这个模型。2. 环境准备与数据集选型2.1 依赖版本怎么选Informer对依赖版本没有特别苛刻的要求但有几个版本问题需要注意。官方源码是基于PyTorch 1.x写的复现环境建议这样配Python 3.8或3.93.10也能跑但部分老旧依赖可能报错PyTorch 1.10~2.02.0需要少量API兼容处理后面讲NumPy 1.21Pandas 1.3Matplotlib画loss和预测图用scikit-learn计算指标CUDA版本根据自己的显卡来。我测试时用CUDA 11.3配PyTorch 1.12最稳没有遇到什么诡异问题。如果你用PyTorch 2.0注意LSTM、utils.kl_div这类API问题不大但部分内部函数签名变了后面遇到什么问题再说。提示如果你用的是PyTorch 2.x在导入模型后运行训练可能会出现AttributeError: module torch has no attribute irfft之类的报错这是老代码兼容性的问题后面讲具体怎么修。2.2 ETT数据集Informer复现的标准配置ETTElectricity Transformer Temperature数据集是Informer论文里最常用的基准数据记录的是电力变压器的油温、负载等指标。数据集有四个变体ETTh1、ETTh2、ETTm1、ETTm2。h表示小时粒度m表示15分钟粒度。训练Informer标准做法是用前12个月的数据做训练集后4个月做验证集最后4个月做测试集。如果直接跑官方代码脚本里默认的划分是0.7训练、0.2验证、0.1测试这个可以根据自己的需求改。ETT数据的特点是带有明显周期性和趋势性但信号相对平稳比较适合作为复现验证。我在复现时建议先用ETTh1因为它序列中等、训练速度快、结果稳定方便快速验证模型各个模块是否正确。等模型整体跑通了再上ETTm1或者更长序列的实验。除了ETTInformer官方还支持electricity电力负荷、exchange-rate汇率、traffic流量和weather天气这几个数据集。它们的预处理方式略有不同但架构上无需改动。2.3 数据预处理和时间特征标记Informer源码里对时间特征的处理我认为非常值得学习。它的做法是把时间戳分解成多个特征分量比如月份、日期、星期、小时、分钟等然后和数值特征拼接在一起喂给模型。具体来说源码中的time_features函数会根据数据频率返回不同的特征组合。以小时数据为例时间特征维度是4维月份、日期、星期、小时。如果是15分钟粒度的数据会额外加上分钟维度变成5维。这里有个细节容易忽略预测时如果要用真实未来时间戳来生成时间特征这些特征本身就能为模型提供很强的周期信息。比如预测未来24小时模型知道明天是几点、星期几温度负荷这类强周期数据就能直接受益。复现时一定要把时间特征这块保留完整很多人复现效果差就是因为把时间特征简化掉了。数据标准化的方式也要注意。Informer用的是z-score标准化统计量在训练集上计算然后应用到验证集和测试集。这个操作看似简单但顺序不能错一旦把测试集的统计量混进来就是数据泄漏指标会虚高。# 正确的标准化方式 from sklearn.preprocessing import StandardScaler scaler StandardScaler() train_data scaler.fit_transform(train_data) valid_data scaler.transform(valid_data) test_data scaler.transform(test_data)3. 完整模型核心模块复现实操3.1 Embedding层数据进入网络的第一步Informer的Embedding层由三部分组成数值特征映射、时间特征映射、位置编码。三者相加后作为模型输入。官方源码这个模块写得挺干净的包含了三个子层class DataEmbedding(nn.Module): def __init__(self, c_in, d_model, embed_typetimeF, freqh, dropout0.1): super(DataEmbedding, self).__init__() self.value_embedding nn.Linear(c_in, d_model) self.position_embedding PositionalEmbedding(d_model) self.temporal_embedding TemporalEmbedding(d_model, embed_type, freq) \ if embed_type ! timeF else TimeFeatureEmbedding(d_model, embed_type, freq) self.dropout nn.Dropout(pdropout) def forward(self, x, x_mark): x self.value_embedding(x) self.temporal_embedding(x_mark) self.position_embedding(x) return self.dropout(x)value_embedding把原始数值比如油温、负载从输入维度映射到d_model维度。c_in是特征数量d_model是模型宽度Informer源码里默认512。position_embedding标准的正弦位置编码提供位置信息。temporal_embedding处理时间戳特征这里有两个分支TimeFeatureEmbedding和TemporalEmbedding。前者把时间特征线性映射后者用固定位置编码。通常embed_typetimeF走的是TimeFeatureEmbedding这也是论文里的默认配置。这段代码里有一个容易被忽略的点temporal_embedding只处理x_mark时间戳特征不处理数值特征。所以x进入Encoder时的维度是[batch_size, seq_len, d_model]其中d_model512。在后续的注意力机制计算中这个维度会作为所有子层的维度基准。3.2 ProbSparse自注意力核心中的核心ProbSparse注意力是整个Informer最精华的部分复现时关键是处理采样、稀疏性计算和mask。class ProbAttention(nn.Module): def __init__(self, mask_flagTrue, factor5, scaleNone, attention_dropout0.1, output_attentionFalse): super(ProbAttention, self).__init__() self.factor factor self.scale scale self.mask_flag mask_flag self.output_attention output_attention self.dropout nn.Dropout(attention_dropout) def prob_qk(self, Q, K, sample_k, n_top): # Q: [B, H, L, D] B, H, L, D Q.shape # 采样部分key用采样方式近似计算稀疏性得分 K_expand K.unsqueeze(-3).expand(B, H, L, L, D) index_sample torch.randint(L, (L, sample_k)) K_sample K_expand[:, :, torch.arange(L).unsqueeze(1), index_sample, :] Q_K_sample torch.matmul(Q.unsqueeze(-2), K_sample.transpose(-2, -1)).squeeze(-2) M Q_K_sample.max(-1)[0] - torch.div(Q_K_sample.sum(-1), L) # 选top n_top个query M_top M.topk(n_top, sortedFalse)[1] # 用这些query计算完整注意力 Q_reduce Q[torch.arange(B)[:, None, None], torch.arange(H)[None, :, None], M_top, :] Q_K torch.matmul(Q_reduce, K.transpose(-2, -1)) return Q_K, M_top这里最关键的细节是sample_k的选取。源码中sample_k d_model // factor在d_model512、factor5时sample_k102。而n_top c * ln(L_q)其中c通常取factor5。例如输入长度L96时每个head只取前5 * ln(96) ≈ 23个query参与全量计算。这样做为什么有效因为稀疏性得分M衡量的是query和所有key之间注意力分布的“不均匀程度”。如果某个query对所有key的注意力都差不多那它的M值很低说明它是“低信息量”query用均值近似即可。反之如果某个query只对少数key有高注意力它的M值很高就需要精确计算。全量计算时对于没有选中的query直接用整个注意力矩阵的行均值近似这在实现上等价于“让这类query对所有位置一视同仁”。实际计算时要注意为了保证梯度回传Q_reduce和K的值要直接通过采样索引从原始Q、K中获取不能对Q_K_sample直接做detach()。因为最终输出的注意力矩阵要参与V的加权求和梯度必须从输出反传到Q和K上。得到Q_K之后接下来是缩放和maskD 100 # 论文中设置的缩放常数 Q_K Q_K / math.sqrt(D) if self.mask_flag: attn_mask torch.zeros(L, L, dtypetorch.bool) attn_mask torch.triu(attn_mask, diagonal1) Q_K Q_K.masked_fill(attn_mask, -np.inf) A torch.softmax(Q_K, dim-1) # 将未选中的query按均值填充 A self._fill_with_mean(A, M_top, L) context torch.matmul(A, V)掩码机制这里要特别强调Decoder里的mask是因果掩码保证t时刻只能看到t之前的信息。而Encoder里不需要mask所以mask_flagFalse。源码里通过torch.triu生成上三角掩码把未来位置填成-inf经过softmax后变成0。3.3 编码器自注意力蒸馏和多层堆叠Encoder的结构就是“注意力层 蒸馏层”交替堆叠。每一层Encoder包含一个ProbSparse多头自注意力子层和一个前馈网络子层与标准Transformer类似但区别在于每个注意力层之后跟了一个ConvLayer做长度减半。class ConvLayer(nn.Module): def __init__(self, c_in): super(ConvLayer, self).__init__() self.downConv nn.Conv1d(in_channelsc_in, out_channelsc_in, kernel_size3, padding1, padding_modecircular) self.norm nn.BatchNorm1d(c_in) self.activation nn.ELU() self.maxPool nn.MaxPool1d(kernel_size3, stride2, padding1) def forward(self, x): x self.downConv(x.permute(0, 2, 1)) x self.norm(x) x self.activation(x) x self.maxPool(x) x x.permute(0, 2, 1) return x这里有几个实现细节Conv1d的kernel_size取3padding1并用circular模式保证卷积之后长度不变。随后用MaxPool1d(kernel_size3, stride2, padding1)把序列长度减半。padding1在MaxPool里加上之后长度变化是(L 2 * 1 - 3) / 2 1 L/2——正好完成一半的降采样。ConvLayer放在每层Encoder的输出之后。堆叠过程中序列长度从L降到L/2再降到L/4。如果Encoder有2层第一层输出96第二层输出48总计算量下降一半以上。这个蒸馏设计是Informer能处理超长序列的另一个关键保障。实现Encoder时有一个细节要处理干净第一层Encoder和后续层使用的ConvLayer是不同的。源码中Encoder类的构造函数接收一个conv_layers列表第一个元素是None其余是ConvLayer即在每层后面接卷积蒸馏。3.4 生成式Decoder一次前向输出全部预测Decoder是Informer和标准Transformer差距最大的地方。Informer的Decoder并不逐个生成token而是通过一个“已知部分占位符”的拼接输入一次前向计算出所有预测步的结果。Decoder输入由两部分拼接而成start token和placeholder。具体来说如果我们要预测未来48个时间点Decoder的输入长度是label_len pred_len。其中前label_len是已知的真实序列尾部后pred_len是用0填充的占位部分。class Decoder(nn.Module): def __init__(self, layers, norm_layerNone, projectionNone): super(Decoder, self).__init__() self.layers nn.ModuleList(layers) self.norm norm_layer self.projection projection def forward(self, x, cross, x_maskNone, cross_maskNone): for layer in self.layers: x layer(x, cross, x_maskx_mask, cross_maskcross_mask) if self.norm is not None: x self.norm(x) if self.projection is not None: x self.projection(x) return xDecoder中的注意力机制和标准Transformer的Decoder类似有两个注意力子层Masked ProbSparse自注意力只关注Decoder输入内已有的位置确保预测某个位置时看不到后面的占位符信息。交叉注意力query来自Decoder当前序列key和value来自Encoder的最终输出保持和标准Transformer一样的信息流向从编码器取上下文信息。最后通过一个全连接投影层把d_model映射回预测的目标维度c_out得到最终的预测序列。整个Decoder前向一次输出形状为[batch_size, pred_len, c_out]不需要自回归循环。生成式Decoder带来的收益是巨大的训练时可以直接用真实序列的后半段作为start token让模型学习如何基于历史的真实值快速过渡到预测值。推理时即使没有真实值也可以用Encoder的未来时间特征来填补占位照样一次生成全部预测。4. 训练配置、实验结果与避坑指南4.1 参数配置按数据集分类的经验值Informer的主要超参数有d_model、n_heads、e_layersEncoder层数、d_layersDecoder层数、d_ff前馈隐层维度、dropout、learning_rate、batch_size、seq_len输入长度、label_lenstart token长度、pred_len预测长度。官方实验最常用的一组配置是参数取值说明d_model512模型宽度n_heads8多头数量e_layers3Encoder层数d_layers2Decoder层数d_ff2048前馈隐层维度dropout0.05丢弃率learning_rate0.0001初始学习率batch_size32批大小seq_len96输入序列长度label_len48Decoder已知部分长度pred_len48或96预测长度不同规模的数据集参数要适当调整。我的经验是ETTh用小batch16~32电力负荷数据可以用64d_model在ETT上512够用如果数据量大比如traffic这种可以加到512甚至768。学习率用Adam优化器配0.0001训练时用ReduceLROnPlateau按验证集loss衰减步长2衰减系数0.5。训练步数方面Informer论文里的基准实验一般训练10~20个epoch就收敛了。ETTh1上如果显存允许batch_size32、训练15个epoch单卡V100大约跑15~20分钟。在个人电脑上跑慢一点但也不会超过1小时。如果你的机器显存有限seq_len可以先用48验证整个流程模型跑通了再加大这是调试阶段的通行做法。4.2 训练过程常见的三个坑先说一个我在复现时遇到的最典型的问题loss不降反升。这种现象通常出现在PyTorch 2.x环境下原因是老版源码里的learning_rate默认值设置过低配合新版本优化器行为导致模型更新非常缓慢甚至前期波动掩盖了下降趋势。解决方式是查看训练日志如果loss在一个区间反复波动没有明显下降试着把学习率调大一倍观察几轮再确定是数据问题还是参数问题。第二个坑是显存不足。Informer对长序列非常友好但如果把seq_len设到1000以上batch_size又设得大叠加d_model512和3层Encoder显存还是会吃紧。我的做法是把batch_size降到8或者直接缩短seq_len。注意蒸馏层会把每层长度减半所以显存峰值在输入那一层压缩输入长度往往效果立竿见影。第三个坑是数据泄漏问题。很多人复现的指标比论文还好但一看代码发现标准化时把完整数据集的均值和方差都算进去了或者数据切分时没按时间顺序而是随机切分。时间序列必须严格按照时间顺序切分任何随机洗牌都是在作弊。标准化的统计量只能用训练集计算。这一点没到位后面所有指标都没有参考意义。4.3 复现结果怎么对齐评估指标和可视化Informer官方评估用的是MSE均方误差和MAE平均绝对误差。这两个指标在预测任务里很常用计算方式为def metric(pred, true): mae np.mean(np.abs(pred - true)) mse np.mean((pred - true) ** 2) return mae, mse需要注意测试时Informer会将连续滑窗的预测结果拼接然后统一和真实值对比。窗口与窗口之间是有重叠的拼接时直接对重叠位置取平均。官方源码predict函数里会把preds和trues存下来再统一算指标而不是逐batch算完再平均这一点对结果复现影响很大。我自己在ETTh1上复现的一组结果为seq_len96、pred_len48时MSE约为0.142MAE约为0.243。如果你复现的结果和这个差距在5%以内基本可以认为代码没问题。超出这个范围优先检查数据预处理、时间特征、batch size这些容易出偏差的地方。可视化时官方代码提供visualize函数会画出预测值和真实值的对比曲线。建议在训练完成后选几个代表性的时间窗口画出来看看。模型如果学好了预测曲线应该能跟上真实曲线的整体趋势尤其在周期性明显的时段。如果预测曲线是一条直线贴近均值说明模型没学到周期性信息大概率时间特征没有正确加入。4.4 几个容易忽略的实现细节Informer源码里有一个优化值得注意学习率调度策略用的不是CosineAnnealing而是ReduceLROnPlateau。这个调度器在验证集loss不下降时自动降低学习率对于时间序列这种loss曲线波动较大的场景更稳定。训练日志里如果看到验证集loss长时间不动不妨手动调低学习率重试。另一个细节是模型保存的时机。官方的训练脚本是每个epoch结束都保存一次模型权重文件名带epoch编号和loss值。复现时建议把验证集loss最低的那个epoch单独留一份权重文件因为测试时加载的往往不是最后一个epoch的权重而是历史上val loss最小的那个。这个做法类似early stopping能有效避免过拟合导致测试指标变差。最后就是随机种子的问题。Informer源码里提供了set_seed函数但我强烈建议在训练脚本开头手动固定Python、NumPy、PyTorch的随机种子包括CUDA的随机种子。不固定种子的情况下即便同样的参数两次跑出来的MSE也可能差1%左右对精确对比实验影响不小。5. 踩坑记录与调试经验我在复现Informer过程中踩过不少坑有些问题排查了很久才发现原因这里单独整理出来。如果你复现时遇到类似情况可以少走很多弯路。第一个比较隐蔽的问题是mask的实现。Decoder里做因果掩码时如果用torch.triu(torch.ones(L, L) * -np.inf, 1)生成掩码矩阵要确保对角线是0而不是-inf。如果掩码把对角线也遮住了模型每个位置连自己都看不到训练会非常不稳定loss会一直在高位震荡。检查方式很简单把mask矩阵打印出来看一眼就行。第二个问题是输出维度和真实标签维度不匹配。Informer的预测输出形状是[batch_size, pred_len, c_out]而数据加载器返回的标签batch_y的形状取决于tgt窗口。在训练循环里如果不注意把batch_y截取为pred_len长度计算loss时会得到形状不匹配的错误。源码中有个f_dim0的处理逻辑如果数据有多列特征但只需要预测其中一列务必确认输出只对第一列计算loss。第三个问题是个老生常谈但还是踩了数据增强。时间序列预测任务里我不建议用类似于图像领域的随机裁剪、翻转这类数据增强操作因为时间序列的时域顺序和值域范围都承载着物理含义。Informer的效果提升主要靠模型本身对长序列的建模能力而不是靠人为扰动数据。加了不必要的增强反而可能破坏周期性。此外还有一个细节在torch.utils.data.DataLoader中源码默认设置了drop_lastTrue。这个选项的意思是如果最后一个batch数据不够batch_size就丢弃。这样做是为了保证每个batch形状一致尤其在Decoder的placehoder长度不整除时类似问题很容易出现。如果复现时报错说某个维度长度不一致优先检查drop_last是否设置。6. 多数据集上的复现效果对比代码完整跑通ETTh1之后我在另外几个数据集上也做了验证这里把大致的复现结果和配置列出来供大家做横向对比。数据集seq_lenpred_lenMSE我的复现MAE我的复现ETTh196480.1420.243ETTh196960.1910.285ETTh296480.1780.298ETTm196480.0810.198Electricity96960.2110.326不同数据集上同样的超参数表现差异很大。比如ETTh2的预测难度明显高于ETTh1相同配置下MSE高出不少。ETTm1因为时间粒度更细、数据点更多MSE反而更低。如果你的复现结果和表中数据差异明显可以尝试调整d_model、n_heads、dropout这三个参数它们对结果影响最敏感。电力负荷数据和ETT不同数值范围差异更大标准化后训练效果会好很多。traffic数据集周期性强但整体波动幅度小模型容易学到一个均值预测训练时要格外关注loss是否真的在下降必要时把label_len调长一些让Decoder看到更多真实历史。我当时做完这些对比实验最大的感触是模型结构只决定了效果的上限真正的差距是在数据处理和训练配置这些细节上拉开的。同样一个Informer不同人复现出来的结果差10%以上是很常见的事原因就在于这些细节取舍。最后再说一个可以扩展的方向。Informer源码里有个Exp基类把训练、验证、预测三种模式封装成独立方法。如果你想在这个基础上做实验比如换注意力机制、加外部特征、做多步滚动预测直接在Exp子类上改就行。我后续做一些时序预测对比实验时也都是在这个框架上扩展的。
返回列表