☰
Informer代码注释版:ProbSparse注意力与蒸馏层实战调参指南
2026/9/26 4:32:39 网站建设 项目流程

简介:面向深度学习时序预测方向的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(description='Informer') # ---------- 数据相关参数 ---------- parser.add_argument('--data', type=str, default='ETTh1', help='数据集名称,决定 data_provider 加载哪个文件') parser.add_argument('--features', type=str, default='M', help='M: 多变量预测多变量,S: 单变量预测单变量,MS: 多变量预测单变量') parser.add_argument('--seq_len', type=int, default=96, help='输入历史窗口长度') parser.add_argument('--label_len', type=int, default=48, help='解码器里真实序列拼接的长度') parser.add_argument('--pred_len', type=int, default=96, help='预测长度') # ---------- 模型结构参数 ---------- parser.add_argument('--d_model', type=int, default=512, help='编码器/解码器内部特征维度') parser.add_argument('--n_heads', type=int, default=8, help='多头注意力的头数') parser.add_argument('--e_layers', type=int, default=3, help='编码器层数') parser.add_argument('--d_layers', type=int, default=2, help='解码器层数') parser.add_argument('--distil', type=bool, default=True, help='是否使用编码器里的自注意力蒸馏') parser.add_argument('--attn', type=str, default='prob', help='prob 或 full,选择稀疏注意力或全量注意力') parser.add_argument('--factor', type=int, default=5, help='ProbSparse 采样因子,越大越接近全量注意力') # ---------- 训练策略参数 ---------- parser.add_argument('--learning_rate', type=float, default=0.0001, help='学习率') parser.add_argument('--train_epochs', type=int, default=6, help='训练轮数') parser.add_argument('--batch_size', type=int, default=32, help='批大小') parser.add_argument('--patience', type=int, default=3, 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 model

enc_in、dec_in、c_out这三个值非常容易填错。enc_in是数据里参与预测的输入特征数,如果数据集有 7 列,--enc_in 7。c_out是最终预测的目标维度,features=M时通常等于列数,features=S时是 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_len+pred_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_len,r_begin = s_end - label_len可能变成负数。比如--seq_len 24 --label_len 48,当index=0时,r_begin=-24,Python 会按倒数索引取数据,不会报错,但取到的数据完全是错的。我一般会建议至少保证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_mask=None, dec_self_mask=None, dec_enc_mask=None): # x_enc: [B, seq_len, enc_in] # x_mark_enc: [B, seq_len, time_feature_dim] # x_dec: [B, label_len+pred_len, c_out] # x_mark_dec: [B, label_len+pred_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_mask=enc_self_mask) dec_out = self.dec_embedding(x_dec, x_mark_dec) # dec_out: [B, label_len+pred_len, d_model] dec_out = self.decoder(dec_out, enc_out, ...) return dec_out

embedding 在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_model;temporal_embedding接收的是x_mark,里面每个时间字段会被当成离散索引,查一个nn.Embedding,再相加得到d_model。所以x_mark的字段顺序很重要,数据加载器里生成哪些时间特征,embedding 里就必须按同样顺序接收。如果自定义数据集改了时间特征,只用官方训练脚本容易维度错位。

3.2 ProbSparse Attention:采样为什么要用 factor 和 sample_k

ProbSparse 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, sorted=False)[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 数量

默认factor=5,当L_Q=96时,n_top大约是 23。也就是说,96 个 query 里只有约 23 个会走完整 softmax,其余 query 的注意力权重直接用整个注意力的均值填充。这个近似让复杂度从 O(L^2) 降到 O(L log L)。

factor是一个很关键的参数。设太小,比如 1,采样不足,稀疏性度量不稳定,loss 可能偏高;设太大,比如 20,n_top接近L_Q,ProbSparse 退化成全量注意力,性能和 full attention 一样,但代码里还是走随机采样,反而损失了稳定性。我在实际项目里一般从 5 起步,如果数据噪声很大,会调到 7 或 10,但不会超过 12。

3.3 蒸馏层与解码器:维度变化如何影响层数

编码器部分,Informer 在每个注意力层后面接了一个蒸馏层,作用是让序列长度逐步减半。这是用一维卷积加池化实现的。

# models/encoder.py 中的蒸馏层 self.distil = nn.Sequential( nn.Conv1d(d_model, d_model, kernel_size=3, stride=1, padding=1, bias=False), nn.GELU(), nn.MaxPool1d(kernel_size=3, stride=2, padding=1) ) # 前向: # 输入 [B, L, d_model] -> 转置为 [B, d_model, L] -> 卷积 -> 池化 -> 转回 # 输出长度约为 L/2

长度变化可以直接算:MaxPool1d的输出长度是floor((L + 2*padding - kernel_size) / stride) + 1。代入L=96, kernel=3, stride=2, padding=1,得到floor(95/2)+1 = 48。每层减半,所以e_layers不能随便加。比如seq_len=48, e_layers=4,长度变化是 48→24→12→6,最后一层池化勉强能跑;如果e_layers=6,最后一层长度大约是 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 个油电特征,features=M表示用全部历史特征预测全部未来特征。--freq h表示小时粒度,这必须和数据集的真实时间戳对齐。如果数据集是 15 分钟粒度的,--freq h会让时间特征编码错位,但不容易报错,只会让精度变差。

下面是几个我在调参时最关心的参数:

参数默认值作用我的建议
factor5控制 ProbSparse 采样的 query 数量小数据用 3-5,长序列大数据用 5-10
d_model512模型宽度,直接决定显存数据量小时用 256,收敛更快
e_layers3编码器层数,每层序列长度减半不要超过log2(seq_len),否则蒸馏后长度不够
distilTrue是否开启蒸馏训练阶段别关,推理阶段如果长度不够可以关掉
label_len48解码器真实启动 token 长度通常是seq_len的一半,不要大于seq_len
learning_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_len=24而e_layers=4,长度会从 24 变成 12、6、3,最后一层的卷积核大小为 3,在长度为 3 的序列上做MaxPool1d(kernel=3, stride=2, padding=1)时,实际计算出的输出长度可能是 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(drop=True) # 如果是 15 分钟采样,训练脚本里用 --freq t

5.3 features=S 但 enc_in 没改成 1,预测结果是一条水平线

现象:训练 loss 能下降,验证集 loss 也不差,但最终预测曲线几乎是最近一段历史值的平移,或者完全没有形状。

原因:--features S表示单变量预测,但很多人把--enc_in、--dec_in、--c_out都留成默认的 7。模型内部会把 7 维数据都作为输入,却只取第一列作为预测目标。结果模型学到了一个“把所有列平均值当成下一帧预测”的退化解。

解决:确认单变量任务时,features=S搭配enc_in=1, dec_in=1, c_out=1。还有,数据加载器里的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(f'attn_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%,说明稀疏假设成立,attn=prob是合理的;如果只有 40%-50%,说明这个数据集的注意力本来就比较平滑,换成attn=full效果可能更好。第二个验证是看采样是否稳定:连续跑两个 epoch,比较同一个 batch 的M_top索引,如果每次选出的 top query 都不一样,说明factor太小,采样噪声太大,可以适当调大。

我现在拿新数据集做实验时,已经养成一个习惯:先开output_attention跑一个 epoch,把注意力分布打印出来,再决定用 prob 还是 full,而不是默认相信论文里的稀疏假设。这个小步骤帮我省掉了很多盲目调参的时间。希望帮到你。

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

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询