简介:本资源是Informer时间序列预测模型的代码详细注释版,面向深度学习初学者与时间序列建模实践者,旨在降低Transformer类长序列预测模型的理解门槛。压缩包共63个文件,涵盖17个核心Python源码(含models/、exp/、data/等模块)、4个Shell脚本(支持ETTh1/WTH等数据集一键运行)、5个CSV数据样例、4个PNG模型结构图与实验结果图、1个Dockerfile及环境配置文件(environment.yml、requirements.txt),整体大小62.33MB,结构清晰、开箱即用。已有663人学习下载,适合需要深入理解Informer自注意力机制改进(如ProbSparse Attention)、Encoder-Decoder架构设计、时间特征嵌入及长序列预测工程实现的学习者。注释覆盖全部关键函数与模块逻辑,辅以README.md说明与ipynb示例,可直接用于复现实验、调试模型或教学讲解。
1. Informer代码详细注释版:不是“能跑就行”的复现包,而是你真正看懂ProbSparse自注意力、长序列时序预测黑匣子的逐行解剖刀
你有没有试过 clone 下来一个号称“SOTA”的时序模型仓库,pip install -r requirements.txt后python main_informer.py一跑——loss 下降了,test MSE 打印出来了,但合上终端那一刻,脑子里只剩下一个问号:它到底在哪一步把 96 步历史压缩成 48 维隐状态?mask 是怎么在 encoder 里悄悄跳过 70% 的 QK 计算的?为什么 decoder 的 self-attn 不用 ProbSparse,而 cross-attn 又必须用?
这个「Informer代码详细注释版」就是为解决这种“玄学复现”而生的。它不是原始论文代码的简单打包,而是对Informer2020-main仓库中全部 37 个 Python 文件、5 个核心 shell 脚本、2 个关键配置文件(environment.yml / Makefile)进行了逐函数、逐循环、逐 if 分支的中文注释覆盖,注释密度达 1:1.8(即平均每 1.8 行代码配 1 行注释),重点标注了ProbSparse Attention 的采样逻辑、timefeatures.py 中 7 种时间编码的物理意义、data_loader.py 里 multivariate 数据如何被切片为 (batch, seq_len, features) 张量、以及 checkpoints 目录下模型文件名中每个字段(sl96_ll48_pl24_dm512_nh8_el2_dl1_df2048_atprob_fc5_ebtimeF_dtTrue_mxTrue)对应的实际超参含义。适合两类人:刚接触长序列时序预测的新手,想绕过 Transformer 黑箱直接理解 Informer 设计哲学;也适合已调通 baseline 但卡在指标提升瓶颈的熟手,靠注释反向定位attn.py中prob_mask生成时机或exp_informer.py中inverse_transform是否漏掉归一化逆操作。这不是一份“能跑就行”的资源,而是一份你愿意打印出来、贴在显示器边框上、边 debug 边划重点的源码地图。
2. 从零跑通 ETTh1 单变量预测:环境搭建、数据准备与训练命令的完整链路
2.1 环境隔离与依赖安装:为什么必须用 environment.yml 而非 requirements.txt?
原始仓库同时提供了requirements.txt和environment.yml,但实测发现:仅用 pip 安装 requirements.txt 会导致 PyTorch 与 CUDA 版本错配,进而触发torch.cuda.is_available()返回 False,即使显卡正常工作。根本原因在于requirements.txt中只写了torch>=1.7.0,未约束 CUDA 编译版本;而environment.yml显式声明了pytorch=1.7.1=py3.8_cuda11.0.221_cudnn8.0.3_0,强制匹配 CUDA 11.0 工具链。这是 Informer 训练中第一个隐形断点。
提示:不要跳过 conda 环境重建。我曾因复用旧环境导致
torch.fft在 decoder 中报RuntimeError: fft: ATEN not compiled with MKL support,耗时 3 小时排查才发现是 MKL 库版本冲突。
执行以下命令创建纯净环境:
# 创建并激活新环境(conda 4.12+) conda env create -f environment.yml conda activate informer-env # 验证关键依赖 python -c "import torch; print(f'PyTorch: {torch.__version__}, CUDA: {torch.version.cuda}, Available: {torch.cuda.is_available()}')" # 正常输出应为:PyTorch: 1.7.1, CUDA: 11.0.221, Available: Trueenvironment.yml中还锁定了numpy=1.19.2和pandas=1.1.5,这是为兼容data_loader.py中pd.read_csv(..., parse_dates=['date'])的日期解析逻辑——新版 pandas 在parse_dates处理空值时行为变更,会导致ETTh1.csv中部分缺失时间戳被转为NaT,后续timefeatures.py的time_features函数调用.dt.hour时抛出AttributeError。
2.2 数据下载与目录结构校验:ETT 数据集的三个隐藏约定
Informer 论文使用的 ETT(Electricity Transformer Temperature)数据集并非直接内嵌在代码包中,需手动下载。官方提供地址为 GitHub Release(https://github.com/zhouhaoyi/Informer2020/releases/download/v1.0/ETT-small.zip),但实际使用中必须注意三点:
- 文件名大小写敏感:解压后必须得到
ETTh1.csv、ETTh2.csv、ETTm1.csv、ETTm2.csv四个文件,且扩展名全为小写.csv。若下载包内为ETTh1.CSV,Linux/macOS 下data_loader.py的os.path.join(data_path, f'{flag}.csv')将返回None,引发FileNotFoundError; - 时间列名硬编码:所有 ETT 文件首列为
date,第二列为预测目标(如OT)。data_loader.py第 42 行df_raw = pd.read_csv(os.path.join(data_path, f'{flag}.csv'))后,第 45 行border1s = [0, 12*30*24 - self.seq_len, 12*30*24+4*30*24 - self.seq_len]直接按固定索引切分训练/验证/测试集,不读取文件头判断列数。若你误将ETTh1.csv替换为自定义数据,且首列非date或目标列非第二列,df_raw.values将包含时间字符串,导致model.py中x_enc = x_enc.float()报ValueError: could not convert string to float; - 数据路径必须严格匹配脚本参数:
ETTh1.sh中--data_path ./data/ETT-small/指向的./data/ETT-small/目录下,必须存在ETTh1.csv。若你将文件放在./data/ett-small/(小写 ett),os.path.exists(data_path)返回False,data_loader.py第 38 行assert os.path.exists(data_path), f'data file not found at {data_path}'直接中断。
校验命令(Linux/macOS):
# 进入项目根目录后执行 mkdir -p data/ETT-small wget https://github.com/zhouhaoyi/Informer2020/releases/download/v1.0/ETT-small.zip unzip ETT-small.zip -d data/ # 检查文件名与内容 ls -l data/ETT-small/ETTh1.csv head -n 3 data/ETT-small/ETTh1.csv # 应输出:date,OT,...(两列,首行是表头)2.3 启动单变量预测训练:从 shell 脚本到核心参数的映射解析
scripts/ETTh1.sh是启动 ETTh1 单变量预测的标准入口。其内容看似简单,但每个参数都直指 Informer 架构的关键设计:
# scripts/ETTh1.sh 关键片段 python -u main_informer.py \ --model informer \ --data ETTh1 \ --root_path ./data/ETT-small/ \ --data_path ETTh1.csv \ --features S \ # ← 核心!S=Single-variate, M=Multivariate --target OT \ # ← 当 features=S 时,target 必须指定单列名 --freq h \ # ← 时间频率:h=hourly, t=15min, d=daily --seq_len 96 \ # ← encoder 输入长度(历史窗口) --label_len 48 \ # ← decoder 输入长度(带 mask 的起始 token) --pred_len 24 \ # ← decoder 输出长度(预测步长) --enc_in 1 \ # ← encoder 输入特征维度(S 模式下恒为 1) --dec_in 1 \ # ← decoder 输入特征维度(含 target + covariates) --c_out 1 \ # ← decoder 输出维度(单变量预测为 1) --d_model 512 \ # ← embedding 维度(也是 attention head 的输入维度) --n_heads 8 \ # ← attention head 数量(d_model 必须被 n_heads 整除) --e_layers 2 \ # ← encoder 层叠数(含 ProbSparse attn + FFN) --d_layers 1 \ # ← decoder 层叠数(含 masked self-attn + cross-attn + FFN) --d_ff 2048 \ # ← feed-forward 网络隐藏层维度 --dropout 0.05 \ # ← dropout rate(应用于 attn output 和 FFN output) --attn prob \ # ← attention 类型:prob=ProbSparse, full=标准 Transformer --factor 5 \ # ← ProbSparse 中 top-k 的 k 值(k = d_model // factor) --embed timeF \ # ← 时间特征嵌入方式:timeF=Fourier, fixed=fixed embedding --activation gelu \ # ← FFN 激活函数 --output_attention False \ --distil True \ # ← 是否启用蒸馏模块(decoder 中的额外 attention) --mix True \ # ← 是否混合 encoder 输出(cross-attn 中 query 来自 decoder,key/value 来自 encoder) --des 'Exp' \ --itr 1 \ --train_epochs 6 \ --patience 3这里需要强调两个易错参数:
--features S与--target OT是强绑定的。若设--features S但漏写--target,data_loader.py第 102 行cols = list(df_raw.columns)后,df_raw = df_raw[['date'] + cols[1:]]会错误地将所有列(包括HUFL,HULL,MUFL等)都作为输入特征,导致enc_in实际为 7 而非 1,model.py初始化self.encoder时enc_in=7与d_model=512不匹配,报RuntimeError: size mismatch;--attn prob必须与--factor 5配合。attn.py第 127 行scores_top = torch.topk(scores, top_k, sorted=False)[0]中top_k = d_model // factor = 512 // 5 = 102(向下取整),若factor设为 3,则top_k=170,但scores张量第二维(key 长度)为seq_len=96,torch.topk将因k > dim_size报错。
运行命令:
chmod +x scripts/ETTh1.sh ./scripts/ETTh1.sh训练日志中关键验证点:
Encoder input shape: torch.Size([32, 96, 1])→ 确认 batch=32, seq_len=96, enc_in=1;ProbSparseAttention: top_k=102, sparse_ratio=0.105→102/96≈1.06,说明 top_k 被自动 clip 到seq_len,此时 sparse_ratio 无意义,属正常现象;vali mse: 0.1234, mae: 0.2567→ 首轮验证 loss 应在 0.1~0.3 区间,若 >1.0 说明数据加载异常。
3. 注释深度解析:attn.py中 ProbSparse Attention 的四层实现逻辑
3.1 从公式到代码:ProbSparse 的数学本质与prob_mask生成机制
Informer 论文公式 (4) 定义 ProbSparse Self-Attention 的核心思想:不计算全部 QK^T 矩阵,而是对每个 query,只保留与其最相关(score 最高)的 top-u 个 key,其余置为负无穷(mask 掉)。其中 u = ⌈log(L)⌉ * d,L 为序列长度,d 为 embedding 维度。但代码中并未直接实现该公式,而是采用更鲁棒的采样策略——这正是注释版的价值所在。
models/attn.py第 89 行开始的_prob_QK函数,是 ProbSparse 的心脏。我们逐段解析其注释逻辑:
def _prob_QK(self, Q, K, sample_k, n_top): # Q: [B, H, L, D], K: [B, H, S, D] # Step 1: 计算 QK^T 得到原始相似度矩阵 scores: [B, H, L, S] # 注意:此处未除以 sqrt(d_k),因后续 softmax 会归一化,省略不影响 top-k 选择 B, H, L, D = Q.shape _, _, S, _ = K.shape scores = torch.einsum("bhld,bhsd->bhls", Q, K) # [B, H, L, S] # Step 2: 对每个 query(L 维),随机采样 sample_k 个 key(而非全部 S 个) # sample_k = 25(由 factor=5, d_model=512 推出:sample_k = d_model // factor = 102 → 但实际设为 25) # 为何是 25?注释版指出:这是作者经验性设定,避免 top-k 过大导致内存爆炸 U_part = torch.div(scores, np.sqrt(D)) # [B, H, L, S],为后续采样做准备 U_part = U_part.clone() # 防止原地修改影响梯度 U_part = U_part.permute(0, 1, 3, 2) # [B, H, S, L],将 key 维度前置以便采样 # Step 3: 对每个 key(S 维),随机选取 sample_k 个 query 位置,计算其 score 均值 # 这是 ProbSparse 的精髓:用局部统计量(均值)代替全局最大值,降低方差 scores_top = torch.zeros(B, H, L, n_top).to(Q.device) # [B, H, L, n_top] index = torch.zeros(B, H, L, n_top).to(Q.device).long() # 对每个 batch 和 head,独立采样 for i in range(B): for j in range(H): # 从 S 个 key 中随机选 sample_k 个索引 idx = torch.randint(0, S, (sample_k,)) # 取出这些 key 对应的 scores(即 U_part[i,j,idx,:]),形状 [sample_k, L] scores_i = U_part[i, j, idx, :] # [sample_k, L] # 计算每个 query(L 维)在这 sample_k 个 key 上的 score 均值 scores_i_mean = torch.mean(scores_i, dim=0) # [L] # 对每个 query,取其 top-n_top 个 score_i_mean 值对应的 key 索引 # 注意:此处 top-k 是在 sample_k 个 key 的均值上选,而非全量 S 个 _, top_idx = torch.topk(scores_i_mean, n_top, sorted=True) # [n_top] # 将 top_idx 扩展为 [n_top, 1],与 scores_i 索引对齐 scores_top[i, j, :, :] = scores_i[:, top_idx].t() # [L, n_top] index[i, j, :, :] = idx[top_idx].unsqueeze(0).repeat(L, 1) # [L, n_top]这段代码揭示了两个关键事实:
- ProbSparse 并非严格按公式 (4) 实现,而是用
sample_k个随机 key 的 score 均值来近似全量 key 的分布,再从中选 top-n_top。sample_k=25是经验值,远小于S=96,大幅降低计算量; n_top(即factor=5决定的top_k)作用于采样后的子集,而非全量 key。这意味着实际参与计算的 key 数量是n_top,但采样过程引入了随机性,这也是 ProbSparse 具有正则化效果的原因。
3.2masking.py中的三种 mask:encoder、decoder self、decoder cross 的差异化应用
Informer 的 mask 机制比标准 Transformer 更精细,utils/masking.py定义了三类 mask,注释版明确标出了它们在模型各处的调用位置:
| Mask 类型 | 生成函数 | 形状 | 应用位置 | 注释关键点 |
|---|---|---|---|---|
TriangularCausalMask | __init__(self, B, L, device='cpu') | [B, L, L] | decoder.py第 112 行dec_self_mask = TriangularCausalMask(B, L, device) | 仅用于 decoder 的 self-attention,确保预测t时刻时不看到t+1及之后;L是label_len + pred_len = 48+24=72,非seq_len=96 |
ProbMask | __init__(self, B, H, L, index, scores, device='cpu') | [B, H, L, S] | attn.py第 127 行prob_mask = ProbMask(B, H, L, index, scores, device) | 专为 ProbSparse 设计,将index中未选中的 key 位置置为float('-inf'),scores参数用于调试(打印 mask 前后 score 分布) |
FullAttentionMask | __init__(self, B, L, S, device='cpu') | [B, L, S] | attn.py第 145 行mask = FullAttentionMask(B, L, S, device) | 仅当attn=full时启用,生成全Falsemask(即无 mask),但代码中仍调用mask.mask方法,体现架构一致性 |
特别注意ProbMask的构造逻辑(masking.py第 45 行):
def __init__(self, B, H, L, index, scores, device='cpu'): # index: [B, H, L, n_top],记录每个 query 选中的 key 索引 # scores: [B, H, L, S],原始 QK^T 分数 super(ProbMask, self).__init__() self.mask = torch.ones(B, H, L, S, dtype=torch.bool, device=device) # 将选中的 index 位置设为 False(即不 mask),其余为 True(mask 掉) for i in range(B): for j in range(H): self.mask[i, j, torch.arange(L), index[i, j, :, :].t()] = False这里self.mask是布尔型,True表示该位置被 mask(置为-inf),False表示保留。attn.py第 132 行scores.masked_fill_(self.mask, -np.inf)即完成最终屏蔽。注释版在此处添加了调试技巧:在masked_fill_前插入print(f'Mask ratio: {(self.mask.sum() / self.mask.numel()).item():.3f}'),可实时监控当前 batch 的稀疏比例,验证 ProbSparse 是否生效。
3.3timefeatures.py中的七种时间编码:为什么timeF比fixed更适配电力负荷预测?
utils/timefeatures.py实现了 Informer 支持的全部时间特征嵌入方式,注释版对每种方法的物理意义和适用场景做了标注:
def time_features(df, time_col='date', freq='h'): # freq 取值:'h'(hourly), 't'(15min), 'd'(daily), 'b'(business day), 'w'(weekly), 'm'(monthly), 'y'(yearly) df['month'] = df[time_col].dt.month df['day'] = df[time_col].dt.day df['weekday'] = df[time_col].dt.weekday df['hour'] = df[time_col].dt.hour df['minute'] = df[time_col].dt.minute // 15 # 仅当 freq='t' 时有效 # 关键区别:timeF 使用 Fourier 变换,fixed 使用 learnable embedding if freq == 't': # 15min 数据:周期为 96(24h/15min) df['microsecond'] = df[time_col].dt.microsecond // 150000 # 150000ms = 15min elif freq == 'h': # hourly 数据:周期为 24(日周期)、168(周周期=24*7) df['dayofweek'] = df[time_col].dt.dayofweek df['dayofyear'] = df[time_col].dt.dayofyear # ... 其他 freq 处理 # timeF 核心:对每个周期性特征,生成 sin/cos 对 # 例如 hour: sin(2π*hour/24), cos(2π*hour/24) # weekday: sin(2π*weekday/7), cos(2π*weekday/7) # 这种编码具有平移不变性,且能表达任意周期长度 feat_set = ['month','day','weekday','hour'] if freq == 't': feat_set.append('minute') elif freq == 'h': feat_set.extend(['dayofweek','dayofyear']) # 注释版强调:Fourier 编码无需训练,对长序列泛化更好;而 fixed embedding 需为每个周期值(如 hour=0~23)学习一个向量,在 ETTh1 这种跨年数据中,不同年份的 hour 分布可能偏移,Fourier 更鲁棒 return df[feat_set]在ETTh1.sh中--embed timeF指定使用 Fourier 编码。注释版指出:若你更换为--embed fixed,必须同步修改embed.py第 62 行self.time_embedding = nn.Embedding(24, d_model)中的24为实际周期长度(如dayofweek周期为 7,则需nn.Embedding(7, d_model)),否则forward中time_emb = self.time_embedding(time_feat)会因索引越界报错。而timeF自动适配所有周期,无需手动配置。
4. 避坑指南:训练与推理中五个高频翻车现场及血泪解决方案
4.1 现象:训练 loss 为 nan,且vali mse从第一轮就显示nan
原因:data_loader.py第 132 行scaler.fit(train_data)中,train_data包含NaN值。ETT 数据集虽宣称无缺失,但ETTh1.csv中HUFL列在 2016-07-01 前有连续 12 行为空,pd.read_csv默认将空字符串转为NaN,StandardScaler对NaN调用.mean()返回NaN,后续x_enc = scaler.transform(x_enc)输出全NaN,model.py中x_enc = self.enc_embedding(x_enc)输入NaN导致loss.backward()梯度爆炸。
解决:在data_loader.py第 128 行df_raw = pd.read_csv(...)后插入清洗代码:
# 清洗 NaN:用前向填充(ffill)处理时间序列缺失 df_raw = df_raw.fillna(method='ffill').fillna(method='bfill') # 双重填充防首尾 NaN或在ETTh1.sh中预处理数据:sed -i 's/^,,/0,0,/g' data/ETT-small/ETTh1.csv(Linux)。
4.2 现象:test anything.ipynb运行到model.load_state_dict(torch.load(checkpoint_path))报Missing key(s) in state_dict
原因:检查点文件checkpoints/informer_ETTh1_ftM_sl96_ll48_pl24_dm512_nh8_el2_dl1_df2048_atprob_fc5_ebtimeF_dtTrue_mxTrue_test_0/checkpoint.pth中的模型权重键名,与当前model.py中Informer类的__init__定义不一致。常见于你修改了encoder.py中EncoderLayer的子模块名(如将self.attention改为self.attn),但未更新state_dict的load逻辑。
解决:在加载前打印键名对比:
checkpoint = torch.load(checkpoint_path) print("Checkpoint keys:", list(checkpoint.keys())[:5]) print("Model keys:", list(model.state_dict().keys())[:5]) # 若发现 model 有 'encoder.layers.0.attention...' 而 checkpoint 是 'encoder.layers.0.attn...', # 则需手动映射:checkpoint = {k.replace('attn', 'attention'): v for k, v in checkpoint.items()} model.load_state_dict(checkpoint)4.3 现象:--features M多变量预测时,vali mse极低(<0.01)但test mse高达 5.0+
原因:data_loader.py第 102 行cols = list(df_raw.columns)获取所有列名后,df_raw = df_raw[['date'] + cols[1:]]将date列置于首位,但--target OT指定的目标列OT在原始ETTh1.csv中是第二列(索引 1),而多变量模式下cols[1:]包含OT,HUFL,HULL,MUFL,MULL,LUFL,LULL共 7 列。--target OT仅用于确定c_out=1,但data_loader.py第 148 行data = df_raw[cols[1:]].values加载全部 7 列作为输入特征,OT列被当作普通 covariate,而非监督信号。真正的监督信号来自df_raw[cols[1:]]的第二列(即HUFL),导致训练目标错位。
解决:修改data_loader.py第 148 行,显式提取target列:
# 原代码:data = df_raw[cols[1:]].values # 修改为: target_col_idx = cols.index(self.target) # 获取 target 列索引 # 输入特征:除 date 和 target 外的所有列 feature_cols = [col for col in cols[1:] if col != self.target] data = df_raw[feature_cols].values # X data_y = df_raw[[self.target]].values # y(监督信号)4.4 现象:Dockerfile构建镜像后,python main_informer.py报ModuleNotFoundError: No module named 'utils'
原因:Dockerfile第 10 行COPY . /app/将整个项目复制到/app/,但未执行pip install -e .或设置PYTHONPATH。main_informer.py中from utils.metrics import metric依赖相对导入,而 Docker 容器内 Python 解释器默认不将/app加入sys.path。
解决:在DockerfileCOPY后添加:
WORKDIR /app ENV PYTHONPATH="/app:${PYTHONPATH}" # 或更规范:安装为可编辑包 # RUN pip install -e .4.5 现象:--pred_len 48预测 48 步时,result_univariate.png图中预测曲线在 24 步后突然变平成直线
原因:exp/exp_informer.py第 186 行pred = inverse_transform(pred)调用utils/tools.py的inverse_transform函数,但该函数默认只对pred的最后一维(即c_out=1)进行逆变换。当pred_len=48时,pred形状为[B, 48, 1],逆变换正确;但若你在main_informer.py中误将predreshape 为[B, 1, 48],inverse_transform会错误地对dim=1(即 batch 维)做逆变换,导致所有样本共享同一组逆变换参数,输出失真。
解决:在exp_informer.py第 185 行后插入形状校验:
print(f"pred shape before inverse: {pred.shape}") # 应为 [B, pred_len, c_out] if len(pred.shape) == 3 and pred.shape[1] != args.pred_len: raise ValueError(f"pred shape {pred.shape} mismatch with args.pred_len {args.pred_len}") pred = inverse_transform(pred)5. 模型诊断与结果可视化:用test anything.ipynb深度验证你的训练是否真正收敛
5.1 从 checkpoint 提取中间层输出:定位 attention 权重异常的黄金三步法
test anything.ipynb不仅是测试脚本,更是模型诊断利器。当你发现 test MSE 高于预期,不要急于调参,先用以下三步定位问题根源:
Step 1:加载模型并开启output_attention=True
修改ETTh1.sh中--output_attention False为True,重新训练 1 epoch。这会强制model.py第 172 行return dec_out, attns返回 attention 权重字典attns,其中attns['encoder']包含每层 encoder 的 ProbSparse mask 结果。
Step 2:在 notebook 中提取并分析attns
# 加载训练好的模型(确保 --output_attention=True) model = Informer(...) model.load_state_dict(torch.load('checkpoints/.../checkpoint.pth')) model.eval() # 构造 dummy input(与训练时 shape 一致) x_enc = torch.randn(32, 96, 1) # [B, L, enc_in] x_dec = torch.randn(32, 72, 1) # [B, label_len+pred_len, dec_in] x_mark_enc = torch.randn(32, 96, 4) # time features x_mark_dec = torch.randn(32, 72, 4) # 前向传播获取 attention with torch.no_grad(): dec_out, attns = model(x_enc, x_dec, x_mark_enc, x_mark_dec) # 分析 encoder 第一层的 attention mask enc_attn_layer0 = attns['encoder'][0] # [B, H, L, S] print(f"Encoder layer 0 attn shape: {enc_attn_layer0.shape}") print(f"Sparsity ratio: {(enc_attn_layer0 == float('-inf')).float().mean().item():.3f}") # 正常值应在 0.7~0.9 之间(70%~90% 的位置被 mask)Step 3:可视化 attention 热力图,识别 collapse 现象
import matplotlib.pyplot as plt import seaborn as sns # 取 batch=0, head=0 的 attention map attn_map = enc_attn_layer0[0, 0].cpu().numpy() # [96, 96] # 将 -inf 替换为 0 以便可视化 attn_map = np.where(attn_map == float('-inf'), 0, attn_map) plt.figure(figsize=(10, 8)) sns.heatmap(attn_map, cmap='viridis', cbar_kws={'label': 'Attention Score'}) plt.title('Encoder Layer 0 Attention Map (Batch 0, Head 0)') plt.xlabel('Key Position') plt.ylabel('Query Position') plt.show()若热力图显示所有 query 都集中在少数几个 key(如第 10、20、30 位)上形成强烈亮斑,其余区域全黑,说明 ProbSparse 的采样失效,模型退化为关注固定时间点(如每天 0 点、12 点),这是过拟合或数据泄露的征兆。此时应检查data_loader.py的数据切分逻辑,确认border1s划分未将测试集未来信息混入训练。
5.2 多变量预测的指标拆解:为什么metric.py中的mae比mse更值得信任?
utils/metrics.py提供了metric函数计算mae,mse,rmse,
本文还有配套的精品资源,点击获取