news 2026/9/20 23:46:35

Informer代码逐行解析:数据切片、ProbSparse注意力与生成式解码

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Informer代码逐行解析:数据切片、ProbSparse注意力与生成式解码

简介:针对Informer长序列时间序列预测模型精心整理的逐行注释代码包,面向深度学习研究者、学生与工业界时序建模工程师,适合在理解Transformer类模型、复现实验或做二次改造时使用。代码基于官方Informer2020项目逐模块补充详细中文注释,覆盖数据加载、掩码机制、时间特征构造、概率稀疏注意力、编码器/解码器、训练与评估等核心部分。包体共63个文件,压缩后约62.33MB,以Python源码(py)、编译缓存(pyc)为主,辅以shell训练脚本、CSV数据集、XML/yml等配置及环境说明(Dockerfile、requirements.txt)等,结构与官方仓库保持一致,便于对照原版逐项阅读。目前已有662人学习下载,借助这些注释和配套的notebook示例、预训练权重(pth)、结果图片,能显著降低Informer源码的上手门槛,帮助厘清长序列预测中稀疏注意力等关键实现细节。

1. 为什么 Informer 代码详细注释版比论文正文更值得逐行读

Informer 论文的公式密度不高,但第一次读开源代码的人总是卡在同样的地方:解码器输入为什么是一半真实值加一半零;ProbSparse 的 top_k 到底选的什么;编码器里为什么夹了一层卷积往下采样。这些问题在论文里只有一两句话,代码里却是几十行带坑的实现。这篇博文默认你已经装好 PyTorch、能跑通一个最小训练脚本,直接按数据侧、注意力、结构侧、训练侧、注释侧这条线路,把 Informer 代码里最值得写明白注释的点逐段拆开。适合准备把这份代码改成自己二开版本的工程师和研究生,读完可以直接在自己 fork 里动手补注释、调参数。

2. 数据侧代码:滑动窗口与标签构造是第一个容易卡住的位置

2.1 读懂 Dataset 类里 seq_len、label_len、pred_len 的关系

在 Informer 的公开 PyTorch 实现里,数据侧入口几乎都长成同一个骨架:一个继承torch.utils.data.Dataset的类,在__init__里完成读文件、标准化,在__getitem__里按索引切出编码器输入、解码器输入和标签三段。最先要看清楚的不是读文件的路径,而是构造样本时的三个长度参数。

seq_len = 96 # 编码器回看窗口,用多长一段历史做输入 label_len = 48 # 解码器里真实已知的 token 数,也就是 start token 长度 pred_len = 24 # 要预测的未来步数 s_begin = idx # 当前样本在整条序列上的起点 s_end = s_begin + seq_len r_end = s_end + pred_len # 编码器输入:最近 seq_len 个观测点 seq_x = data_x[s_begin:s_end] # 解码器输入:真实历史末尾 label_len 个点 + 待预测位置占位 0 seq_y = data_y[s_end - label_len:r_end] # 标签:从 seq_len 位置往后 pred_len 个点的真实值 label = data_y[s_begin + seq_len:r_end]

这段代码暴露了 Informer 与普通 Transformer 预测 python 代码的最大差异:解码器不是从零开始生成,而是吃一段“半个已知 + 半个未知”的序列。seq_y里未知位置填 0 只是占位,模型用生成式解码一次把整段输出算出来,而不是像自回归模型那样逐点滚动预测。这 0 不会进入 loss,因为 loss 只计算label对应位置的输出与真实值的差。

三个长度参数之间没有硬性约束,但不能乱配。最常见做法是label_lenpred_len的一半左右,例如96/48/24336/168/336。如果label_len太小,解码器能参考的真实尾部信息太少,预测结果容易整体偏移;如果seq_lenlabel_len差距过大,训练时大量计算浪费在几乎用不到的长历史上。数据侧三个张量的对应关系可以用下面这张表记住:

张量shape构造方式
seq_x[seq_len, C]从起点连续截取一段历史
seq_y[label_len + pred_len, C]真实尾部 label_len 个点 + 补零占位
label[pred_len, C]待预测位置对应的真值

2.2 标准化与逆标准化在代码里的位置

数据侧另一个容易被注释漏掉的细节是标准化。Informer 代码里通常先在训练集上算meanstd,做一次StandardScaler,然后用同一组参数去 transform 验证集和测试集。这一点本身不难,但很容易写错:如果对整个数据集而不是仅对训练集求均值,验证集的信息会被当作已知信息泄漏到训练过程中,测出来的指标会虚高。

from sklearn.preprocessing import StandardScaler scaler = StandardScaler() # 只在训练集上 fit,验证集和测试集只使用同一套参数 transform data_train = scaler.fit_transform(data_train) data_val = scaler.transform(data_val) data_test = scaler.transform(data_test)

模型输出的是标准化之后的值,所以训练时 loss 在标准化空间里算,验证时要把预测结果inverse_transform回原始尺度再计算 MSE、MAE。如果你发现验证集指标数值小得反常,先检查是不是拿标准化后的输出直接算了指标。不少第一次跑代码的人会在这里对着异常小的数字自我怀疑半天,问题往往就是少了一次逆标准化。

2.3 滑动窗口采样时容易忽略的两个边界

采样方式是数据侧最后一个值得写注释的点。公开实现里常见做法是训练集用随机起点:每次__getitem__在合法区间里随机选一个起点,而不是固定按步长滑动。这样做等价于给训练集做了在线增强,模型见到的片段更多;验证和测试则固定从序列开头切,保证结果可复现。

边界上最常见的 bug 有两种。一是s_end + pred_len超出数组长度,处理方式一般是对总长度做整段丢弃,或者直接丢掉尾部不完整的样本。二是索引切错位,最常见的是把seq_y切成[s_end - label_len : s_end]而漏掉后面的pred_len占位段,这会让解码器输出长度比实际标签短,报错时看到 size mismatch 先回来查这一段。

提示:调代码时把三个长度参数打印出来,检查seq_x.shapeseq_y.shapelabel.shape是否分别为[seq_len, C][label_len+pred_len, C][pred_len, C],这一步能挡掉一半的维度错误。

3. ProbSparse 自注意力代码:这里才是 Informer 名字的由来

3.1 普通注意力与 ProbSparse 在代码上的差异

普通 Transformer 的自注意力代码很短,短到让人误以为 Informer 只是把其中一行换成稀疏算子。对照着看会更清楚:

# 普通多头注意力 attn = torch.matmul(Q, K.transpose(-2, -1)) / (d_k ** 0.5) attn = torch.softmax(attn, dim=-1) out = torch.matmul(attn, V) # ProbSparse 的核心改动:不再对所有 query 计算 # 1) 先计算每个 query 的稀疏度得分 # 2) 只保留得分最高的 u 个 query 做完整注意力

Informer 这篇 transformer 预测 python 代码里最核心的注意力模块,思路是“大部分 query 对应的注意力分布是均匀且无信息的,只有少数 query 能产生显著分布”。代码上把它实现成两步:先用一个廉价公式给全部 query 打分,再对 top_u 个 query 做完整注意力计算,其余位置用均值近似。

3.2 稀疏度得分代码:一行 max 减 mean 的来历

论文里的稀疏度度量来自 KL 散度,代码里却常常只有一行max - mean。原因在于 KL 散度展开后有一部分对同一个 query 是常数,比较大小的时候可以约掉,剩下的就是“点积结果的最大值减均值”。这一行是读代码时最容易忽略但最值得注释的地方。

def prob_sparse_attention(Q, K, V, factor=5): B, H, L_Q, L_K, D = *Q.shape[:2], Q.shape[2], K.shape[2], Q.shape[-1] # 1. 全部 query 先与 K 做一次点积,得到 logits 矩阵 scores = torch.matmul(Q, K.transpose(-2, -1)) / (D ** 0.5) # 2. 稀疏度得分近似式:max - mean # 原论文用 KL 散度度量每个 query 与均匀分布的差异, # 公式展开后常数项可约去,代码里退化成这一行。 sparse_score = scores.max(dim=-1).values - scores.mean(dim=-1) # [B, H, L_Q] # 3. 保留 top-u 个 query 索引 u = min(factor * int(math.ceil(math.log(L_K))), L_Q) top_k = sparse_score.topk(u, dim=-1).indices # [B, H, u] # 4. 只对选中的 query 做 softmax 注意力 Q_hat = torch.gather( Q, 2, top_k.unsqueeze(-1).expand(-1, -1, -1, D) ) attn = torch.softmax( torch.matmul(Q_hat, K.transpose(-2, -1)) / (D ** 0.5), dim=-1 ) out_selected = torch.matmul(attn, V) # [B, H, u, D] # 5. 未选中的 query 不参与完整计算,输出用 V 的均值近似 out_full = torch.zeros(B, H, L_Q, D, device=Q.device) out_full.scatter_( 2, top_k.unsqueeze(-1).expand(-1, -1, -1, D), out_selected ) out_full = out_full + V.mean(dim=2, keepdim=True).expand(-1, -1, L_Q, -1) * ( out_full == 0 ) return out_full

这段逻辑里有两个参数值得专门说明。factor控制保留 query 的比例,u = factor * log(L_K)意味着序列越长、保留比例越小,这正解释了论文题目里强调的 O(L log L) 复杂度。math.log是自然对数,多数公开实现对底数不敏感,因为factor可以吸收这个差异;但如果你的序列长度特别短,比如 L_K 只有 8,log(L_K)约等于 2,factor=5u=10超过 L_Q,代码里的min会把它压回来。

3.3 掩码与 factor 参数:两个最容易踩的坑

解码器里的 ProbSparse 要和掩码一起工作。注意力的 mask 是二值的bool张量,shape 通常是[L_Q, L_K],在 softmax 之前把非法位置加-inf。因为只对选中的 query 计算注意力,mask 也必须按top_k做同样的gather,否则选出来的 query 行和 mask 对不上,会出现该掩的位置没掩掉、不该掩的位置被置零的隐蔽错误。

factor 参数的调法比很多人想得简单。固定其他条件,factor 从 5 往 3 调是让模型更“稀疏”,训练更快,但可能丢掉关键 query;往 10 调则更接近普通注意力,计算量上升。数据噪声大、趋势弱的场景建议从 7 起步,预测步数很长如 336 或 720 时保持 5 左右,不要为了提速把 factor 压到 3 以下。

4. 编码器蒸馏与生成式解码器的结构代码:把 forward 函数当维度说明书读

4.1 蒸馏层为什么长这样

Informer 编码器的每个 stage 都由注意力、蒸馏、下采样组成。代码上常见是一个单独的DistilLayer,原因在于它和 Transformer 里常见的 LayerNorm 加 FFN 结构差异太大,值得单独拆开看。

class DistilLayer(nn.Module): """蒸馏层:把序列长度减半,同时保留主要特征""" def __init__(self, d_model): super().__init__() self.conv = nn.Conv1d(d_model, d_model, kernel_size=3, padding=1) self.pool = nn.MaxPool1d(kernel_size=3, stride=2, padding=1) self.act = nn.GELU() def forward(self, x): # x 的初始 shape 是 [B, L, D],但 Conv1d 需要 [B, D, L] x = x.transpose(1, 2) x = self.act(self.conv(x)) x = self.pool(x) # 长度约从 L 减到 L/2 return x.transpose(1, 2)

这里容易误解的是“蒸馏”的对象:它不是在蒸馏注意力分数,而是通过卷积把上一层特征浓缩到更短的序列上。padding=1保证卷积不改变长度,真正下采样由MaxPool1dstride=2完成,序列长度从 L 变成约 L/2,d_model通道数全程不变。多层编码器叠加后,序列长度呈 96 → 48 → 24 → 12 这样递减,模型整体复杂度也因此下降。

4.2 解码器输入:start token 是真实尾巴而不是零

解码器输入构造是 Informer 代码里最反直觉的一段,值得单独注释。训练时的解码器输入由两部分拼接而成:真实历史尾部label_len个点,加上一段长度pred_len的零占位。

# 解码器输入:真实观测尾巴 + 待预测位置补零 dec_inp = torch.cat([ batch_y[:, :label_len, :], # 真实值,作为 start token torch.zeros(batch_size, pred_len, features), # 待预测位置,先用 0 占位 ], dim=1)

start token 这个词在论文里没有展开,代码里其实就是“上一段真实观测”.预测部分填 0 不是让模型输出 0,而是给生成式解码器一个位置占位;loss 只截取解码器输出的最后pred_len步与真实标签比对。注释这一行时,建议直接写“真实尾巴 + 占位符”,后面的人读起来少一层翻译。

4.3 用 shape 把 forward 流程串起来

把 forward 完整看一遍的最快方式,是把每层的输入输出 shape 列成一张表,再对照模型定义检查哪里维度不匹配。以batch=32, seq_len=96, label_len=48, pred_len=24, d_model=512, features=7为例:

张量shape说明
x[32, 96, 7]编码器原始输入
编码器第 1 层输出[32, 48, 512]注意力加蒸馏,长度减半
编码器第 2 层输出[32, 24, 512]继续减半
编码器第 3 层输出[32, 12, 512]最后一级编码特征
dec_inp[32, 72, 7]48 真实 + 24 占位
解码器输出[32, 72, 512]每个位置都对应 d_model 维特征
线性投影[32, 72, 7]映射回原始特征维度
预测结果[32, 24, 7]只取最后 24 步做 loss

这张表的价值在于定位问题。比如报错说注意力里的 Q 和 K 长度不一致,先回头看dec_inp的第二维是不是label_len + pred_len;比如最终 loss 的 tensor shape 对不上,先确认线性层输出维度是c_out而不是d_model。Informer 源码的 forward 函数并不复杂,只是维度流动比常规 Transformer 多跳了几次,shape 表就是那根线。

5. 训练评估与参数调优:用代码确认你关心的几个数字

5.1 训练循环里三个值得注意的小细节

Informer 的训练循环从表面看和常规 PyTorch 模型没有区别,但有三个细节决定长序列训练能不能稳定收敛。

for epoch in range(epochs): model.train() for i, (batch_x, batch_y) in enumerate(train_loader): # 1. 解码器输入在模型内部构造,或由 Dataset 返回 pred = model(batch_x) # [B, label_len+pred_len, C] # 2. loss 只截取预测段,不能拿整段解码器输出计算 loss = criterion(pred[:, -pred_len:, :], batch_y[:, :, :]) # 3. 长序列训练容易梯度爆炸,clip 是标配 torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0) optimizer.zero_grad() loss.backward() optimizer.step()

第一个细节是 loss 截取。解码器输出包含label_len + pred_len个位置,只有最后pred_len个位置和标签对齐,前面的 start token 位置不能算进 loss,否则模型会花大量容量去“复读”已知值。第二个细节是梯度裁剪,clip_grad_norm_的参数max_norm=1.0是常见做法,不加的话长序列在训练后期容易出现单步 loss 突跳。第三个细节是学习率,常见配置是 Adam 加ReduceLROnPlateau,让验证 loss 连续两个 epoch 不下降时自动降学习率。

5.2 验证集指标为什么必须在逆标准化之后计算

验证阶段的代码最容易出现“对比的坐标系不一致”问题。如果训练在标准化空间里进行,验证时预测结果也是标准化值,必须逆变换回原始尺度再算指标,否则 MSE 和 MAE 会比真实值小几个数量级,让人误以为模型已经完美拟合。

# 验证阶段:先预测,再逆标准化,最后算指标 pred = model(batch_x).detach().cpu().numpy() pred = pred[:, -pred_len:, :] pred = scaler.inverse_transform(pred.reshape(-1, C)).reshape(B, pred_len, C) true = scaler.inverse_transform(batch_y.cpu().numpy().reshape(-1, C)) \ .reshape(B, pred_len, C) mse = ((pred - true) ** 2).mean() mae = np.abs(pred - true).mean()

这组代码比模型本身的 forward 更容易被复用。Informer 用的是标准生成式解码,验证时不需要像自回归模型那样滚动预测,一次 forward 就能得到整段预测,所以验证流程很短。计算指标前先检查truepred的 scale 是否一致,这份注释可以写进验证脚本的 docstring 里,避免下次打开代码时重复踩坑。

5.3 参数速查与调优顺序

调 Informer 参数时,很多人一上来就动d_modeldropout,效果却不好。常见做法的调参顺序是:先确定序列长度相关的三个参数,再调模型宽度,最后才动稀疏度相关参数。

参数常见默认调参优先级说明
seq_len96 或 3361编码器回看窗口,决定单样本有效信息量
label_len48 或 1681解码器已知尾巴长度,建议不低于 pred_len 一半
pred_len24 / 48 / 7201预测步长,改它必须同步改标签切分
d_model5122模型宽度,增大不必然提升,显存涨幅最明显
factor53ProbSparse 保留 query 比例,影响速度大于精度
dropout0.054长序列易过拟合,通常最后才动

seq_lenlabel_len是第一优先级的原因是它们直接改变数据切分,而数据切分错了后面所有调参都无法收敛。d_modelfactor这两个参数在验证集上的影响一般是 2% 到 10% 的 MSE 波动,不要指望它们带来质的飞跃。

6. 把“详细注释版”变成可维护的中文注释代码

6.1 注释模板与三个落点

给 Informer 代码补注释时,最常见的做法是先套一个方法注释模板,类似在 IDE 里设置方法注释模板的格式,把每个参数单独拆行说明。

class Informer(nn.Module): """Informer 模型封装(中文注释版) 参数说明: enc_in: 编码器输入特征数,等于数据列数 C d_model: 模型内部统一特征维度,注意力输出都投影到此大小 factor: ProbSparse 注意力保留的 query 比例系数,范围通常 3~10 e_layers: 编码器层数,每层包含一次蒸馏,序列长度减半 d_layers: 解码器层数,控制生成式解码的深度 """ pass

第一个值得落注释的位置是__init__里的参数字段注释。enc_ind_modele_layers这几个字段光看名字容易理解,但label_lenpred_len之间的比例关系、factor 的取值范围,不写注释下次必然要重新翻论文。第二个位置是 forward 函数入口处,把输入张量的 shape 和物理含义用四行注释列清楚,就像前面那张 shape 表一样。第三个位置是 ProbSparse 的max - mean那行,这行是全文最容易被误删的代码,删掉它模型也能跑,但稀疏度度量完全失效。

6.2 验证注释版没有改坏逻辑

如果“详细注释版”是通过重构得到的,写完注释后要跑一遍逻辑比对。做法是保留一份原始脚本,同一批数据、同一个随机种子、同一个 checkpoint,分别跑一遍训练和验证,对比两个版本的 loss 曲线和验证集指标。差异应该为 0;如果有细微差别,先检查是不是改了随机种子或 DataLoader 的采样顺序,再看注意力里的 mask 处理是否一致。

提示:注释不应该只描述“代码在做什么”,而要说明“为什么这么做”。对 Informer 这类结构代码,最保值的一句注释格式是:# 这里保持长度不变,真正下采样在后面 MaxPool 的 stride=2 完成,比单纯的# 卷积有用得多。

最后落一个张量形状断言,把 forward 里最容易出错的维度检查固化下来,用一笔注释说明每个 shape 是从哪一行推导出来的。把张量形状、参数边界和实现缘由钉在代码里,这份注释版才有继续改下去的实际价值。

本文还有配套的精品资源,点击获取

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/20 23:45:57

chezmoi 模板函数 protonPassJSON:从 Proton Pass 提取结构化密钥数据

开发工具CLI配置管理 【免费下载链接】chezmoi Manage your dotfiles across multiple diverse machines, securely. 项目地址: https://gitcode.com/gh_mirrors/ch/chezmoi 点击查看 免费下载 本篇技术指南讲解 chezmoi 点文件管理器中内置的模板函数 protonPassJ…

作者头像 李华
网站建设 2026/9/20 23:44:46

构建可进化的AI编程工作台:从工具堆砌到个人知识操作系统

1. 这不是“AI编程工具合集”,而是一套可生长的个人工作台系统我第一次把“AI编程”当真,是在一个凌晨三点的调试现场——本地跑不通的单元测试,被我喂给刚搭好的本地模型,它不仅指出了mock对象初始化顺序的bug,还顺手…

作者头像 李华
网站建设 2026/9/20 23:44:24

Java 6/7/8历史版本官方下载指南与配置避坑手册

你在帮一个上了年纪的金融项目换开发机,或者刚接手一套十年前写的老系统,大概率会被同一个问题卡住:Java 历史版本从哪里下载?尤其是 Java 6、Java 7、Java 8 这种早就被官方“藏”起来的版本,网上搜出来一堆垃圾站、捆…

作者头像 李华