简介:这份中文注释版代码面向想要深入理解 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源码还停留在“能跑但不懂内部”的工程师和算法研究员。读完后你不仅知道factor=5是什么,更能说出它背后的采样逻辑和经过逐层传播后的维度演变,在参数配置、显存优化和结构复用上达到真正拿来即用的程度。
2. ProbSparse注意力机制:核心代码逐行拆解与采样参数的影响
2.1 从标准自注意力到稀疏度度量:为什么是max-mean
Informer对标准自注意力的改动集中在一点:不再让每个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_factor=5, mask=None): # 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, dtype=torch.float32))) u = max(u, 1) # 随机采样:对每个query随机挑选部分key计算稀疏度分数 index = torch.randint(0, L_K, (B, H, L_Q, u), device=query.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, sorted=False) 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, keepdim=True).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是最值得调的参数之一。以输入长度L=96为例,标准自注意力中每个query要和96个key计算,而ProbSparse只采样5*ln(96)≈23个key来评估稀疏度,最终完整注意力也只对这23个选中query展开。显存占用从O(L^2)降到了O(L·lnL),在L=1000时两者差距接近两个数量级。但如果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_model=512、H=8时每个头的维度是64。工程上如果显存吃紧,优先减少D_model而不是H,因为头数过少会直接削弱多头在不同子空间捕捉依赖的能力,而维度的降低在效果上相对平滑。关键经验:如果输出序列中低频周期成分预测不准,问题往往不在头部数量,而在稀疏度度量对周期型依赖不敏感——低频信号对应的query稀疏度分数普遍偏低,容易被采样丢弃,这时应该加大factor而不动H。
3. 编码器与自注意力蒸馏:Informer在代码里如何裁剪特征层
3.1 多层编码器的堆叠策略与空洞卷积下采样
Informer把Transformer编码器做深了,同时引入了“蒸馏”操作来控制特征图尺寸。编码器由多个EncoderLayer组成,每个层里除了ProbSparse注意力外,还有一个一维空洞卷积+最大池化的蒸馏模块。蒸馏的作用是在时间维度上做降采样:每经过一个蒸馏层,序列长度减半。官方代码里的默认设置是三个EncoderLayer,蒸馏分别在第二层和第三层之后进行,最终序列长度从L降到L/4。这种方式比直接做平均池化好在一维卷积带可学习参数,能够在降采样过程中保留局部时序模式,而stride=1+padding=1的配置确保卷积不会引入额外的时间偏移。
3.1.1 蒸馏层的详细注释代码
import torch.nn as nn class DistillingLayer(nn.Module): def __init__(self, d_model, kernel_size=3, dropout=0.5): super().__init__() # 空洞卷积:dilation=1表示标准卷积,但保留卷积核宽度3的局部感知 self.conv = nn.Conv1d( in_channels=d_model, out_channels=d_model, kernel_size=kernel_size, stride=1, padding=kernel_size // 2, dilation=1, groups=d_model # 深度可分离:每个通道独立卷积,降低参数 ) self.norm = nn.BatchNorm1d(d_model) self.act = nn.GELU() self.maxpool = nn.MaxPool1d(kernel_size=3, stride=2, padding=1) 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]这段代码里两个参数容易忽略:groups=d_model把标准卷积换成了深度可分离卷积,参数量从D×D×k降到D×k,在D=512时参数少两个数量级;MaxPool1d的stride=2配合padding=1保证了长度从L到L/2时能整除,如果原始序列长度是奇数,需要在编码器入口做截断或padding,否则最后一层池化会丢步。实际运行中如果x的长度在池化后不符合预期,可以用assert x.shape[1] == seq_len // 2在debug模式下手动验证。
3.1.2 蒸馏对感受野的扩张效果
空洞卷积在这里的作用是扩大感受野而不增加参数。实践中如果直接叠加两个标准卷积层,第二层卷积只能看到原始序列的局部窗口;而蒸馏模块借助空洞卷积和池化的组合,让高层特征的一个点能够对应输入序列中更大范围的上下文。在序列长度L=336时,三层蒸馏后特征长度变为336/8=42,注意力计算量进一步降低。在实际代码里,蒸馏层的个数必须和序列长度配合:比如输入长度是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, mask=None): # x: [B, L, D] 编码器输入特征 attn_maps = [] for i, (layer, distill) in enumerate(zip(self.layers, self.distillings)): x, attn = layer(x, mask=mask) # 先做ProbSparse注意力 x = distill(x) # 再做蒸馏降采样 attn_maps.append(attn) # 每个蒸馏层都接了BatchNorm,所以这里不再额外加归一化 return x, attn_mapsattn_maps收集每层稀疏注意力的索引和分数,用于可视化分析。调试时如果发现attn_maps里分数分布过于均匀,说明factor设置偏大,模型趋向于平均注意力,失去了稀疏选择的意义;反之如果某个头出现了极端峰值,需要检查是否出现了某个query主导全部注意力的问题。蒸馏层使用BatchNorm在训练和推理间存在差异:BatchNorm在训练时用batch内统计量,推理时用滑动均值;batch_size=1的在线推理场景下,滑动均值可能因统计量积累不足而出现偏差,建议在加载预训练权重时打印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], dim=1)拼接后解码器输入长度是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_len=48、pred_len=24,则掩码的前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_norm=1.0) optimizer.step()梯度裁剪在长序列预测里几乎是必须的,Informer的深度编码器加上多层蒸馏,梯度在反向传播时经过多个池化层后数值容易爆炸。如果loss曲线在训练初期出现“断崖式”上升再到NaN,max_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_len=96, label_len=48, pred_len=24,train/valid/test按0.7/0.1/0.2切分。不同pred_len下需要调factor和蒸馏层数,如下表:
| pred_len | 推荐 enc_len | 推荐 label_len | factor | 蒸馏层数 |
|---|---|---|---|---|
| 24 | 96 | 48 | 5 | 3 |
| 48 | 192 | 96 | 5 | 3 |
| 96 | 336 | 96 | 6 | 3 |
| 168 | 512 | 168 | 7 | 2 |
预测长度比较大时,factor从5升到6或7,因为更长的预测需要解码器从编码特征中检索更多有效信息,稀疏采样的覆盖范围必须扩大。蒸馏层数从3降到2,是因为enc_len=512经过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_in=D, dec_in=D, d_model=D, d_ff=2048, n_heads=H, e_layers=3, d_layers=2, factor=5, pred_len=24, label_len=48, dropout=0.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-5 | softmax维度错误、mask位置偏移 |
| 蒸馏层输出 | 1e-4 | 池化padding方式不同、卷积权重初始化差异 |
| 最终预测值 | 1e-3 | 解码器输入拼接顺序不一致、归一化统计量不同 |
误差阈值定得太严没有意义,浮点计算的累加顺序本身就会带来微小差异,所以1e-6的阈值只适合定位逻辑错误,不适合做精细比对。如果发现预测值差异在1e-2量级但方向一致(比如整体偏大或偏小),多半是归一化阶段的均值统计差异,不是模型结构问题。
最后提一个我在对照注释读代码时的习惯:凡是在注释里写了“为了XXX而设计”的地方,都手动改掉再跑一次训练。比如把factor=5改成factor=0(即所有key都参与),看loss和显存的变化,你才能真正理解稀疏采样带来的效率边界在哪里。把这当作验证手段而不是调参建议,读一遍注释版的收益会超过直接跑通三个模型。
本文还有配套的精品资源,点击获取