尧图网络科技YAOTU DIGITAL 获取报价
获取报价
首页 / 资讯中心 / 文章详情

Informer代码注释版:ProbSparse注意力与蒸馏层实战调参指南

发布时间:2026/9/26 4:32:40

资讯中心
01
ARTICLE

Informer代码注释版:ProbSparse注意力与蒸馏层实战调参指南

Informer代码注释版:ProbSparse注意力与蒸馏层实战调参指南
简介面向深度学习时序预测方向的Informer模型代码逐行注释资源由CSDN作者qq_40957277整理适合想从源码层面理解Informer原理的研究者、学生或开发者使用。资源包共包含63个文件压缩后约62.33MB其中以Python脚本17个py为核心覆盖模型定义、注意力机制、编码器解码器、数据加载与训练评估等模块同时附带sh运行脚本、csv数据集、环境配置文件、模型权重pth以及说明文档便于直接跑通实验。内容预览显示其目录结构完整包含ETT数据集、实验模块exp、模型models、工具utils、脚本scripts等标准工程组织注释细致到逐行可帮助使用者省去对照论文阅读源码的精力快速掌握长序列时间序列预测的实现细节。目前该资源已有662人学习对希望深入Informer内部机制或进行二次开发的研究者具有不错的参考价值。1. Informer代码详细注释版从跑通到改得动Informer代码详细注释版不是把官方仓库的注释补全那么简单它要回答一个实际问题当你把Informer跑起来之后某个张量为什么是这个形状、某个mask为什么这样写、解码器为什么可以一步生成出了问题去哪里改。很多第一次复现Informer的人卡在ProbSparse Attention的采样和蒸馏层的长度变化上这两处在论文里只有两页公式在代码里却要处理五六个张量维度。这篇笔记按我自己的阅读习惯先把代码骨架理清再逐个拆核心模块最后给出能直接抄的训练命令和几个最容易掉进去的坑。适合正在做长序列时间序列预测、想用Informer跑自己的数据又不想只把代码当黑匣子调参的读者。2. 源码骨架从main_informer.py到数据加载器的调用链Informer的官方实现本身是一个实验框架不是库。它把训练、验证、测试的逻辑都写在exp/exp_main.py模型定义在models/数据处理在data_provider/。注释版的价值在于把main_informer.py里的每一个 argparse 参数和exp_main.py里真正使用它的地方对应起来。不然你改一个--d_model很多时候根本不知道影响哪个张量。2.1 入口如何初始化参数解析、模型构造和数据加载打开main_informer.py注释版通常会把 argparse 参数分成数据相关、模型结构、训练策略三类。下面这段是后面所有实验的源头。# main_informer.py 关键片段详细注释版 import argparse parser argparse.ArgumentParser(descriptionInformer) # ---------- 数据相关参数 ---------- parser.add_argument(--data, typestr, defaultETTh1, help数据集名称决定 data_provider 加载哪个文件) parser.add_argument(--features, typestr, defaultM, helpM: 多变量预测多变量S: 单变量预测单变量MS: 多变量预测单变量) parser.add_argument(--seq_len, typeint, default96, help输入历史窗口长度) parser.add_argument(--label_len, typeint, default48, help解码器里真实序列拼接的长度) parser.add_argument(--pred_len, typeint, default96, help预测长度) # ---------- 模型结构参数 ---------- parser.add_argument(--d_model, typeint, default512, help编码器/解码器内部特征维度) parser.add_argument(--n_heads, typeint, default8, help多头注意力的头数) parser.add_argument(--e_layers, typeint, default3, help编码器层数) parser.add_argument(--d_layers, typeint, default2, help解码器层数) parser.add_argument(--distil, typebool, defaultTrue, help是否使用编码器里的自注意力蒸馏) parser.add_argument(--attn, typestr, defaultprob, helpprob 或 full选择稀疏注意力或全量注意力) parser.add_argument(--factor, typeint, default5, helpProbSparse 采样因子越大越接近全量注意力) # ---------- 训练策略参数 ---------- parser.add_argument(--learning_rate, typefloat, default0.0001, help学习率) parser.add_argument(--train_epochs, typeint, default6, help训练轮数) parser.add_argument(--batch_size, typeint, default32, help批大小) parser.add_argument(--patience, typeint, default3, help早停轮数)argparse 只是声明参数真正的解析发生在main函数里调用exp_main.Exp之后。exp_main.py里的_build_model会按args.model选择模型然后调用_get_data加载数据。这里有一个容易看漏的点--data决定数据集的目录和文件名--freq决定时间特征编码的粒度很多人只改了--data没改--freq结果时间戳解析失败。模型构造的简化代码在exp/exp_main.py中长这样# exp/exp_main.py 中 _build_model 的简化逻辑 def _build_model(self): model_dict { Informer: Informer, Autoformer: Autoformer, } model model_dict[self.args.model].Model( self.args, self.args.enc_in, # 编码器输入特征数通常是数据列数 self.args.dec_in, # 解码器输入特征数通常等于 c_out 时间特征维度 self.args.c_out, # 输出特征数多变量预测时等于要预测的列数 self.args.d_model, self.args.n_heads, self.args.e_layers, self.args.d_layers, self.args.distil, self.args.dropout, self.args.attn, self.args.factor, ) return modelenc_in、dec_in、c_out这三个值非常容易填错。enc_in是数据里参与预测的输入特征数如果数据集有 7 列--enc_in 7。c_out是最终预测的目标维度featuresM时通常等于列数featuresS时是 1。dec_in在官方实现里是解码器输入维度它要能容纳c_out和拼接的时间特征所以常见的注释版会提醒你dec_in不能只是为了对齐enc_in而随便填否则前向传播会在 embedding 附近报维度错。2.2 数据加载器时间戳、切分与标准化Informer 的数据加载逻辑在data_provider/data_loader.py。它的职责不只是读文件还有一个特别容易忽略的动作在训练集上做标准化然后把这个标准化器保存下来给验证集和测试集用。如果你自己写数据加载器最容易犯的错是在整个数据集上做 StandardScaler这会造成信息泄漏验证集和测试集的 loss 会虚低等到上线才翻车。# data_provider/data_loader.py 中 Dataset_Custom 的 __read_data__ def __read_data__(self): df_raw pd.read_csv(self.root_path self.data_path) # 确保日期列能被解析 df_raw[self.date_col] pd.to_datetime(df_raw[self.date_col]) # 常见切分比例训练 70%验证 10%测试 20% train_size int(len(df_raw) * 0.7) val_size int(len(df_raw) * 0.1) test_size int(len(df_raw) * 0.2) # 对训练集做标准化 self.scaler StandardScaler() self.scaler.fit(df_raw[train_start:train_end][self.target_cols]) # 验证集和测试集都用训练集的 mean/std df_raw[train_start:train_end][self.target_cols] self.scaler.transform(...) df_raw[val_start:val_end][self.target_cols] self.scaler.transform(...) df_raw[test_start:test_end][self.target_cols] self.scaler.transform(...)标准化之后__getitem__会按照seq_len、label_len、pred_len切窗口。注释版通常会把这段窗口关系画出来因为它是理解解码器输入的关键。# data_provider/data_loader.py 中的 __getitem__ def __getitem__(self, index): # 编码器输入窗口 s_begin index s_end s_begin self.seq_len # 解码器输入窗口往前回退 label_len 个位置 r_begin s_end - self.label_len r_end r_begin self.label_len self.pred_len seq_x self.data[s_begin:s_end] # 编码器输入 [seq_len, enc_in] seq_y self.data[r_begin:r_end] # 解码器输入 [label_lenpred_len, c_out] seq_x_mark self.data_stamp[s_begin:s_end] # 编码器时间特征 seq_y_mark self.data_stamp[r_begin:r_end] # 解码器时间特征 return seq_x, seq_y, seq_x_mark, seq_y_mark这里seq_y的长度是label_len pred_len其中前label_len个位置是真实值后pred_len个位置会在训练时被置零。这个设计是 Informer 的生成式解码器核心模型不是一步一步预测而是把一整段未来序列作为解码器输入用 mask 让注意力只关注前面已知的真实部分。还必须注意一个边界条件如果seq_len小于label_lenr_begin s_end - label_len可能变成负数。比如--seq_len 24 --label_len 48当index0时r_begin-24Python 会按倒数索引取数据不会报错但取到的数据完全是错的。我一般会建议至少保证seq_len label_len或者让数据加载器做边界裁剪。2.3 用注释标记快速定位张量形状读注释版代码时我最依赖的是张量形状注释。官方源码很多变量名很抽象比如batch_x、batch_y、x_mark单看名字不知道维度。注释版会在每个关键张量后面标形状我自己维护项目时也养成了这个习惯。# 编码器输入 batch_x: [B, seq_len, enc_in] # 历史序列 batch_x_mark: [B, seq_len, time_feature_dim] # 历史时间特征 # 解码器输入 batch_y: [B, label_len pred_len, c_out] # 真实值 待预测位置的占位 batch_y_mark: [B, label_len pred_len, time_feature_dim] # 模型输出 outputs: [B, label_len pred_len, c_out] # 与 batch_y 同形状time_feature_dim取决于--freq。对 ETTh1 这种小时级数据常见的时间特征是 hour、weekday、day、month 等通常 4 到 5 维。很多人在自定义数据时把enc_in设成特征列数 时间特征维数导致输入维度翻倍。实际上时间特征是通过另外的 embedding 通道处理的不会占用enc_in。新手看到这里容易蒙注释版的价值就是把这类约定直接写在代码旁边。3. 核心模块逐段注释ProbSparse Attention、蒸馏层与生成式解码器读完数据流之后再啃模型。Informer 最值得逐行看的三个地方是 embedding、ProbSparse Attention 和编码器的蒸馏层。这三个地方都直接把论文公式转成了张量操作注释版的意义在于把公式里的字母对应到代码里的Q、K、V和维度上。3.1 从 embedding 到三份输入模型的前向入口在models/model.py的Informer.forward。它接收的是数据加载器返回的四个张量然后先做 embedding。# models/model.py 的 Informer 类前向 def forward(self, x_enc, x_mark_enc, x_dec, x_mark_dec, enc_self_maskNone, dec_self_maskNone, dec_enc_maskNone): # x_enc: [B, seq_len, enc_in] # x_mark_enc: [B, seq_len, time_feature_dim] # x_dec: [B, label_lenpred_len, c_out] # x_mark_dec: [B, label_lenpred_len, time_feature_dim] enc_out self.enc_embedding(x_enc, x_mark_enc) # enc_out: [B, seq_len, d_model] enc_out self.encoder(enc_out, attn_maskenc_self_mask) dec_out self.dec_embedding(x_dec, x_mark_dec) # dec_out: [B, label_lenpred_len, d_model] dec_out self.decoder(dec_out, enc_out, ...) return dec_outembedding 在models/embed.py里。Informer 使用的是三类嵌入叠加数值嵌入、位置嵌入、时间特征嵌入。# models/embed.py 中的 DataEmbedding class DataEmbedding(nn.Module): def __init__(self, d_model, dropout): super().__init__() self.value_embedding nn.Linear(enc_in, d_model) # 把数值特征投影到 d_model self.position_embedding PositionalEmbedding(d_model) # sin/cos 位置编码 self.temporal_embedding TemporalEmbedding(d_model) # 时间特征查表 self.dropout nn.Dropout(dropout) def forward(self, x, x_mark): # x: [B, L, enc_in], x_mark: [B, L, time_feature_dim] x self.value_embedding(x) self.position_embedding(x) self.temporal_embedding(x_mark) return self.dropout(x)这里有一个容易看漏的细节value_embedding输入是enc_in输出是d_modeltemporal_embedding接收的是x_mark里面每个时间字段会被当成离散索引查一个nn.Embedding再相加得到d_model。所以x_mark的字段顺序很重要数据加载器里生成哪些时间特征embedding 里就必须按同样顺序接收。如果自定义数据集改了时间特征只用官方训练脚本容易维度错位。3.2 ProbSparse Attention采样为什么要用 factor 和 sample_kProbSparse Attention 是 Informer 最核心的改动。全量注意力对每个 query 都要和所有 key 做点积复杂度是 O(L^2)。Informer 先估计每个 query 的稀疏性只对信息量最大的部分 query 做全量注意力其余 query 用局部均值代替。对应代码在models/attn.py的ProbAttention里。# models/attn.py 中 ProbAttention 的核心方法 def _prob_QK(self, Q, K, sample_k, n_top): B, H, L_Q, D Q.shape # Q: [B, H, L_Q, D] # K: [B, H, L_K, D] L_K K.shape[-2] # 1. 每个 query 随机采样 sample_k 个 key计算近似得分 K_sample K[:, :, torch.randint(0, L_K, (L_Q, sample_k))] # K_sample: [B, H, L_Q, sample_k] Q_K_sample torch.matmul(Q, K_sample.transpose(-2, -1)) # Q_K_sample: [B, H, L_Q, sample_k] # 2. 稀疏性度量 M max - mean越大说明这个 query 越可能支配注意力 M Q_K_sample.max(-1).values - Q_K_sample.mean(-1).values # M: [B, H, L_Q] # 3. 每个头只保留 top n_top 的 query 做完整注意力 M_top M.topk(n_top, sortedFalse)[1] # M_top: [B, H, n_top] ...注释版里一般会额外标出sample_k和n_top的计算方式。常见实现里sample_k min(int(self.factor * np.log(L_K)), L_K) # 对 key 的采样数量 n_top min(int(self.factor * np.log(L_Q)), L_Q) # 保留的 query 数量默认factor5当L_Q96时n_top大约是 23。也就是说96 个 query 里只有约 23 个会走完整 softmax其余 query 的注意力权重直接用整个注意力的均值填充。这个近似让复杂度从 O(L^2) 降到 O(L log L)。factor是一个很关键的参数。设太小比如 1采样不足稀疏性度量不稳定loss 可能偏高设太大比如 20n_top接近L_QProbSparse 退化成全量注意力性能和 full attention 一样但代码里还是走随机采样反而损失了稳定性。我在实际项目里一般从 5 起步如果数据噪声很大会调到 7 或 10但不会超过 12。3.3 蒸馏层与解码器维度变化如何影响层数编码器部分Informer 在每个注意力层后面接了一个蒸馏层作用是让序列长度逐步减半。这是用一维卷积加池化实现的。# models/encoder.py 中的蒸馏层 self.distil nn.Sequential( nn.Conv1d(d_model, d_model, kernel_size3, stride1, padding1, biasFalse), nn.GELU(), nn.MaxPool1d(kernel_size3, stride2, padding1) ) # 前向 # 输入 [B, L, d_model] - 转置为 [B, d_model, L] - 卷积 - 池化 - 转回 # 输出长度约为 L/2长度变化可以直接算MaxPool1d的输出长度是floor((L 2*padding - kernel_size) / stride) 1。代入L96, kernel3, stride2, padding1得到floor(95/2)1 48。每层减半所以e_layers不能随便加。比如seq_len48, e_layers4长度变化是 48→24→12→6最后一层池化勉强能跑如果e_layers6最后一层长度大约是 3卷积核都比输入长直接报错或产生 NaN。解码器的核心在生成式输入构造。前面数据加载器提到seq_y的长度是label_len pred_len前label_len是真实值后pred_len是占位。模型内部并不是让解码器回归整段而是通过 mask 让每个位置只能看到它之前的位置最后只取pred_len部分的输出作为预测。这和传统 Transformer 的 decoder 不同它不需要循环生成一次前向就能得到全部预测。注释版里通常会特别强调label_len不是给模型自由发挥的预热区它是真实值的“启动 token”。所以评估指标应该只关注pred_len部分如果把label_len部分也算进 MSE数值会异常好看但那是模型在复读真实值不是预测能力。4. 用注释版跑通训练最小命令与必调参数读完代码下一步是把它跑起来。Informer 的官方训练入口是main_informer.py用命令行参数控制几乎一切。下面这套命令是我在 ETTh1 上验证过的最小复现组合适合第一次跑通。python -u main_informer.py \ --model Informer \ --data ETTh1 \ --freq h \ --features M \ --seq_len 96 \ --label_len 48 \ --pred_len 96 \ --enc_in 7 \ --dec_in 7 \ --c_out 7 \ --d_model 512 \ --n_heads 8 \ --e_layers 3 \ --d_layers 2 \ --attn prob \ --factor 5 \ --distil True \ --dropout 0.05 \ --learning_rate 0.001 \ --loss mse \ --train_epochs 6 \ --batch_size 32 \ --patience 3这套命令对应的是 ETTh1 的 7 个油电特征featuresM表示用全部历史特征预测全部未来特征。--freq h表示小时粒度这必须和数据集的真实时间戳对齐。如果数据集是 15 分钟粒度的--freq h会让时间特征编码错位但不容易报错只会让精度变差。下面是几个我在调参时最关心的参数参数默认值作用我的建议factor5控制 ProbSparse 采样的 query 数量小数据用 3-5长序列大数据用 5-10d_model512模型宽度直接决定显存数据量小时用 256收敛更快e_layers3编码器层数每层序列长度减半不要超过log2(seq_len)否则蒸馏后长度不够distilTrue是否开启蒸馏训练阶段别关推理阶段如果长度不够可以关掉label_len48解码器真实启动 token 长度通常是seq_len的一半不要大于seq_lenlearning_rate0.0001优化器学习率长预测任务用 0.0001短预测可以试 0.001我在实际项目里发现d_model和e_layers的关系比想象中更紧密。很多人把e_layers加到 6期望模型表达能力更强但输入长度只有 96每层减半后只剩 6 个 token注意力头和卷积核都没有足够的位置做局部建模。这种情况下模型不是变强了而是退化成池化器。我一般的判断标准是蒸馏后的最小长度不要小于 12否则宁可减小e_layers或把distil关了。训练开始后日志里每一行会打印类似epoch 1/6, train loss 0.432, vali loss 0.398, test loss 0.401, mse 0.245, mae 0.312这样的内容。这里test loss不是只在最终测试阶段出现而是在每个 epoch 结束后都对测试集做一次评估所以它能不能跟着训练集下降是判断过拟合的第一个信号。如果test loss持续上升而train loss还在降说明该用早停或者调大dropout。还有一个参数容易被忽略--use_amp。OpenAI 时代的炼丹习惯是看到显存不够就打开混合精度但 Informer 的 ProbSparse Attention 里存在topk和随机采样混合精度下梯度偶尔会出现不稳定的 NaN。注释版代码通常会把use_amp标注成“谨慎开启”。我自己的做法是先用纯 FP32 跑通确认结果稳定后再开 AMP 做加速一旦出现 NaN先关掉 AMP 而不是先调学习率。5. 避坑Informer代码注释版里最常踩的五个运行时问题代码注释看得懂不代表运行时不踩坑。这一章整理了我见过最多的五个问题每个都按现象、原因、解决来写。5.1 蒸馏层报错seq_len 乘以 e_layers 之后长度不够现象训练到第一个 batch 就报错错误信息类似Expected 3D tensor, got 2D或者Calculated padded input size per kernel invalid。原因--distil True时编码器每一层会把序列长度减半。如果seq_len24而e_layers4长度会从 24 变成 12、6、3最后一层的卷积核大小为 3在长度为 3 的序列上做MaxPool1d(kernel3, stride2, padding1)时实际计算出的输出长度可能是 0 或负数底层就会抛出形状错误。解决把--seq_len调到 48 以上或者减少--e_layers也可以直接设--distil False绕过长度减半。但需要知道关掉distil后模型复杂度会上升小数据集上容易过拟合。我一般优先调整e_layers而不是关蒸馏。5.2 自定义数据时间戳解析失败freq 和日期格式没对齐现象数据加载阶段报错比如time data 2024-01-01 00:00:00 does not match format %Y-%m-%d %H:%M:%S或者模型能跑但预测精度比论文差很远。原因Informer 的时间特征编码依赖--freq参数。--freq h表示小时级--freq t表示分钟级数据加载器会用不同的时间字段去生成特征。如果 CSV 里的日期是分钟级数据却设了--freq h或者日期列有缺失值解析就会失败。解决先把日期列清洗成统一格式再根据实际采样粒度传freq。常见做法是在进模型之前用 pandas 做一次预处理# 自定义数据预处理示例 import pandas as pd df pd.read_csv(your_data.csv) df[date] pd.to_datetime(df[date], format%Y-%m-%d %H:%M:%S) df df.sort_values(date).reset_index(dropTrue) # 如果是 15 分钟采样训练脚本里用 --freq t5.3 featuresS 但 enc_in 没改成 1预测结果是一条水平线现象训练 loss 能下降验证集 loss 也不差但最终预测曲线几乎是最近一段历史值的平移或者完全没有形状。原因--features S表示单变量预测但很多人把--enc_in、--dec_in、--c_out都留成默认的 7。模型内部会把 7 维数据都作为输入却只取第一列作为预测目标。结果模型学到了一个“把所有列平均值当成下一帧预测”的退化解。解决确认单变量任务时featuresS搭配enc_in1, dec_in1, c_out1。还有数据加载器里的target_cols也要只取目标列否则前面 7 维数据仍然会参与拼装。最好在数据预处理阶段就把其他列删掉让数据文件本身就是单变量。5.4 训练 loss 持续 NaN学习率、AMP 和数据泄漏现象第一个 epoch 的 loss 是正常数值到第二个 epoch 突然变成 NaN或者一开始就是 NaN。原因概率注意力里用了topk和随机采样这些操作在 FP16 混合精度下容易出现梯度爆炸。另一个常见原因是数据里有 NaN 或无穷值StandardScaler 一 fit 就拿到了一个 NaN 的均值。还有一个更隐蔽的原因是学习率太大导致 embedding 层的数值在反向传播后发散。解决先检查原始 CSV 是否有空值再把--use_amp关掉最后把学习率从 0.001 降到 0.0001。如果这三个都做了还是 NaN把--batch_size减半再看。不要一开始就用 AdamW 的默认参数Informer 的官方建议是 0.0001 起步尤其在长预测任务上。5.5 预测结果看起来像滞后一期label_len 部分被计入了评估现象测试集 MSE 很低但画出来的预测曲线比真实值慢一个周期像是把上一段历史搬到了未来。原因Informer 的解码器输入包含前label_len个真实值模型输出的前label_len个位置本质上是在复读这些真实值。如果评估时把整段pred_len都算进去前几个预测点会非常接近真实值导致整体误差虚低但真正需要预测的未来部分可能并不准。解决评估和画图时只取输出中从label_len之后开始的pred_len部分。在exp_main.py的测试函数里通常会有类似outputs outputs[:, -self.args.pred_len:, :]这一步。如果你自己写的评估脚本没有这行等于把启动 token 的复读也算成了预测能力。这是注释版代码里最容易被我标红的一行。6. 进阶把注意力分数导出来验证当前数据的稀疏假设Informer 的加速前提是“注意力分数服从长尾分布”但这并不是所有数据集都成立。如果数据本身平稳性很差ProbSparse 的采样估计可能不准这时与其盲目调factor不如直接把注意力矩阵导出来看。注释版代码里通常保留了一个--output_attention开关。打开后模型前向会额外返回注意力张量。在exp_main.py的_process_one_batch里可以这样把它存下来# exp_main.py 中带 attention 导出的训练片段 if self.args.output_attention: outputs, attn self.model(batch_x, batch_x_mark, dec_inp, batch_y_mark) # attn 形状通常是 [层数, B, 头数, L_Q, L_K] np.save(fattn_epoch{epoch}_batch{i}.npy, attn.detach().cpu().numpy()) else: outputs self.model(batch_x, batch_x_mark, dec_inp, batch_y_mark)拿到.npy文件后我一般会做两个快速验证。第一个是看每个 query 的注意力分布是否足够尖取最后一个 batch 的注意力矩阵按头求平均统计每一行的 top5 概率和。如果 top5 占比超过 80%说明稀疏假设成立attnprob是合理的如果只有 40%-50%说明这个数据集的注意力本来就比较平滑换成attnfull效果可能更好。第二个验证是看采样是否稳定连续跑两个 epoch比较同一个 batch 的M_top索引如果每次选出的 top query 都不一样说明factor太小采样噪声太大可以适当调大。我现在拿新数据集做实验时已经养成一个习惯先开output_attention跑一个 epoch把注意力分布打印出来再决定用 prob 还是 full而不是默认相信论文里的稀疏假设。这个小步骤帮我省掉了很多盲目调参的时间。希望帮到你。本文还有配套的精品资源点击获取
02
RELATED NEWS

相关资讯

更多网站建设与数字化升级内容

03
WHY YAOTU

想打造同款高转化官网?

懂行业、懂生意,从建站到增长一站式陪跑

◈

场景化定制

不做模板站,围绕你的业务场景量身设计,小众不撞款。

◐

营销型架构

以转化目标组织内容与路径,让官网真正带来询盘。

▲

全周期服务

设计、开发、运营、运维一体,上线只是开始。

免费获取你的建站方案

留下需求,专属顾问 24 小时内为你输出方案建议。