ARTICLE DETAIL

资讯详情

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

Informer源码详解:ProbSparse注意力与长序列预测实践

Informer源码详解:ProbSparse注意力与长序列预测实践 简介这份中文注释版代码面向想要深入理解 Informer 模型的读者。Informer 是面向长序列时间序列预测的高效 Transformer原始开源代码使用 PyTorch 实现结构紧凑但阅读门槛较高代码对数据加载、特征编码、稀疏注意力、编码器/解码器、训练和预测等关键环节进行了逐行注释可帮助研究者、算法工程师和学生快速建立从理论到代码的映射。压缩包约 62.33MB共 63 个文件以 Python 源码为主17 个 .py还包含 16 个编译缓存、7 个工程配置、5 张原理示意图、4 个 Shell 训练脚本、4 个数据文件以及运行环境与依赖说明等辅助文件目录结构基本保留官方工程划分便于对照原始仓库学习。目前已有 662 人学习使用借助注释和示意图可以明显缩短源码反复排查的时间。这份代码也适合作为毕设复现、论文实验或二次开发的基础工程注释风格清晰对关键参数与设计意图有明确提示。1. Informer代码详细注释版从源码读懂长序列预测的每一个算子如果你只把Informer当成一个能跑通的模型那就错过了一份难得的时间序列工程教材。Informer的核心价值不在于“比Transformer快多少”而在于它针对长序列预测LSTF在复杂度、长期依赖、解码延迟三个维度上做了系统性改造ProbSparse自注意力替代标准点积注意力、自注意力蒸馏压缩特征层、生成式解码器用一步前向替代逐步递归。这份代码详细注释版要做的正是把这些改造逐行拆开让你清楚每条张量从加载到输出的形状变化、每个超参在损失函数和内存占用上的实际影响。适用对象是已经跑过Transformer或LSTM预测代码、但对Informer源码还停留在“能跑但不懂内部”的工程师和算法研究员。读完后你不仅知道factor5是什么更能说出它背后的采样逻辑和经过逐层传播后的维度演变在参数配置、显存优化和结构复用上达到真正拿来即用的程度。2. ProbSparse注意力机制核心代码逐行拆解与采样参数的影响2.1 从标准自注意力到稀疏度度量为什么是max-meanInformer对标准自注意力的改动集中在一点不再让每个query与所有key做点积而是先通过稀疏度评估找出“活跃”的query只让这些query参与完整注意力计算。标准注意力的第i个query对所有key的注意力分布是p(k_j|q_i)Informer用KL散度来判断该分布与均匀分布的差异差异越大说明这个query的选择性越强越值得保留完整计算。代码实现中这个度量被简化成了max(q·k^T) - mean(q·k^T)的形式因为原始KL散度需要逐query遍历全部key计算对数求和代价太高这个近似度量在数学上做了放缩处理保留了排序能力但把复杂度降到了O(L_Q log L_K)。2.1.1 稀疏度评估的PyTorch实现样例import torch import torch.nn.functional as F def prob_sparse_attention(query, key, value, sampling_factor5, maskNone): # query/key/value: [B, H, L, D] B, H, L_Q, D query.shape _, _, L_K, _ key.shape # 计算采样数量控制参与完整注意力的query子集大小 u int(sampling_factor * torch.log(torch.tensor(L_K, dtypetorch.float32))) u max(u, 1) # 随机采样对每个query随机挑选部分key计算稀疏度分数 index torch.randint(0, L_K, (B, H, L_Q, u), devicequery.device) key_sample key.gather(-2, index.expand(-1, -1, -1, D)) q_k torch.matmul(query, key_sample.transpose(-2, -1)) # 稀疏度近似度量max - mean替代完整KL散度 m q_k.max(-1).values - q_k.mean(-1) # [B, H, L_Q] # 选出Top-u个高稀疏度query _, top_indices torch.topk(m, u, dim-1, sortedFalse) top_indices top_indices.unsqueeze(-1).expand(-1, -1, -1, D) query_selected query.gather(-2, top_indices) # 仅对选中query做完整注意力 scores torch.matmul(query_selected, key.transpose(-2, -1)) / (D ** 0.5) if mask is not None: scores scores.masked_fill(mask 0, -1e9) attn F.softmax(scores, dim-1) context_selected torch.matmul(attn, value) # [B, H, u, D] # 将未选中的query用value的均值填充保持输出形状不变 context value.mean(-2, keepdimTrue).expand(B, H, L_Q, D).clone() context.scatter_(-2, top_indices, context_selected) return context这段代码还原了ProbSparse注意力的核心数据流。sampling_factor控制采样数量官方默认是5对应公式u factor * ln(L_K)。topk操作选出的top_indices在后续scatter_回填时必须严格保持索引一致否则输出张量里的位置会错位。注释里特意标出了value.mean(-2)这个操作未参与完整注意力的query直接复用所有value的均值这是保证输出形状不变的关键trick。2.1.2 factor参数对显存和精度的实际影响factor是最值得调的参数之一。以输入长度L96为例标准自注意力中每个query要和96个key计算而ProbSparse只采样5*ln(96)≈23个key来评估稀疏度最终完整注意力也只对这23个选中query展开。显存占用从O(L^2)降到了O(L·lnL)在L1000时两者差距接近两个数量级。但如果factor设得过大比如到20采样数量会超过L_K的一半稀疏度评估本身就失去意义代码中的randint会在L_Q·factor·ln(L_K)远大于L_Q·L_K时产生重复索引相当于变相做了近似全量计算设得过小小于3则可能漏掉真正活跃的query预测曲线会出现明显的滞后和毛刺。我的建议是从5起步在验证集上观察attention分布的熵如果熵值普遍偏高说明采样过于均匀适当减小factor。提示不同版本的Informer代码在实现ProbSparse时可能存在细节差异。部分实现会用torch.randperm代替randint来避免重复采样但前者在L_K很大时的张量分配开销不小工程上我更倾向保留randint因为在稀疏度评估阶段重复索引的影响可控。2.2 多头稀疏注意力的拼接与残差连接ProbSparse注意力在实现上是多头并行的每个头独立做上述采样和计算。拼接时要注意的是context张量的维度恢复顺序[B, H, L, D]要先transpose(1,2)变回[B, L, H, D]再reshape(B, L, H*D)然后经过线性投影。这一步看似简单但对新手而言最容易在这里出错——reshape和view在张量内存不连续时会直接报错需要用contiguous()做一次内存整理。残差连接放在投影之后out layer_norm(x dropout(proj(context)))这里有个容易忽略的细节x是注意力层输入形状为[B, L, D_model]而context经过多头拼接后也是这个形状两者直接相加没有问题但如果你在某个实现里看到先归一化再进注意力Pre-LN残差路径上就不用再接LayerNorm避免重复归一化导致梯度不稳定。2.2.1 多头注意力的维度检查清单张量形状说明query/key/value 输入[B, L, D_model]编码器输入D_model为特征维度多头拆分后[B, H, L, D_head]D_head D_model / H稀疏度分数[B, H, L_Q, u]u为采样后的key数量context_selected[B, H, u, D_head]仅选中query的注意力输出回填后context[B, L, D_model]未选中query用value均值填充多头数量H要和D_model整除对应起来比如D_model512、H8时每个头的维度是64。工程上如果显存吃紧优先减少D_model而不是H因为头数过少会直接削弱多头在不同子空间捕捉依赖的能力而维度的降低在效果上相对平滑。关键经验如果输出序列中低频周期成分预测不准问题往往不在头部数量而在稀疏度度量对周期型依赖不敏感——低频信号对应的query稀疏度分数普遍偏低容易被采样丢弃这时应该加大factor而不动H。3. 编码器与自注意力蒸馏Informer在代码里如何裁剪特征层3.1 多层编码器的堆叠策略与空洞卷积下采样Informer把Transformer编码器做深了同时引入了“蒸馏”操作来控制特征图尺寸。编码器由多个EncoderLayer组成每个层里除了ProbSparse注意力外还有一个一维空洞卷积最大池化的蒸馏模块。蒸馏的作用是在时间维度上做降采样每经过一个蒸馏层序列长度减半。官方代码里的默认设置是三个EncoderLayer蒸馏分别在第二层和第三层之后进行最终序列长度从L降到L/4。这种方式比直接做平均池化好在一维卷积带可学习参数能够在降采样过程中保留局部时序模式而stride1padding1的配置确保卷积不会引入额外的时间偏移。3.1.1 蒸馏层的详细注释代码import torch.nn as nn class DistillingLayer(nn.Module): def __init__(self, d_model, kernel_size3, dropout0.5): super().__init__() # 空洞卷积dilation1表示标准卷积但保留卷积核宽度3的局部感知 self.conv nn.Conv1d( in_channelsd_model, out_channelsd_model, kernel_sizekernel_size, stride1, paddingkernel_size // 2, dilation1, groupsd_model # 深度可分离每个通道独立卷积降低参数 ) self.norm nn.BatchNorm1d(d_model) self.act nn.GELU() self.maxpool nn.MaxPool1d(kernel_size3, stride2, padding1) self.dropout nn.Dropout(dropout) def forward(self, x): # x: [B, L, D] - 转成 [B, D, L] 供Conv1d使用 x x.transpose(1, 2) x self.conv(x) x self.norm(x) x self.act(x) x self.maxpool(x) # 长度从L变为L/2 x self.dropout(x) return x.transpose(1, 2) # 转回 [B, L/2, D]这段代码里两个参数容易忽略groupsd_model把标准卷积换成了深度可分离卷积参数量从D×D×k降到D×k在D512时参数少两个数量级MaxPool1d的stride2配合padding1保证了长度从L到L/2时能整除如果原始序列长度是奇数需要在编码器入口做截断或padding否则最后一层池化会丢步。实际运行中如果x的长度在池化后不符合预期可以用assert x.shape[1] seq_len // 2在debug模式下手动验证。3.1.2 蒸馏对感受野的扩张效果空洞卷积在这里的作用是扩大感受野而不增加参数。实践中如果直接叠加两个标准卷积层第二层卷积只能看到原始序列的局部窗口而蒸馏模块借助空洞卷积和池化的组合让高层特征的一个点能够对应输入序列中更大范围的上下文。在序列长度L336时三层蒸馏后特征长度变为336/842注意力计算量进一步降低。在实际代码里蒸馏层的个数必须和序列长度配合比如输入长度是96三层蒸馏后变12还能支撑后续处理但如果输入只有48三层蒸馏后剩6特征过于稠密信息损失严重需要减少一个蒸馏层或调整池化stride。这是Informer代码注释版里最常被忽视的约束条件。3.2 编码器输出的特征聚合与全连接映射编码器最后一层的输出形状是[B, L/4, D_model]。如果直接把整个特征图交给解码器解码器的交叉注意力要处理的时间步仍然不少。官方实现的做法是只取特征图的最后一个时间步x[:, -1, :]然后经过一个全连接层把D_model映射到预测长度对应的维度。这个设计很直接概率稀疏注意力已经在各层内做了充分的时间依赖建模最后一步用“最后一个状态代表整个序列”虽然粗暴但结合蒸馏后特征中已经包含多尺度信息效果上完全够用。3.2.1 编码器前向传播的完整流程class Encoder(nn.Module): def __init__(self, layers, distilling_layers): super().__init__() self.layers layers # 注意力层列表 self.distillings distilling_layers # 蒸馏层列表与注意力层交替 def forward(self, x, maskNone): # x: [B, L, D] 编码器输入特征 attn_maps [] for i, (layer, distill) in enumerate(zip(self.layers, self.distillings)): x, attn layer(x, maskmask) # 先做ProbSparse注意力 x distill(x) # 再做蒸馏降采样 attn_maps.append(attn) # 每个蒸馏层都接了BatchNorm所以这里不再额外加归一化 return x, attn_mapsattn_maps收集每层稀疏注意力的索引和分数用于可视化分析。调试时如果发现attn_maps里分数分布过于均匀说明factor设置偏大模型趋向于平均注意力失去了稀疏选择的意义反之如果某个头出现了极端峰值需要检查是否出现了某个query主导全部注意力的问题。蒸馏层使用BatchNorm在训练和推理间存在差异BatchNorm在训练时用batch内统计量推理时用滑动均值batch_size1的在线推理场景下滑动均值可能因统计量积累不足而出现偏差建议在加载预训练权重时打印running_mean的数值范围若异常则改用LayerNorm。4. 生成式解码器与长时间序列预测训练阶段和推理阶段的代码差异4.1 解码器输入拼接start token与预测位置的动态掩码Informer解码器采用生成式结构输入由三部分拼接而成序列最后一段已知值start token、今天之前对应周期的已知值如果做的是周维度预测就用上周同时段数据、以及需要预测的位置用0填充。这个拼接发生在data_loader的batchify函数里代码注释版里一般写成# seq_x: 原始序列 [B, L_total, D] # label_len: 解码器已知序列长度默认48 # pred_len: 预测长度默认24/48/96 enc_input seq_x[:, :enc_in_len, :] # 编码器输入取序列前半段 dec_input torch.zeros_like(seq_x[:, -pred_len:, :]) # 预测位置初值 dec_input torch.cat([seq_x[:, label_len:enc_in_len, :], dec_input], dim1)拼接后解码器输入长度是label_len pred_len。生成式解码器中预测位置的0填充在训练阶段会被mask掉——通过注意力掩码让解码器只能看到已给的真实片段不能看到未来位置。这个掩码是下三角矩阵的变体和标准Transformer解码器掩码的关键区别在于Informer的掩码还要覆盖掉预测位置内部的自注意力防止解码器在训练时看到预测位置的“答案”。4.1.1 掩码实现的边界条件def generate_mask(dec_len, pred_len): # 掩码矩阵形状 [dec_len, dec_len] mask torch.ones(dec_len, dec_len).tril() # 下三角为1 # 预测区域全部置0该区域内部不允许互相看到 mask[:, -pred_len:] 0 # 已知区域可以看到所有已知区域但不能看预测区 mask[:label_len, :label_len] 1 return mask实际操作中label_len和pred_len的边界最容易写错如果label_len48、pred_len24则掩码的前48行前48列为下三角实际应该全1因为是已知段后24行所有位置为0。很多魔改版本想要让解码器“自回归式生成”会把后24行改成下三角形式但这样做导致训练和推理的数据流不一致——训练阶段解码器能看到未来位置而推理时这些位置根本没有输入模型在测试集上的误差会急剧放大。如果必须改造要保证训练、推理使用同一套掩码逻辑。4.2 训练损失与推理流程一条代码路径里的两种模式Informer的损失函数用的是MSE有的版本也会加上MAE做组合loss mse_loss 0.5 * mae_loss。训练时解码器输出形状是[B, label_len pred_len, D]由于预测位置在输入时是0填充损失只计算后半段loss criterion(output[:, -pred_len:, :], target[:, -pred_len:, :])。这里有版本在损失计算前会对输出做inverse标准化把归一化后的预测还原成原始量纲再算误差这会让loss数值反映真实量纲便于监控但反向传播时梯度会经过标准化的逆变换梯度尺度可能被放大建议用单个batch实测梯度的norm如果大于10则应在inverse前截断。推理阶段和训练共用同一个前向函数。关键区别是推理时enc_input使用最新可用的完整序列dec_input的已知段取自序列最后label_len个点预测段全0。代码注释版中常见的predict函数会把model.eval()和with torch.no_grad()包在一起但要注意BatchNorm和Dropout的行为差异——Dropout在eval模式下自动关闭BatchNorm则继续用滑动均值这两者在推理时都不会重新计算。如果你的Informer代码在训练好之后做推理时结果明显异常优先检查模型是否真的切到了eval模式而不是怀疑参数出了问题。4.2.1 训练循环中梯度裁剪的作用optimizer.zero_grad() output model(enc_input, dec_input) loss criterion(output[:, -pred_len:, :], batch_y[:, -pred_len:, :]) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0) optimizer.step()梯度裁剪在长序列预测里几乎是必须的Informer的深度编码器加上多层蒸馏梯度在反向传播时经过多个池化层后数值容易爆炸。如果loss曲线在训练初期出现“断崖式”上升再到NaNmax_norm设成0.5就能解决但如果设得太小低于0.1模型收敛会极慢loss下降变成一条平缓直线。我通常会先不裁剪跑5个epoch观察梯度norm的分布如果中位数超过10再加裁剪。提示推理阶段不要调用torch.no_grad()后就直接把模型输出当成最终预测。Informer的decoder输入中的start token部分在官方实现里用的是label_len长度的真实值有部分优化版本会在推理时用前一次预测结果替换start token实现多步滚动预测。在代码注释版中这两种模式通常由参数is_training区分训练、验证、测试三个阶段分别设置不要混淆。5. 数据维度与超参数对应关系从数据加载到模型配置的代码注释5.1 输入特征归一化与反归一化Informer代码里有标准的归一化处理训练集上计算mean和std验证集和测试集沿用训练集的统计量防止数据泄漏。inverse_transform函数在推理结束后调用还原预测值到原始单位。这一组操作的顺序不能反先做归一化再划分数据集还是先划分再归一化看起来只差两行代码但后者会保留时间维度上的分布漂移信息在非平稳序列上能提升5%左右的预测精度。代码注释版里一般建议先切分再计算统计量而不是对全量数据做归一化再切。5.1.1 数据加载器中的维度匹配from torch.utils.data import Dataset, DataLoader class TimeSeriesDataset(Dataset): def __init__(self, data, enc_len, dec_len, label_len, pred_len): self.data data self.enc_len enc_len # 编码器输入长度如96 self.dec_len dec_len # 解码器输入长度 label_len pred_len self.label_len label_len # 已知token长度如48 self.pred_len pred_len # 预测长度如24 def __len__(self): return len(self.data) - self.enc_len - self.pred_len 1 def __getitem__(self, idx): s_begin idx s_end s_begin self.enc_len r_end s_end self.pred_len # 编码器输入从s_begin到s_end enc_input self.data[s_begin:s_end] # 解码器输入从s_end - label_len到r_end预测部分自动0填充 dec_input self.data[s_end - self.label_len:r_end] # 预测目标从s_end到r_end target self.data[s_end:r_end] return enc_input, dec_input, target__len__的计算是滑动窗口不重叠时最容易出错的地方。比如总长1000、编码长度96、预测长度24窗口数应该是1000 - 96 - 24 1 881如果你的代码用len(data) // (enc_len pred_len)来做会直接丢掉末尾的完整序列预测段和真实值的对齐也会错位。另外一个常见问题是dec_input里如果包含NaN或inf数据源常见的缺失值填充会把训练loss变成NaN且这个错误不会在loss曲线早期暴露而是在某个batch突然爆炸。建议在__getitem__里加一行类型检查代码assert not torch.isnan(enc_input).any()定位到具体哪条样本出了问题。5.2 编码长度、预测长度和label_len的组合建议参数组合没有固定的最优值但有几个边界条件值得记录enc_len建议取预测长度的4到8倍比如预测24步编码96~192步能捕捉到足够的周期上下文label_len设置为预测长度的一半到两倍之间过短时解码器的start token信息不足生成的序列起点误差偏大过长时解码器的输入张量变大注意力计算量上升但精度增益有限。我从实践角度给的默认启动配置是enc_len96, label_len48, pred_len24train/valid/test按0.7/0.1/0.2切分。不同pred_len下需要调factor和蒸馏层数如下表pred_len推荐 enc_len推荐 label_lenfactor蒸馏层数2496485348192965396336966316851216872预测长度比较大时factor从5升到6或7因为更长的预测需要解码器从编码特征中检索更多有效信息稀疏采样的覆盖范围必须扩大。蒸馏层数从3降到2是因为enc_len512经过3层蒸馏后变64特征长度已经足够紧凑再加一层会压缩到32过度丢细节。6. 验证代码注释版的正确性让模型复现你注释过的每一行6.1 用单元测试锁定张量形状防止注释和实际行为脱节拿到一份Informer代码详细注释版后如何验证注释和代码真的对应我的做法是先写一组形状断言把每层输入输出张量的shape固化下来。这样既能确认注释里的描述与实际运行一致又能在后期调参时及时发现维度变化带来的连锁影响。def test_informer_shapes(): B, L, D, H 4, 96, 512, 8 x torch.randn(B, L, D) model Informer( enc_inD, dec_inD, d_modelD, d_ff2048, n_headsH, e_layers3, d_layers2, factor5, pred_len24, label_len48, dropout0.1 ) enc_input torch.randn(B, L, D) dec_input torch.randn(B, 48 24, D) out model(enc_input, dec_input) assert out.shape (B, 48 24, D), f解码器输出形状异常: {out.shape} print(形状验证通过编码器输出、蒸馏层输出、解码器输出全部符合预期)如果out.shape报错先检查e_layers和蒸馏层的数量配合关系再用torchsummary库或者手动打印各层forward的shape逐层定位是哪个模块改变了维度。注释版的价值就在于你不需要从零读源码但必须让注释里描述的shape和模型实际跑出来的结果完全对齐否则注释就没有意义。6.2 用固定随机种子复现一个batch的数值结果验证注释是否准确的第二个方法是数值复现固定torch.manual_seed(0)取同一个batch的数据分别用原始模型和注释版模型各跑一次比较两者输出差异是否在一个极小的容差范围内比如1e-6。如果注释版的代码里改动了任何一个算子的实现细节比如把torch.matmul换成了torch.einsum或者把F.softmax换成了手动除温度后再softmax数值误差会放大到1e-3以上。这个测试在模型蒸馏、注意力掩码、梯度裁剪三个位置尤其有效因为这些地方的等价变换最容易出错。6.2.1 实际验证的误差判定标准比对对象误差阈值可能原因注意力输出1e-5softmax维度错误、mask位置偏移蒸馏层输出1e-4池化padding方式不同、卷积权重初始化差异最终预测值1e-3解码器输入拼接顺序不一致、归一化统计量不同误差阈值定得太严没有意义浮点计算的累加顺序本身就会带来微小差异所以1e-6的阈值只适合定位逻辑错误不适合做精细比对。如果发现预测值差异在1e-2量级但方向一致比如整体偏大或偏小多半是归一化阶段的均值统计差异不是模型结构问题。最后提一个我在对照注释读代码时的习惯凡是在注释里写了“为了XXX而设计”的地方都手动改掉再跑一次训练。比如把factor5改成factor0即所有key都参与看loss和显存的变化你才能真正理解稀疏采样带来的效率边界在哪里。把这当作验证手段而不是调参建议读一遍注释版的收益会超过直接跑通三个模型。本文还有配套的精品资源点击获取
返回列表