news 2026/10/2 2:38:27

Informer长序列预测实战:解决OOM与训练不收敛问题

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
Informer长序列预测实战:解决OOM与训练不收敛问题

简介:本资源是一份面向深度学习与人工智能初学者及进阶实践者的Informer模型时间序列预测实战教学包,聚焦长序列预测这一典型工业场景(如电力负荷、气象趋势、设备故障预警等)。资源包含完整可运行代码、多组实测数据集(ETTh1等)、详细参数配置说明及训练结果文件,帮助读者从零掌握Informer核心机制——ProbSparse自注意力与自注意力蒸馏技术,并基于个人数据集完成端到端建模。压缩包共64个文件,涵盖17个Python源码(含模型定义、数据加载、实验调度等模块)、17个Numpy数据文件、2个PyTorch模型权重(.pth)及环境配置(.yml),整体大小为115.95MB,结构清晰、模块解耦,便于理解Transformer改进思路与工程落地细节。目前已有2060人学习下载,配套代码已通过实际运行验证,开箱即用,显著降低复现门槛。

1. Informer不是“又一个Transformer”:它专治长序列预测里那些卡死、爆显存、训不动的玄学问题

你有没有试过用标准Transformer跑ETTh1(电力负荷)或Traffic(高速车流)这类长度动辄500+的时间序列?模型一跑就OOM,调小batch_size后loss飘得像没系安全带,验证集MAE比LSTM还高——不是你代码写错了,是原始Transformer的O(L²)自注意力根本扛不住长序列。Informer在2020年ICLR拿下Best Paper,靠的不是堆参数,而是实打实把时间序列预测的工程瓶颈给捅穿了:它用ProbSparse自注意力把计算复杂度从O(L²)压到O(L log L),再用自注意力蒸馏砍掉冗余token,让单卡3090能稳训126步输入、24步预测的完整任务。这份实战包不是教学Demo,而是作者在真实工业场景(比如某电网负荷预测系统)落地后拆出来的最小可运行闭环:含完整训练/推理脚本、ETTh1原始数据、预训练checkpoint、结果可视化脚本,连environment.yml都锁死了PyTorch 1.10+和CUDA 11.3——你不用再猜“哪个版本不报错”,直接conda env create -f environment.yml就能复现论文级指标。适合两类人:一是手头有自己时序数据(传感器、IoT、金融tick)想快速验证Informer效果的工程师;二是被Transformer内存墙卡住、急需一个开箱即用长序列方案的研究者。


2. 从零启动Informer训练:环境搭建、数据准备与核心参数含义逐行拆解

2.1 环境隔离与依赖安装:为什么必须用environment.yml而不是pip install?

Informer对PyTorch版本极其敏感。实测发现:PyTorch 1.12以上会触发torch.nn.functional.scaled_dot_product_attention的默认启用,而Informer的ProbSparse模块依赖手动实现的稀疏mask逻辑,一旦被自动替换就会导致attention权重全为0。environment.yml中明确锁定pytorch=1.10.2=cuda113py39h7e861b5_0,并禁用torchvision的自动升级:

# 执行前确认当前无其他conda环境干扰 conda env create -f environment.yml conda activate informer_env # 验证关键依赖版本 python -c "import torch; print(torch.__version__)" # 必须输出 1.10.2+cu113 python -c "import numpy; print(numpy.__version__)" # 必须输出 1.21.5

提示:若遇到ModuleNotFoundError: No module named 'torch._C',说明CUDA版本不匹配。检查nvidia-smi输出的驱动支持最高CUDA版本,若为11.6,则需修改environment.yml中cudatoolkit=11.3为cudatoolkit=11.6,并重装环境。

2.2 数据结构解析:ETTh1.csv不是普通CSV,它的字段顺序和缺失值处理决定模型成败

Informer要求输入数据严格遵循[timestamp, target_var, covariate_1, ..., covariate_n]格式。ETTh1.csv实际结构如下:

dateOTHUFLHULLMUFLMULLLUFLLULLOT_1...
2016-07-0131.228.127.530.229.832.131.931.5...

其中:

  • OT(Oil Temperature)是目标变量,必须放在第二列(索引1)
  • HUFL/HULL/MUFL/MULL/LUFL/LULL是6个协变量(High/Ultra/Medium/Low Level Load)
  • OT_1起是滞后特征(lagged features),Informer默认忽略所有以_结尾的列,仅用前7列
# data_loader.py中关键逻辑(已验证) def __read_data__(self): df_raw = pd.read_csv(os.path.join(self.root_path, self.data_path)) # 只取前7列:date + OT + 6 covariates cols = list(df_raw.columns); cols.remove('date') # 移除时间戳列 df_raw = df_raw[['date'] + cols[:6]] # 强制取前6个非date列作为covariates # 注意:此处隐含假设——目标变量OT必须是cols[0],否则需手动调整cols索引

注意:若你的数据目标变量不在第二列,必须修改data_loader.py第42行cols = list(df_raw.columns)[1:]为cols = [df_raw.columns.get_loc('your_target_col')] + other_covariates_indices,否则模型永远在预测错误变量。

2.3 核心参数详解:每个命令行参数背后都是论文里的一个技术决策

Informer的训练命令形如:

python main_informer.py \ --model informer \ --data ETTh1 \ --data_path ETTh1.csv \ --features M \ # M=Multivariate, S=Single, MS=Mixed --target OT \ --freq h \ --seq_len 126 \ # 输入长度,对应论文Table 1的"Input Length" --label_len 64 \ # Decoder输入长度(含历史信息),非预测长度! --pred_len 24 \ # 真正预测长度,对应论文"Horizon" --d_model 512 \ # 模型隐藏层维度,影响显存占用 --n_heads 8 \ # Attention头数,必须整除d_model --e_layers 2 \ # Encoder层数,Informer论文推荐2~3层 --d_layers 1 \ # Decoder层数,通常为1 --d_ff 2048 \ # FeedForward中间层维度 --attn prob \ # ProbSparse注意力,不可改为'full' --factor 5 \ # ProbSparse采样因子,越大越稀疏(默认5) --dropout 0.05 \ # Dropout率,过高会导致收敛慢 --embed timeF \ # 时间特征嵌入方式:timeF=时间戳分解,fixed=固定位置编码 --activation gelu \ --distil False \ # 是否启用蒸馏(原论文True,但此包设为False因已集成) --output_attention False \ --train_epochs 6 \ --patience 3 \ --batch_size 32 \ --learning_rate 0.0001

关键参数避坑点:

  • --label_len不是预测起点偏移量,而是Decoder的输入长度(含最后label_len步真实值+pred_len步待预测)。例如label_len=64, pred_len=24时,Decoder输入为64步真实值+24步mask,输出24步预测。
  • --features M必须与数据列数匹配:ETTh1有1目标+6协变量=7列,选M;若只预测OT单变量,需设S并注释掉协变量列。
  • --attn prob是Informer灵魂,设为full将退化为标准Transformer,显存暴涨3倍。

3. 训练全流程实操:从数据切分到checkpoint保存的每一步验证点

3.1 数据集自动切分逻辑:为什么val/test比例固定为20%/20%且不可改?

Informer源码采用固定切分策略,而非按比例随机划分。以ETTh1(8544条记录)为例:

  • train: 前60% → 0~5126行(5127条)
  • val: 接着20% → 5127~6835行(1709条)
  • test: 最后20% → 6836~8543行(1708条)

该逻辑硬编码在data/data_loader.py的__load_dataset__函数中:

# line 102-105 num_train = int(len(df_raw) * 0.6) num_test = int(len(df_raw) * 0.2) num_val = len(df_raw) - num_train - num_test # 剩余全给val border1s = [0, num_train - self.seq_len, len(df_raw) - num_test - self.seq_len] border2s = [num_train, num_train + num_val, len(df_raw)]

提示:若需自定义切分(如按时间点切分),必须修改此处三行代码。例如按2017-01-01前为train,之后为test,需替换border1s/border2s为时间戳索引位置。

3.2 训练过程监控:如何判断模型是否真正收敛而非假收敛?

Informer的loss曲线极易出现“伪收敛”:前3轮loss骤降,后续停滞在0.8~1.0(MAE量级)。真收敛标志有三:

  1. 验证集MAE持续下降:results/informer_*.npy中的metrics.npy第0列(MAE)在patience=3内至少下降0.02;
  2. Attention权重可视化正常:运行exp/exp_informer.py生成attention_weights.png,应看到清晰的对角线稀疏模式(ProbSparse特征),若全白或全黑则attention失效;
  3. 预测结果分布合理:pred.npy与true.npy的差值std应<0.15(ETTh1量纲下),过大说明过拟合。
# 实时监控验证集MAE(每epoch输出) grep "vali" train.log | tail -10 # 输出示例:vali 1.2345 | test 1.3456 | MAE 0.8765 # 关键看第三列MAE是否阶梯式下降

3.3 Checkpoint保存机制:为什么每次训练只保留最优模型且不覆盖?

Informer采用torch.save保存完整state_dict,路径为checkpoints/informer_custom_*.pth。其命名规则包含全部超参:

informer_custom_ftMS_sl126_ll64_pl24_dm512_nh8_el2_dl1_df2048_atprob_fc5_ebtimeF_dtTrue_mxTrue_test_0.pth

其中:

  • ftMS: features=M, target=OT
  • sl126: seq_len=126
  • ll64: label_len=64
  • pl24: pred_len=24
  • dm512: d_model=512
  • nh8: n_heads=8
  • el2: e_layers=2
  • dl1: d_layers=1
  • df2048: d_ff=2048
  • atprob: attn=prob
  • fc5: factor=5
  • ebtimeF: embed=timeF
  • dtTrue: distil=True(注意:此包实际为False,命名有误)
  • mxTrue: mixed=True(多变量混合预测)

注意:test_0表示第0次运行,避免覆盖。若需复用checkpoint,需手动复制到新路径并修改main_informer.py第127行model.load_state_dict(torch.load(...))的路径。


4. 预测与结果分析:如何用训练好的模型跑自己的数据并解读metrics.npy

4.1 单样本预测脚本:绕过完整pipeline直接调用模型

当需要快速验证某条新序列时,无需重跑整个test流程。新建predict_single.py:

import numpy as np import torch from models.model import Informer from data.data_loader import Dataset_ETT_hour # 加载模型 model = Informer( enc_in=7, dec_in=7, c_out=1, seq_len=126, label_len=64, pred_len=24, d_model=512, n_heads=8, e_layers=2, d_layers=1, d_ff=2048, attn='prob', factor=5, dropout=0.05, embed='timeF', activation='gelu' ).cuda() model.load_state_dict(torch.load('checkpoints/informer_custom_ftMS_sl126_ll64_pl24_dm512_nh8_el2_dl1_df2048_atprob_fc5_ebtimeF_dtTrue_mxTrue_test_0.pth')) model.eval() # 构造单样本输入(shape: [1, 126, 7]) sample_input = np.random.randn(1, 126, 7).astype(np.float32) # 替换为你的真实数据 sample_input = torch.from_numpy(sample_input).cuda() # 预测 with torch.no_grad(): pred = model(sample_input, None, None, None) # decoder_input等设为None print("Prediction shape:", pred.shape) # [1, 24, 1]

4.2 metrics.npy深度解读:MAE/RMSE/MAPES不只是三个数字

results/informer_*.npy/metrics.npy是1x3数组,对应:

  • [0]: MAE(Mean Absolute Error)→ 绝对误差均值,对异常值鲁棒
  • [1]: MSE(Mean Squared Error)→ 平方误差均值,放大大误差影响
  • [2]: MAPE(Mean Absolute Percentage Error)→ 百分比误差均值,要求真实值≠0

但关键在true.npy和pred.npy的物理意义:

  • true.npy: shape=(N, 24, 1),N为测试样本数,每行是24步真实值
  • pred.npy: shape=(N, 24, 1),对应预测值
  • real_prediction.npy: shape=(N*24, 1),展平后的预测序列(用于画图)
# 验证MAPE计算逻辑(避免分母为0) true_vals = np.load('results/.../true.npy').flatten() pred_vals = np.load('results/.../pred.npy').flatten() mask = true_vals != 0 # 过滤真实值为0的点 mape = np.mean(np.abs((true_vals[mask] - pred_vals[mask]) / true_vals[mask])) * 100 print(f"Manual MAPE: {mape:.2f}%") # 应与metrics.npy[2]一致

4.3 可视化预测效果:用matplotlib画出真实vs预测曲线

import matplotlib.pyplot as plt import numpy as np true = np.load('results/informer_custom_.../true.npy') # (N, 24, 1) pred = np.load('results/informer_custom_.../pred.npy') # (N, 24, 1) # 取第一个样本画图 plt.figure(figsize=(12, 4)) plt.plot(true[0].flatten(), label='True', color='blue') plt.plot(pred[0].flatten(), label='Predicted', color='red', linestyle='--') plt.title('Informer Prediction vs True (Sample 0)') plt.xlabel('Time Step') plt.ylabel('Value') plt.legend() plt.grid(True) plt.savefig('prediction_sample0.png', dpi=300, bbox_inches='tight') plt.show()

提示:若曲线完全不重合,先检查true.npy和pred.npy维度是否一致(必须同为3D),再验证data_loader.py中inverse_transform是否启用(ETTh1无需逆变换,但自定义数据需开启)。


5. 避坑指南:Informer训练中90%失败案例的根源与血泪解决方案

5.1 现象:训练loss为nan,且从第一轮就开始

原因:--learning_rate 0.0001在某些GPU上仍过大,尤其当d_model=512时梯度爆炸。
解决:将--learning_rate降至0.00005,并在main_informer.py第112行添加梯度裁剪:

torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)

5.2 现象:验证集MAE持续上升,训练集MAE下降

原因:--dropout 0.05过低导致过拟合,或--patience 3太短未等到收敛。
解决:增大dropout至0.1,同时将--patience设为5,并观察train.log中vali行是否在第5轮后开始下降。

5.3 现象:预测结果全为常数(如所有pred值=0.321)

原因:--features参数与数据列数不匹配。例如ETTh1有7列却设--features S,模型只读取第一列(date)导致输入全0。
解决:用pandas.read_csv('ETTh1.csv').shape确认列数,--features设为M(≥2列)或S(仅1列目标变量)。

5.4 现象:ImportError: cannot import name 'scaled_dot_product_attention'

原因:PyTorch版本>1.10.2,自动启用了新attention API,与ProbSparse冲突。
解决:严格按environment.yml创建环境,或手动降级:conda install pytorch=1.10.2 torchvision=0.11.3 cpuonly -c pytorch。

5.5 现象:RuntimeError: CUDA out of memory即使batch_size=1

原因:--seq_len 126时ProbSparse仍需O(L log L)内存,但--d_model 512使单层encoder显存达1.2GB。
解决:

  • 降低--d_model至256(显存减半,精度损失<2%)
  • 或增加--factor至10(更稀疏,但可能丢失长程依赖)
  • 终极方案:在models/attn.py第87行scores = torch.softmax(scale * scores, dim=-1)前加scores = scores.masked_fill(attn_mask == 0, -1e9)确保mask生效。

6. 进阶技巧:用Informer做多步滚动预测与工业部署的三个硬核习惯

6.1 多步滚动预测:如何用单次训练模型实现N天连续预测?

Informer原生不支持滚动预测(rolling forecast),需手动实现。核心思想:每次预测24步后,将预测结果拼接到历史序列末尾,滑动窗口重新输入:

def rolling_forecast(model, init_seq, steps=168): # 预测7天(168小时) pred_all = [] current_input = init_seq # shape: [1, 126, 7] for i in range(0, steps, 24): with torch.no_grad(): pred_step = model(current_input, None, None, None) # [1, 24, 1] # 将pred_step拼接到current_input末尾,删除最老24步 # 注意:需保持协变量同步更新(如时间特征) new_input = torch.cat([ current_input[:, 24:, :], # 删除最老24步 torch.cat([pred_step, torch.zeros(1, 24, 6).cuda()], dim=-1) # 拼接预测+0填充协变量 ], dim=1) # 新input shape: [1, 126, 7] pred_all.append(pred_step.cpu().numpy()) current_input = new_input return np.concatenate(pred_all, axis=1) # [1, 168, 1] # 调用 init_seq = torch.randn(1, 126, 7).cuda() # 替换为真实初始序列 result = rolling_forecast(model, init_seq, steps=168)

注意:此方法假设协变量(如温度、湿度)可预测或置0。工业场景中需接入外部预报API填充协变量。

6.2 工业部署 checklist:从checkpoint到ONNX的四步验证

步骤操作验证命令关键指标
1. 模型导出torch.onnx.export(model, dummy_input, "informer.onnx", opset_version=11)onnx.checker.check_model(onnx.load("informer.onnx"))无报错即通过
2. ONNX推理ort_session = onnxruntime.InferenceSession("informer.onnx")ort_session.run(None, {"input": dummy_input.numpy()})输出shape=(1,24,1)
3. TensorRT加速trtexec --onnx=informer.onnx --saveEngine=informer.trttrtexec --loadEngine=informer.trt --shapes=input:1x126x7Latency < 15ms(V100)
4. 内存泄漏检测在循环推理中监控nvidia-smiwatch -n 1 'nvidia-smi --query-gpu=memory.used --format=csv'内存占用稳定不增长

6.3 我的血泪习惯:每次改参数必做的三件事

  1. 改完参数立刻删checkpoint:rm -rf checkpoints/*。Informer的checkpoint命名含全部超参,但旧文件残留会导致torch.load意外加载错误模型;
  2. 训练前强制清空缓存:torch.cuda.empty_cache()+gc.collect()。曾因缓存残留导致seq_len=126时显存占用比seq_len=64还低(假象);
  3. 首次运行加--itr 1:避免多轮重复训练掩盖单次失败。Informer默认--itr 2,若第一轮失败第二轮可能因随机种子不同而成功,造成“偶发性可用”的错觉。

从那以后我每次改--d_model或--attn,都强制走一遍这三步——省下的调试时间够跑完两个完整实验。希望帮到你。

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

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

YOLOv8行人检测实战:数据集处理与PyQt界面集成全流程

简介&#xff1a;面向有深度学习基础的行人检测开发者&#xff0c;这套YOLOv8行人检测工程包整合了标注数据集、训练权重与图形界面三个核心部分&#xff0c;基于YOLOv8算法在数千张街道和交通场景图像上训练&#xff0c;平均精度均值达90%以上&#xff0c;可直接用于行人识别&…

作者头像 李华
网站建设 2026/10/2 2:36:53

基于LSTM的股票价格预测与量化策略实战:从数据到回测的完整链路

简介&#xff1a;这份资源是面向计算机相关专业学生与项目实战学习者的深度学习股票价格预测与量化策略研究完整项目&#xff0c;源自大四毕业设计&#xff0c;经导师指导并获99分评审认可。内容涵盖股票价格预测模型构建与量化策略实现&#xff0c;适合作为毕业设计、课程设计…

作者头像 李华
网站建设 2026/10/2 2:36:26

开源大模型私有化部署与LoRA微调、LangChain应用全链路实战

简介&#xff1a;面向AI大模型应用开发与落地场景&#xff0c;这份资料围绕开源大模型的环境配置、私有化部署、LoRA微调与LangChain应用展开&#xff0c;覆盖DeepSeek、Yi、Qwen、Baichuan、ChatGLM、MiniCPM等主流模型&#xff0c;适合正在学习大模型技术栈并希望动手实践的开…

作者头像 李华
网站建设 2026/10/2 2:35:49

C++中的6种构造函数举例详解

在 C 中&#xff0c;构造函数是一种特殊的成员函数&#xff0c;用于初始化类对象。在对象创建时自动调用&#xff0c;构造函数的主要作用是分配资源、初始化数据成员等。根据不同的功能和使用场景&#xff0c;C 提供了多种类型的构造函数&#xff1a;1. 默认构造函数 (Default …

作者头像 李华
网站建设 2026/10/2 2:33:27

用 CodeArts AI 智能体零依赖打造「鲜拾 FreshKeep」食材保鲜管家

一键开通华为云码道 CodeArts 代码智能体&#xff1a;https://developer.huaweicloud.com/codeartsco.html?sourcedmzntgwatomgit1&sourceaddmzntgwatomgithd 引言&#xff1a;从家庭冰箱里的浪费说起 中国家庭每年因食材过期造成的浪费触目惊心。一项调研显示&#xff0…

作者头像 李华