
简介面向交通流预测领域研究者和深度学习开发者提供时空变换网络ST-Transformer的完整Python实现与配套数据集。模型结合时空卷积模块和注意力机制能够捕捉路网动态时空依赖适用于城市交通流量预测、拥堵预警与智慧交通调度场景。资源包共9个文件其中6个.py脚本覆盖模型定义ST_Transformer.py、layers.py、GCN_models.py、训练验证train.py、validation.py与数据处理One_hot_encoder.py2个CSV文件为PEMSD7路网车流量数据V_25.csv和邻接矩阵W_25.csv另有README指导使用整体仅451KB轻量易部署。已有387人学习下载。通过该资源可深入理解ST-Transformer的架构细节、时空特征提取与注意力权重的实现方式掌握数据预处理、模型训练、性能评估的完整流程还可基于自带数据集复现预测效果并可自行替换数据进一步改进模型。1. 交通流预测的时空变换网络不再依赖预定义邻接矩阵交通流预测通常被建模为给过去一小时的数据预测未来一小时的速度或流量。过去五年的标配组合是图卷积加循环神经网络但它有一个硬约束——必须手工准备邻接矩阵并按固定拓扑建图。路网一旦新增传感器高速路遇到临时封路或大型活动导致流量路径改变这张静态图就失真图卷积的泛化也随之下降。时空变换网络不预设图结构而是让模型用自注意力从数据里学习哪个传感器、哪个时刻在相互影响。覆盖面从上游路段直接跨到下游几个街区比邻接矩阵的语义更贴合真实拥堵的传播方式。围绕这套方案写的源码由三部分组成时空编码、带掩码的多头自注意力、多步回归头。下面以高速公路传感器数据集 METR-LA 为样例把数据准备、模型搭建到训练评估完整讲一遍并给出可以直接复用的参数配置。2. 时空变换网络核心模块与编码器选型2.1 token化把传感器读数变成带时空信息的序列Transformer 处理的是「词序列」而交通流预测中序列里的每个元素是「某个传感器在某个时刻的速度数值」。直接把原始标量丢进注意力层有两个问题一是丢失位置与传感器身份二是数值尺度差异过大。因此常见做法是为每个观测样本构造一个 tokentoken 的维度和传感器数量无关只和嵌入维度 d_model 有关。我一般先用一个全连接层把单通道观测值映射到 d_model 维然后叠加两类可学习嵌入传感器嵌入每个传感器分配一个可学习的 d_model 维向量表示道路身份和地理位置。207 个传感器就是 207 行嵌入矩阵新增一个传感器只需在矩阵里追加一行并做增量训练。时间嵌入交通数据存在双周期特性早高峰和晚高峰在一天内出现工作日与周末规律明显不同。实现时把时间戳拆成两个特征当天内的分钟序号0 到 287步长 5 分钟和星期几0 到 6分别查嵌入表再相加。代码上用 PyTorch 实现如下import torch import torch.nn as nn class SpatioTemporalEmbedding(nn.Module): def __init__(self, num_nodes: int, d_model: int): super().__init__() self.value_proj nn.Linear(1, d_model) # 观测值映射 self.node_emb nn.Embedding(num_nodes, d_model) # 传感器身份 self.timeofday_emb nn.Embedding(288, d_model) # 当天时段 self.weekday_emb nn.Embedding(7, d_model) # 星期几 def forward(self, x, tod, weekday): # x: (B, T, N)tod/weekday: (B, T) 长度 T 的整数索引 B, T, N x.shape x self.value_proj(x.unsqueeze(-1)) # (B,T,N,D) node_ids torch.arange(N, devicex.device) x x self.node_emb(node_ids) # 广播到 (B,T,N,D) x x self.timeofday_emb(tod).unsqueeze(2) # (B,T,1,D) 广播 x x self.weekday_emb(weekday).unsqueeze(2) return x这段代码的要点是把观测值、传感器编号和时间特征叠加成同一个向量空间。timeofday_emb用 288 而不是 1440是因为 METR-LA 每 5 分钟一条记录一天共有 288 个采样点。unsqueeze(2)把时间嵌入从(B,T,D)扩成(B,T,1,D)PyTorch 会在节点维上自动广播不需要显式 repeat省显存也避免写错维度。想压缩参数时可以把value_proj换成卷积核为 1 的nn.Conv1d但收益通常不如把参数留给后面的自注意力。2.2 自注意力与因果掩码决定预测质量的核心得到(B,T,N,D)的 token 序列后需要把它压成标准 Transformer 期望的(B, seq_len, D)。两种接法在源码里都很常见把 T×N 全部平铺成一个大序列注意力在全时空域上两两计算称为全时空注意力或者保持 N 维作为 batch 维只在时间维上做注意力空间信息靠节点嵌入隐式传递。前一种对交通流预测的效果通常更好因为高峰期拥堵传播往往跨越大半个路网注意力需要看到「3 号传感器此刻的速度」与「40 号传感器 25 分钟前的速度」之间的长程耦合。METR-LA 的 N207、T12全时空注意力一次要计算约六万个 token 对单张 8GB 显存就能承受如果是 PEMS 全量数据上千个传感器就需要先按 3 分钟窗口采样或者分块。掩码在自注意力中承担两个职责。第一个是因果掩码预测 t 时刻的值时只能用 t 时刻及其之前的数据否则训练时会把未来信息泄漏给模型第二个是变量掩码交通数据存在大量缺失值某些传感器某时刻是 NaN与其让注意力在 NaN 上计算不如把它挡在 softmax 之前。构造因果掩码的代码如下def make_causal_mask(T: int, N: int) - torch.Tensor: # 仅在时间维上做因果限制节点维不受限 causal torch.tril(torch.ones(T, T, dtypetorch.bool)) causal causal.repeat_interleave(N, dim0).repeat_interleave(N, dim1) return causal # (T*N, T*N)repeat_interleave的含义是把时间轴上的因果约束复制到每个传感器上传感器 i 在 t1 时刻可以看传感器 j 在 t2≤t1 时刻的数据只要 t2t1 就被挡住。这一步如果漏做指标会显得异常好实测 MAE 可能直接下降 15% 以上且模型完全无法用于在线推理因为它偷看了未来。如果后续发现输出全是训练集均值再检查一下掩码是不是和序列排列顺序不一致常见错误是把 flatten 顺序和掩码构建顺序搞混。2.3 编码器堆叠与残差结构自注意力本身不改变数据维度深层结构主要靠前馈网络和残差完成非线性变换。常见做法是堆 2 到 4 个编码器层每层由「多头自注意力 → LayerNorm → 两层 MLP → LayerNorm」串成。相比 NLP 任务动辄 12 层、24 层交通序列的层数不需要太深堆到 6 层以上反而会在小数据集上出现注意力退化所有查询向量收敛到几乎相同的分布多头退化成单头。下表是 METR-LA 这类中等规模路网上的初始配置范围也适合用来给后续源码设定 baseline参数取值说明d_model64 / 128数据量小且显存紧时用 64num_heads4 / 8d_model 必须能被 head 数整除encoder_layers2 / 3超过 4 层收益开始递减dropout0.1 / 0.2大数据集用 0.1小数据集用 0.2feedforward_dimd_model × 4常见的经验倍数残差结构在这里承担一个容易被忽视的作用把原始观测信息直接传向后层避免梯度在深层传播时被非线性激活函数冲刷掉。自定义编码器层时记得先算 x x attn(norm(x))再做 x x ffn(norm(x))顺序反了会明显拖慢收敛。3. 数据集准备从 METR-LA 到可训练样本3.1 原始数据格式与时序切分交通流预测的公开数据集以两类为主METR-LA洛杉矶 207 个传感器2012 年 3 月到 6 月共 4 个月和 PEMS04加州 307 个传感器2018 年 1 月到 2 月。两者的存储结构高度一致一个[num_timesteps, num_sensors]的数值矩阵加上一个时间戳列表矩阵第 i 行第 j 列表示第 i 个时间步每 5 分钟传感器 j 的观测值通常以「速度」为单位。数据集传感器数时间跨度采样间隔METR-LA2072012-03 ~ 065 minPEMS043072018-01 ~ 025 minPEMS081702016-07 ~ 085 min动手写模型前第一件事是确认数据是否已经做了缺失值插补。公开版 METR-LA 已过滤一部分无效记录但仍有约 2% 的 NaN 或 inf。处理缺失值的顺序是先做最近邻插补再做标准化最后进模型。反过来操作会把缺失值信息泄露给归一化层的统计量导致测试集指标偏乐观。切分方式沿用交通流领域的惯例按时间顺序切为 70% 训练、10% 验证、20% 测试。与图像分类不同交通流样本之间强相关随机洗牌会让同一时段出现在训练集和测试集里严格时间顺序切分才能保证「模型预测的是没见过的时间段」。这组比例在 METR-LA 上的测试集大约覆盖最后一个半月难度比前段更大。3.2 滑窗采样与归一化给定一个固定窗口长度 T_in12过去一小时和预测步长 T_out12未来一小时滑窗会把长为样本数的时序矩阵切成无数个重叠样本。重叠是交通流预测中刻意设计的它让模型能在有限数据上见到更多模式缺点是相邻样本高度相似如果不用 shuffle 训练模型容易过拟合最后几个时间步。解决方法是训练集内先随机取窗口再在每个 batch 内打乱。标准化建议用 Z-score而不是 Min-Max。交通流数据近似正态分布但偶尔有 0 值传感器故障和极端值 80 mph。Min-Max 会让平均值附近的微小波动被压缩到几乎不可分辨而 Z-score 保留了这个差异。计算均值方差时只使用训练集统计量测试集沿用训练集算出的 mean 和 std防止测试信息通过统计量流入模型。import numpy as np def normalize(train, val, test): mean, std train.mean(), train.std() return (train - mean)/std, (val - mean)/std, (test - mean)/std def sliding_window(data, T_in12, T_out12, step1): X, Y [], [] for i in range(0, len(data) - T_in - T_out 1, step): X.append(data[i : i T_in]) Y.append(data[i T_in : i T_in T_out]) return np.array(X), np.array(Y)这里step1意味着原始 207 个传感器约 34000 个时间步会被切成三万个重叠样本训练随机打乱换成step6后样本数直接降到原来的六分之一训练快 6 倍但指标通常略差。优先推荐 step1 配 batch_size 64若显存或时间受限再增大 step。X 的形状是(样本数, 12, 207)Y 的形状是(样本数, 12, 207)模型要对 12 个未来时刻、所有传感器输出预测值。3.3 构造带时间特征的 DataLoadersliding_window只产出数值而 2.1 节的时间嵌入还要求样本知道自己每个时间步的「当天时刻」和「星期几」。转换方式是把样本内每个时间步的起点戳换算成整数再包装成一个数据集类。下面这个TrafficDataset把 X、Y、时间特征统一封装from torch.utils.data import Dataset import torch class TrafficDataset(Dataset): def __init__(self, X, Y, start_time, T_in12): self.X, self.Y, self.T_in X, Y, T_in # start_time 是原始数据第一条记录的时间戳datetime64 类型 self.start_ts start_time.astype(datetime64[m]) def __len__(self): return len(self.X) def __getitem__(self, idx): base self.start_ts np.timedelta64(idx * 5, m) step_ts base np.timedelta64(np.arange(self.T_in) * 5, m) tod (step_ts.astype(int64) % 1440) // 5 # 长度 T 的整数 wd step_ts.astype(datetime64[D]).astype(int64) % 7 return (torch.FloatTensor(self.X[idx]), torch.FloatTensor(self.Y[idx]), torch.LongTensor(tod), torch.LongTensor(wd))__getitem__里的tod会算出 0 到 287 的整数序列wd是 0 到 6 的整数序列两者都带着窗口内每个时间步的位置信息和 2.1 节的嵌入表维度严格对应。这里对wd的% 7结果是从 1970-01-01星期四起算的偏移量模型只需要类别一致即可不需要真的对齐到「周一到周日」。最后把数据集丢进DataLoader(dataset, batch_size64, shuffleTrue, num_workers4)训练循环就能稳定跑满 GPU。提示如果发现验证集指标反复横跳先看 DataLoader 的 shuffle 是否只对训练集开启。验证集和测试集必须按时间顺序输出否则相邻窗口的高度相似性会让评估结果虚高。4. 时空变换网络的 PyTorch 源码实现与训练参数4.1 组装一个可运行的 encoder-only 模型把 2.1 节的嵌入层、2.2 节的掩码和 2.3 节的编码器堆叠起来就得到完整模型。为了减少重复造轮子下面直接使用 PyTorch 自带的TransformerEncoderLayer它内部已经实现多头注意力和前馈网络只需在src_mask里传入因果掩码。import torch.nn as nn class STTransformer(nn.Module): def __init__(self, num_nodes, T_in, T_out, d_model64, num_heads4, num_layers3, dropout0.1): super().__init__() self.embed SpatioTemporalEmbedding(num_nodes, d_model) encoder_layer nn.TransformerEncoderLayer( d_modeld_model, nheadnum_heads, dim_feedforwardd_model*4, dropoutdropout, activationgelu, batch_firstTrue ) self.transformer nn.TransformerEncoder(encoder_layer, num_layers) self.regressor nn.Linear(d_model, 1) self.T_in, self.T_out, self.N T_in, T_out, num_nodes def forward(self, x, tod, weekday): src self.embed(x, tod, weekday) # (B,T,N,D) B, T, N, D src.shape src src.reshape(B, T * N, D) # (B,T*N,D) mask make_causal_mask(T, N).to(x.device) out self.transformer(src, maskmask) out self.regressor(out).reshape(B, T, N) return out[:, -self.T_out:] # 取最后 T_out 步src.reshape(B, T*N, D)这一行的顺序很关键reshape 默认按行优先(B,T,N,D)变成先遍历时间再遍历节点的排列和make_causal_mask里的repeat_interleave是严格对应的。如果写成src.permute(0,2,1,3).reshape(...)掩码索引就错位了训练时表面正常指标却会异常偏高。最后一行取out[:, -T_out:]是「输入 12 步、输出最后 12 步」的方式。模型对每个位置都产出了一个预测但因果掩码决定前几个位置根本没有足够历史只有第 T_in 步之后的位置才有完整上下文因此只保留尾部。注意src.reshape的顺序必须与掩码构建时repeat_interleave的语义一致。如果改动了 flatten 方式掩码也要对应修改否则模型会在训练时泄漏未来信息验证集看着很好一到真实在线推理就崩。4.2 训练循环、损失函数与早停损失函数选择上MSE 会放大拥堵时段的误差MAE 则对所有时段一视同仁。交通流预测的评估指标常同时报 MAE 和 RMSE训练若只优化 MAERMSE 会偏高我一般用 Huber Lossdelta1.0等价于误差小于 1 时用二次、大于 1 时用线性兼顾两者。训练过程加入两个细节学习率预热和早停。预热期通常 5 个 epoch 内把学习率从 1e-4 线性提到 1e-3。早停以验证集 MAE 为基准连续 10 个 epoch 不下降就恢复最佳权重并停止。def train_one_epoch(model, loader, opt): model.train() total_loss 0.0 for x, y, tod, wd in loader: x, y x.cuda(), y.cuda() tod, wd tod.cuda(), wd.cuda() pred model(x, tod, wd) # (B, T_out, N) loss torch.nn.functional.huber_loss(pred, y, delta1.0) opt.zero_grad() loss.backward() opt.step() total_loss loss.item() * len(x) return total_loss / len(loader.dataset)pred和y的 shape 都是(B, T_out, N)Huber 损失会在三个维度上做逐元素计算。如果loss.backward()返回 NaN大概率是输入里还有 inf 没清理干净用np.isfinite再扫一遍数据就能定位。如果 loss 正常但训练集 MAE 降到某个值后不动先看学习率是否需要衰减到原来的十分之一而不是直接加大模型。4.3 关键超参数速查很多复现「效果不好」的问题出在掩码维度不对或 batch 组织方式错误而非模型理论缺陷。下面这套参数来自我调 METR-LA 的通用模板可以作为首次运行的基准线超参数推荐值不推荐 / 易错输入 / 输出步长12 / 12输出误取前 T_out 步而非后 T_out 步d_model64256 在小数据集上容易过拟合num_layers2 ~ 34 层以上需更大数据量支撑学习率1e-3 预热后降到 1e-4全程固定 1e-4 收敛过慢dropout0.1 ~ 0.2设 0 时测试集 MAE 明显抬高归一化统计量仅训练集用全量统计会让测试结果失真在 Python 3.9 以上的环境运行即可不必刻意用最新版本PyTorch 1.13 之后的版本对TransformerEncoderLayer的mask参数处理一致代码可以直接迁移到 2.x。5. 评估与排错三个指标和三处值得验证的边界5.1 指标定义与关注点交通流预测中最常汇报的三个指标是 MAE、RMSE、MAPE。MAPE 对真实值为零的传感器是无穷大计算时必须把真实值小于 0.01 的点过滤掉否则一个传感器故障就能把整个 MAPE 拖到无法阅读。具体实现如下def evaluate(model, loader): mae rmae ape cnt 0.0 model.eval() with torch.no_grad(): for x, y, tod, wd in loader: pred model(x, tod, wd).cpu() mae (pred - y).abs().sum().item() rmae ((pred - y) ** 2).sum().item() mask y.abs() 0.01 ape ((pred - y).abs() / y.abs())[mask].sum().item() cnt mask.sum().item() n_sample len(loader.dataset) * 12 * 207 return mae / n_sample, (rmae / n_sample) ** 0.5, ape / cnt评估时容易忽略的一个点模型输出的是归一化之后的 Z-score一定要乘回数据集的 std 再加回 mean 再算指标。网上有些报告的 MAE 在 2 到 3 mph实际上是忘了还原直接拿标准化数据计算的结果数值小一个量级。还原之后再看METR-LA 测试集 MAE 在 13 到 14 mph 之间都是合理的模型差距主要出现在早晚高峰的半小时预测段。5.2 三个有效的验证与排错手段验证模型是否真的学到了时空相关性最简单的做法是保持配置不变训练集中随机抽 20% 时间步做测试。这时模型在没有未来数据的约束下理应成绩变差如果成绩反而更好说明训练集与测试集之间存在时间泄露回头检查滑窗 step 和 DataLoader 的 shuffle 顺序。第二个值得做的是极端高峰验证挑一个工作日早高峰的连续 3 小时切出测试集观察模型在 06:30 到 07:30 的误差是否明显大于夜间。时空变换网络应当比 GCN 在高峰期表现更稳健因为注意力可以跨过中间多个路段直接捕捉低速波的传递。如果这个优势没有出现大概率是掩码写错或注意力退化。第三个技巧是可视化注意力矩阵。取model.transformer.layers[0].self_attn的权重输出shape 是(B, num_heads, T*N, T*N)用 seaborn 画出来。如果对角线占绝对主导模型退化成纯时序自回归空间注意力没学到多半是 d_model 太小或数据没归一化如果注意力分布非常均匀说明传感器身份毫无区分度考虑加大节点嵌入维度或改用基于传感器坐标距离的高斯位置编码。改任何组件之前先跑一遍 4.3 参数表作为 baseline再动模型否则很难判断改动是提升还是引入抵消效应。改完记得固定测试集随机种子跑三遍取均值交通流数据的天数规律会带来波动只跑一次作结论很容易踩到高、低峰分布不均匀的坑。本文还有配套的精品资源点击获取