资讯详情 时空变换网络做交通流预测:源码复现与训练调参实战
📅 2026/10/10 10:46:22
简介面向城市交通流预测任务的时空变换网络ST-TransformerPython实现附完整训练与预测源码及PEMSD7路网交通流量数据集。该模型由时空卷积模块与注意力机制构成可在复杂路网中捕捉动态时空关联替代传统统计方法适合深度学习方向研究者、交通数据分析人员及高校学生用于复现、调优或扩展。压缩包共9个文件其中6个Python脚本分别承担模型结构定义、图卷积网络层编写、训练循环、验证评估、预测推理及独热编码预处理等任务2个csv文件存放路网邻接矩阵与交通速度序列另含1个Markdown文档说明项目结构与使用方法整体大小仅451KB便于快速下载部署。已有387人学习下载覆盖从模型原理到工程实现的完整链路。利用该套代码可快速搭建训练与评估流程不仅能观察时空卷积与注意力模块的协作方式还可通过可视化理解模型行为继而替换自定义数据开展进一步实验是课设、竞赛或相关研究的实用起点。1. 时空变换网络做交通流预测为什么值得复现这套源码交通流预测是个老问题但直到最近两年时空变换网络Spatial-Temporal Transformer Network后文简称 STTN才把 Transformer 从 NLP 搬到路网数据上同时建模传感器之间的空间依赖和早晚高峰的时间周期。这里的价值在于一套干净的 Python 源码加上公开的交通数据集你可以在一个晚上拉通数据预处理、模型训练、误差评估不用从零造轮子。这篇笔记适合手头有交通速度或流量预测任务、想在空间序列上试 Transformer 的读者我会从数据格式、邻接矩阵、模型结构一路写到训练参数与踩坑记录。2. 交通流预测的模型选型从 ARIMA 到时空变换网络为什么选型比调参更重要2.1 时空序列预测的三个难点空间耦合、时间周期与非平稳扰动交通流数据和图像、语音不一样它的样本之间不是独立的。某个传感器检测到的车速既受上游一公里处拥堵的影响又和半小时前自己这条车道的流量相关。这种空间依赖和时间依赖交织在一起让传统时序模型很难招架。ARIMA 和卡尔曼滤波这类经典方法本质上只建模单点的时间变化你把 100 个检测器拼成多维输入ARIMA 依然当成 100 条独立序列分别预测。问题是路网是连通的一条匝道堵死十分钟后相邻干道速度全掉下来这种空间上的传导效应单点模型永远学不到。到了深度学习时代大家开始用 GCN 配合 RNN。GCN 在路网拓扑上做图卷积RNN 在时间维上递推组合起来确实能抓到时空依赖。但实际操作中这套组合有两个别扭的地方一是 GCN 需要你先把邻接矩阵算对算不对整条链路全白做二是 RNN 的序列训练是串行的交通数据集动辄几十万时间步训练速度非常折磨人。时空变换网络的意义在于把这两个维度统一成一个架构。它用注意力机制同时处理传感器节点之间的关系和时间步之间的关系不需要 GCN 那种显式的图卷积核也不需要 RNN 的循环展开。这里想说的选型理由很直白如果你的数据集规模上千节点、时序长度上万步Transformer 的并行训练优势会明显压过 RNN。2.2 Transformer 怎么“时空化”节点即 token时间步即序列把标准 Transformer 用到交通流上第一步是决定 token 是什么。常见做法是让每个传感器节点成为一个 token每个时间步的观测值作为这个 token 的特征。你在 NLP 里理解“一个词是一个 token”在交通流里就把“一个检测器是一个 token”。但这样做有一个缺陷标准 Transformer 的 self-attention 不区分 token 之间的路网距离。两个物理距离十公里的检测器注意力得分可能和相邻检测器一样高。时空变换网络一般会在注意力里融合距离信息或者在输入侧做图卷积让模型知道 3 号节点和 4 号节点是邻居而 3 号节点和 89 号节点隔了两个区。我见过不少项目直接在标准 Transformer 上加一层位置编码就开始训练结果误差居高不下。原因就在于这种空间先验没有注入。反过来的一个反直觉经验是空间先验的作用在交通流预测里比在 NLP 里大得多因为路网的空间结构是物理存在不像文本里词的相对位置那么自由。2.3 源码任务目标先定义输入输出再谈模型结构拿到一套交通流预测源码第一件事不是读模型而是看它的输入输出定义。常见任务设定是用过去 12 个时间步通常 5 分钟一个步长12 步即 1 小时预测未来 12 个时间步的车速或流量。输入张量形状一般是[batch, T_in, N, F]T_in是历史时间步数N是检测器节点数F是特征维度。# 先把任务参数固定下来后面所有代码都以这个为准 T_IN 12 # 观察过去 1 小时 T_OUT 12 # 预测未来 1 小时 STEP 5 # 数据采样间隔单位分钟 N 307 # 传感器节点数量以 METR-LA 为例 F 1 # 特征数量这里用车速设置T_IN12和T_OUT12的原因很实际交通管理需要提前半小时到一小时知道路况变化而 12 步粒度是文献和工程落地之间最常用的折衷。步长再大比如 30 分钟一个时间步早晚高峰的快速变化会被模糊掉步长再小1 分钟一个点数据噪声会明显影响注意力权重计算。输出部分有两种设计。第一种是直接让模型输出[B, T_OUT, N, 1]一次预测出未来 12 步的完整曲线第二种是只输出第一步然后递归地用预测值喂回输入。前者训练稳定但误差会随预测步长累积后者更接近真实使用场景但容易在第二步之后发散。新手复现源码时建议先做第一种把指标跑通后再改成递归式去对比效果差异。3. 数据集准备把 PEMS 系列数据变成模型能吃的张量并构建邻接矩阵3.1 数据集结构与读取先确认维度顺序再谈预处理交通流预测领域最常用的公开数据集是 PEMS04、PEMS08 和 METR-LA。PEMS 系列来自加州高速公路的实时检测器每个检测器每 5 分钟记录一条数据一天 288 条包含流量、速度和占有率三个字段。拿到源码包里的数据文件后第一步不是直接加载训练而是把数据形状和值域打印出来确认。import numpy as np # 常见的 npz 格式里面可能有 data 和 adj 两个 key raw np.load(PEMS04.npz) data raw[data] # 形状可能是 [T, N, F]也可能是 [N, T, F] adj raw[adj] # 形状应该是 [N, N]但有时需要自己构建 print(data shape:, data.shape) print(data range:, data.min(), data.max()) print(data dtype:, data.dtype)这里最容易翻车的是维度顺序。同样一份 PEMS04 数据有人存成[T, N, F]有人存成[N, T, F]还有人把特征维度放在最前面。你不确认顺序就开训练后续所有形状推导全是错的而且这种错不会马上报异常只会让 loss 曲线很怪。我的习惯是打印出来先看如果data.shape[0]是 288 的整数倍一天的记录数那第一维大概率是时间。3.2 归一化与滑动窗口时序数据切分不能随机打乱对交通流做归一化几乎是必须的因为速度和流量不在一个量纲上直接拼进特征会让注意力权重被大数值特征带偏。常用做法是 StandardScaler把每个特征维度的均值归零、方差归一。有一点要注意拟合 scaler 只能用训练集不能用验证集和测试集否则相当于提前把未来数据的分布信息泄露给了模型。from sklearn.preprocessing import StandardScaler # 假设 data 是 [T, N, F]取速度列作为预测目标索引按实际数据调整 speed data[..., 1].astype(np.float32) # 只有一列特征时的写法 scaler StandardScaler() # 先按时间顺序切成三段7:1:2 train_len int(T * 0.7) val_len int(T * 0.1) train_data speed[:train_len] val_data speed[train_len:train_len val_len] test_data speed[train_len val_len:] scaler.fit(train_data) # 只 fit 训练数据 train_norm scaler.transform(train_data) val_norm scaler.transform(val_data) test_norm scaler.transform(test_data)然后是构建滑动窗口样本。交通流预测的输入输出都是连续时间片样本之间天然有重叠这是正常的。真正要注意的是切分时不能随机 shuffle原因在代码注释里def build_samples(seq, input_len12, pred_len12, step1): 把连续序列切成 (X, Y) 样本对。 step 控制窗口滑动的步长step1 表示每个时间步都取一个样本。 如果样本量太大可以调成 step12训练集直接缩小 12 倍。 X, Y [], [] for i in range(0, len(seq) - input_len - pred_len 1, step): X.append(seq[i:i input_len]) Y.append(seq[i input_len:i input_len pred_len]) return np.array(X), np.array(Y)注意这里Y的形状是[样本数, pred_len, N]对应未来 12 个时间步每个节点的速度。构建完样本后你会在源码里看到TensorDataset和DataLoader的组合。有一个细节是DataLoader的shuffle参数时序预测任务里一般设False。很多人习惯图像任务里 shuffleTrue搬到时序上直接照抄结果验证集指标异常地好因为训练集和验证集的样本在时间上重叠了。3.3 邻接矩阵构建距离阈值法和 k 近邻法怎么选邻接矩阵是图模型和时空注意力模型的“路网先验”。PEMS 数据集通常提供传感器之间的距离矩阵你可以根据距离构建邻接矩阵。两种常见做法距离阈值法两点距离小于阈值则连边k 近邻法每个节点连接最近的 k 个传感器。def build_adj_from_dist(distances, sigma0.1, threshold0.5): 用高斯核把距离矩阵转换成权重矩阵。 距离越近权重越大小于 threshold 的直接置 0 来稀疏化。 N distances.shape[0] adj np.exp(-distances ** 2 / (sigma ** 2)) adj[adj threshold] 0.0 np.fill_diagonal(adj, 0.0) # 去掉自环 # 对称归一化D^-1/2 * A * D^-1/2数值更稳定GCN 常用 d adj.sum(axis1) 1e-12 d_inv_sqrt np.power(d, -0.5) norm_adj d_inv_sqrt[:, None] * adj * d_inv_sqrt[None, :] return norm_adjsigma和threshold是两个需要手调的参数。sigma越大权重衰减越慢远距离节点也有较大连接threshold越小邻接矩阵越稠密计算量也越大。我在 PEMS08 上常用sigma0.1、threshold0.5换到 METR-LA 就不一定合适因为两个数据集的传感器分布密度差别很大调试时要打印邻接矩阵的平均度和孤立节点数。如果发现某些节点完全没有连接说明阈值太大或 k 近邻的 k 太小这类节点在注意力机制里基本学不到空间信息。4. 模型结构拆解与实现时空变换网络的核心模块与代码骨架4.1 时空嵌入位置编码和周期编码为什么不能省Transformer 本身没有顺序概念必须把时间步的位置信息编码进去。交通流比文本多一个维度周期。凌晨两点的速度和早上八点的速度规律完全不同但同一周的同一时刻高度相似。所以时间编码通常由两个 part 组成绝对位置编码记录步长序号周期编码记录一天内和一周内的相位。import torch import torch.nn as nn def time_position_encoding(seq_len, hidden_dim, period288): 生成时间步的位置编码。 seq_len 是输入序列长度period 是每天的步长数。 PEMS 数据 5 分钟一个点一天 288 个点换数据集时这个参数必须改。 pe torch.zeros(seq_len, hidden_dim) pos torch.arange(seq_len).unsqueeze(1).float() div_term torch.exp(torch.arange(0, hidden_dim, 2).float() * (-torch.log(10000.0) / hidden_dim)) pe[:, 0::2] torch.sin(pos * div_term) pe[:, 1::2] torch.cos(pos * div_term) return pe # [seq_len, hidden_dim]对应节点维度很多实现会加一个可学习的节点嵌入Node Embedding它的作用类似 NLP 里的 token embedding让模型知道当前处理的是哪个传感器。交通流预测里这个节点嵌入尤其重要因为不同传感器所在路段的车速水平差异很大市中心拥堵点和郊区快速路根本不是一个量级。4.2 多头注意力与门控融合时间维和空间维怎么协同时空变换网络的核心是一个“双分支”结构一条分支对时间维做多头自注意力捕获趋势变化另一条分支用邻接矩阵做空间传播捕获路网耦合。两条分支的输出不直接相加而是通过一个门控网络学融合权重。class TemporalAttentionBlock(nn.Module): 在时间维上做多头注意力。 输入形状 [B, T, N, D]把 T 当成序列维N 当成 batch 维展开。 def __init__(self, hidden_dim, num_heads, dropout0.1): super().__init__() self.attn nn.MultiheadAttention(hidden_dim, num_heads, dropoutdropout, batch_firstTrue) self.norm nn.LayerNorm(hidden_dim) self.dropout nn.Dropout(dropout) def forward(self, x): B, T, N, D x.shape # 把时间维移到序列位置节点数拼进 batch x_flat x.permute(0, 2, 1, 3).reshape(B * N, T, D) attn_out, _ self.attn(x_flat, x_flat, x_flat) attn_out self.dropout(attn_out) out self.norm(x_flat attn_out) # 残差 LayerNorm return out.reshape(B, N, T, D).permute(0, 2, 1, 3)这里有一个容易烧脑的形状变换permute(0, 2, 1, 3)把[B, T, N, D]变成[B, N, T, D]再把 batch 和节点数合并成B * N这样每个节点的 T 个时间步就变成独立的序列去做注意力。处理完后要记得 reshape 回去。这个维度变换在时空注意力模型里几乎是标配理解了它你就能读懂大部分开源实现。空间分支的代码更轻量核心是邻接矩阵乘法class SpatialPropagation(nn.Module): 用归一化邻接矩阵做空间信息传播。 本质上是把每个节点邻居的特征加权求和类似 GCN 的一层。 def __init__(self, hidden_dim, adj): super().__init__() self.adj torch.tensor(adj, dtypetorch.float32) self.linear nn.Linear(hidden_dim, hidden_dim) def forward(self, x): B, T, N, D x.shape # einsum 表示对每个 batch、每个时间步用邻接矩阵聚合邻居特征 x_t x.permute(0, 1, 3, 2) # [B, T, D, N] out torch.einsum(btdn,nm-btdm, x_t, self.adj) # 矩阵乘法 out out.permute(0, 1, 3, 2) # [B, T, N, D] return self.linear(out)torch.einsum这行不熟悉的话容易看成黑匣子它就是把[N, N]的邻接矩阵作用在最后一个维度上。btdn,nm-btdm的意思是对每个 batch b、每个时间步 t、每个输出维度 m把输入特征 d 乘上邻接矩阵的 n 行 m 列累加得到邻居聚合结果。门控融合的作用是让模型自己决定当前时刻该更相信时间分支还是空间分支——class GateFusion(nn.Module): 自适应门控融合时间信息和空间信息各给一个权重。 权重是学习出来的不是手动固定。 def __init__(self, hidden_dim): super().__init__() self.gate nn.Linear(hidden_dim * 2, hidden_dim) self.sigmoid nn.Sigmoid() def forward(self, temporal_out, spatial_out): g self.sigmoid(self.gate(torch.cat([temporal_out, spatial_out], dim-1))) return g * temporal_out (1 - g) * spatial_out固定加权0.5 * temporal 0.5 * spatial看着省事但高峰拥堵时空间传播更重要平峰时时间趋势更值得信赖这个权重本就应该随数据变化。门控机制让模型自己学代价只是多一个线性层和一次 sigmoid非常划算。4.3 把模块串成完整前向流程从输入张量到预测输出整体模型把上面的模块按层堆叠每一层都包含一个时间注意力、一个空间传播和一个门控融合然后通过全连接层输出预测值。class STTN(nn.Module): 时空变换网络的简化骨架。 输入 [B, T_in, N, F]输出 [B, N, T_out]。 def __init__(self, num_nodes, feat_dim, hidden_dim64, num_heads4, num_layers2, adjNone, t_out12): super().__init__() self.node_embed nn.Parameter(torch.randn(num_nodes, hidden_dim - feat_dim) * 0.02) self.t_pe time_position_encoding(12, hidden_dim) self.spatial nn.ModuleList() self.temporal nn.ModuleList() self.gate nn.ModuleList() for _ in range(num_layers): self.temporal.append(TemporalAttentionBlock(hidden_dim, num_heads)) self.spatial.append(SpatialPropagation(hidden_dim, adj)) self.gate.append(GateFusion(hidden_dim)) self.head nn.Linear(hidden_dim, t_out) def forward(self, x): B, T, N, F x.shape # 节点嵌入与时间编码都拼到特征上 node_token self.node_embed.unsqueeze(0).unsqueeze(0).expand(B, T, -1, -1) time_token self.t_pe[:T].unsqueeze(0).unsqueeze(2).expand(B, -1, N, -1) x torch.cat([x, node_token, time_token], dim-1) for i in range(len(self.temporal)): t_out self.temporal[i](x) s_out self.spatial[i](x) x self.gate[i](t_out, s_out) x # 残差连接 # 取最后一个时间步的特征做预测 out self.head(x[:, -1, :, :]) # [B, N, T_OUT] return out代码里有三个细节值得展开。第一node_token和time_token都是直接拼接而不是相加这样模型可以在后续层里自由决定用不用这些信息。第二残差连接在门控融合的输出上又加了一层 x这是 Transformer 训练的稳定器层数超过 4 层时没有残差基本会发散。第三预测头只取了最后一个时间步这种做法在 T_OUT 较短时没问题但 T_OUT 超过 24 步后会损失中间时间维的信息进阶玩法可以改成对时间维做一维平均池化再进全连接效果通常更好。5. 训练实操与避坑让模型在新数据集上收敛的调试经验5.1 最小训练脚本从数据加载到梯度裁剪模型结构搭好后训练脚本反而是最容易写错的地方。数据加载、损失函数、优化器和梯度裁剪每一步都有影响收敛的隐藏参数。下面这份最小脚本是我平常调通一个模型后再精简出来的骨架你可以直接改改变量名套用。import torch from torch.utils.data import DataLoader, TensorDataset device torch.device(cuda if torch.cuda.is_available() else cpu) # 假设 x_train, y_train 来自 build_samples 且已归一化 x_train torch.tensor(X_train, dtypetorch.float32) # [样本数, 12, N, F] y_train torch.tensor(Y_train, dtypetorch.float32) # [样本数, 12, N] dataset TensorDataset(x_train, y_train) loader DataLoader(dataset, batch_size64, shuffleFalse, drop_lastTrue) model STTN(num_nodesN, feat_dimF, hidden_dim64, num_heads4, num_layers2, adjadj, t_out12).to(device) optimizer torch.optim.AdamW(model.parameters(), lr1e-3, weight_decay1e-5) criterion torch.nn.L1Loss() # MAE比 MSE 对尖峰噪声更鲁棒 for epoch in range(100): model.train() total_loss 0.0 for xb, yb in loader: xb, yb xb.to(device), yb.to(device) out model(xb) # [B, N, 12] loss criterion(out, yb.transpose(1, 2)) # [B, N, 12] 对齐形状 optimizer.zero_grad() loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), 5.0) optimizer.step() total_loss loss.item() if (epoch 1) % 10 0: print(fepoch {epoch1}, loss {total_loss / len(loader):.4f})shuffleFalse前面解释过是时序任务的硬要求。clip_grad_norm_设成 5.0 是一个保守值它防止注意力机制偶尔产生巨大梯度把参数直接冲飞。loss 函数的选择上L1 损失MAE对交通流里的瞬时波动不敏感收敛曲线更平滑如果你更关注大误差的惩罚换L1Loss配合SmoothL1Loss的beta1.0是综合折衷。5.2 五个绕不开的坑现象、原因与解决方案这部分是血泪经验汇总。我复现过的多个交通流预测模型新数据集上翻车的原因高度集中在下面五类按排查顺序排列。坑一训练 loss 完全不动前几个 epoch 一直维持在初始值附近。现象loss 曲线是一条水平线数值比正常值大一个量级。原因通常是学习率偏大导致梯度在最优解附近反复震荡无法下降或者初始化时嵌入矩阵方差过大输入特征被噪声淹没。解决把学习率从1e-3降到5e-4同时把节点嵌入的初始化方差从0.02改小到0.01再试一次。这两个参数是玄学但按这个顺序调多数情况下第三个 epoch 就能看到明显下降。坑二训练集误差很好验证集误差远高于训练集且差距越来越大。现象epoch 越大train loss 和 val loss 的剪刀差越明显。原因是 Transformer 参数太多小数据集上极易过拟合。解决给时间注意力加 dropoutMultiheadAttention的dropout参数调到0.2或0.3同时把weight_decay从1e-5提到1e-4。更有效的是早停每 10 个 epoch 记一次验证集 MAE连续 30 个 epoch 不下降就停。坑三换到新数据集后STTN 效果比简单 LSTM 还差。现象同一条数据LSTM 的 MAE 是3.2STTN 跑出3.8。原因八成是邻接矩阵没做对称归一化或者传感器编号和数据顺序对不上。解决先打印adj的前 10 行看看数值范围如果每一行和不是 1 且对角线不为 0说明归一化没写好。再检查adj的节点顺序是否与数据列顺序一致不一致的空间传播完全是在乱传。坑四训练过程中偶尔出现 NaN loss让整个实验白跑。现象loss 某个 epoch 突然变成nan后面全部是nan。原因一般有两个输入数据里有缺失值NaN 或 inf没有清理或者 LayerNorm 输入的数值太大导致梯度爆炸。解决预处理时用线性插值把缺失值补掉确保data.min()和data.max()都是有限值。同时把学习率降到1e-4训练前打印一次 loss 如果已经是 NaN就检查归一化这一步是不是 fit 到了全量数据。坑五损失函数一直下降但最终预测曲线像“延迟版”的实际值。现象可视化预测结果发现预测曲线比真实曲线整体向右平移了一段。原因是模型学到的其实是上一时刻数值的惯性外推没有学到真正的动态变化。解决这个现象意味着时空注意力没有起作用检查你的时间注意力是否真的在 T 维上做而不是在 N 维上做。很多人permute写错后注意力在节点维度上运行模型退化成纯线性回归。5.3 参数怎么调学习率、batch size 与注意力头数的取舍下表给出 STTN 的几个关键参数在新手阶段最稳妥的起点值以及调参时的方向判断。这些参数不是越大越好很多课题组会给出掩码矩阵但交通流数据要收敛到一个好点参数必须相互配合。参数名推荐起点调参方向指标变化特征学习率1e-3验证集 loss 震荡就减半出现平台期则尝试5e-4batch size64显存不足就减半过小时收敛慢注意力头数4显存充裕可加8头数太多增加噪声dropout0.1验证集与训练集差距大就增大最大不建议超过0.5层数2指标饱和后可试3~4超 4 层必须配残差和预归一化学习率的调整最依赖反馈信号。如果训练早期 loss 上下跳动幅度超过 30%直接减半如果平滑下降但后期停滞适当提高一次到2e-3做短时间冲刺再降回来这是学习率热重启的简化版。注意力头数调整上num_heads4在节点数 300 左右的数据上通常够用头数加到 8 提升有限但显存占用几乎翻倍。这些参数之间不是独立的比如 dropout 调大后损失曲线会变高但验证集指标可能反而变好判断依据一定以验证集 MAE 为准。6. 验证与进阶用同一份数据对比基线确认你的模型真的有效模型跑通后最忌讳直接跳到“调结构”先做两件事基线对比和多步预测可视化。基线我建议至少跑一个 LSTM 和一个 GCNGRU 组合不需要调优用默认参数即可。如果 STTN 在测试集上的 MAE 不比 LSTM 低 10% 以上说明你的注意力机制没有真正学到时空依赖可能只是参数多带来的假象。评估指标上交通流预测常用 MAE、RMSE 和 MAPE 三个前两个直接由损失函数得到MAPE 要额外算def compute_metrics(y_true, y_pred): mae torch.abs(y_true - y_pred).mean().item() rmse torch.sqrt(((y_true - y_pred) ** 2).mean()).item() mape (torch.abs(y_true - y_pred) / (y_true.abs() 1e-6)).mean().item() * 100 return {MAE: mae, RMSE: rmse, MAPE(%): mape}MAPE 在车速接近零时数值会爆掉所以分母加了1e-6保护。另一个实用技巧是把预测结果按一天 288 个时间步重排画出某几个节点一天内的预测曲线和真实曲线叠加图重点看早晚高峰那两段。如果高峰期误差比其他时段高出一大截说明模型对突发拥堵的时间动态学习不足此时优先检查周期编码是否把星期信息丢了而不是急着加深网络。进阶方向有三个常见选择。第一把星期几、节假日、天气温度拼进编码层让模型区分工作日和周末的早高峰差异这个对 MAPE 的改善通常立竿见影。第二把训练好的注意力权重导出看哪些节点之间的注意力得分最高与邻接矩阵做对比你会发现模型学到的连接关系比单纯距离近邻更丰富这个方法能让黑匣子变白一点。第三把单步预测改成带教师强制的多步自回归训练时用真实值作为下一步输入推理时改用预测值这样能够在高动态场景下减少误差累积。我自己跑这类模型有个习惯先花两个晚上只调学习率和 dropout不改任何网络结构直到训练曲线稳定收敛确认模型本身没问题后再动嵌入和注意力结构。这套源码加数据集的组合我建议你也按这个流程走一遍先照抄跑通再谈改进。希望帮到你。本文还有配套的精品资源点击获取