时空变换网络:自注意力驱动的交通流预测新方法,无需邻接矩阵
2026/9/16 23:44:23 网站建设 项目流程

简介:面向交通流预测领域研究者和深度学习开发者,提供时空变换网络(ST-Transformer)的完整Python实现与配套数据集。模型结合时空卷积模块和注意力机制,能够捕捉路网动态时空依赖,适用于城市交通流量预测、拥堵预警与智慧交通调度场景。资源包共9个文件,其中6个.py脚本覆盖模型定义(ST_Transformer.py、layers.py、GCN_models.py)、训练验证(train.py、validation.py)与数据处理(One_hot_encoder.py),2个CSV文件为PEMSD7路网车流量数据(V_25.csv)和邻接矩阵(W_25.csv),另有README指导使用,整体仅451KB,轻量易部署。已有387人学习下载。通过该资源可深入理解ST-Transformer的架构细节、时空特征提取与注意力权重的实现方式,掌握数据预处理、模型训练、性能评估的完整流程,还可基于自带数据集复现预测效果,并可自行替换数据进一步改进模型。

1. 交通流预测的时空变换网络:不再依赖预定义邻接矩阵

交通流预测通常被建模为:给过去一小时的数据,预测未来一小时的速度或流量。过去五年的标配组合是图卷积加循环神经网络,但它有一个硬约束——必须手工准备邻接矩阵,并按固定拓扑建图。路网一旦新增传感器,高速路遇到临时封路或大型活动导致流量路径改变,这张静态图就失真,图卷积的泛化也随之下降。时空变换网络不预设图结构,而是让模型用自注意力从数据里学习哪个传感器、哪个时刻在相互影响。覆盖面从上游路段直接跨到下游几个街区,比邻接矩阵的语义更贴合真实拥堵的传播方式。

围绕这套方案写的源码由三部分组成:时空编码、带掩码的多头自注意力、多步回归头。下面以高速公路传感器数据集 METR-LA 为样例,把数据准备、模型搭建到训练评估完整讲一遍,并给出可以直接复用的参数配置。

2. 时空变换网络核心模块与编码器选型

2.1 token化:把传感器读数变成带时空信息的序列

Transformer 处理的是「词序列」,而交通流预测中序列里的每个元素是「某个传感器在某个时刻的速度数值」。直接把原始标量丢进注意力层有两个问题:一是丢失位置与传感器身份,二是数值尺度差异过大。因此常见做法是为每个观测样本构造一个 token,token 的维度和传感器数量无关,只和嵌入维度 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, device=x.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 的 N=207、T=12,全时空注意力一次要计算约六万个 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, dtype=torch.bool)) causal = causal.repeat_interleave(N, dim=0).repeat_interleave(N, dim=1) return causal # (T*N, T*N)

repeat_interleave的含义是把时间轴上的因果约束复制到每个传感器上:传感器 i 在 t1 时刻可以看传感器 j 在 t2≤t1 时刻的数据,只要 t2>t1 就被挡住。这一步如果漏做,指标会显得异常好,实测 MAE 可能直接下降 15% 以上,且模型完全无法用于在线推理,因为它偷看了未来。如果后续发现输出全是训练集均值,再检查一下掩码是不是和序列排列顺序不一致,常见错误是把 flatten 顺序和掩码构建顺序搞混。

2.3 编码器堆叠与残差结构

自注意力本身不改变数据维度,深层结构主要靠前馈网络和残差完成非线性变换。常见做法是堆 2 到 4 个编码器层,每层由「多头自注意力 → LayerNorm → 两层 MLP → LayerNorm」串成。相比 NLP 任务动辄 12 层、24 层,交通序列的层数不需要太深;堆到 6 层以上反而会在小数据集上出现注意力退化,所有查询向量收敛到几乎相同的分布,多头退化成单头。

下表是 METR-LA 这类中等规模路网上的初始配置范围,也适合用来给后续源码设定 baseline:

参数取值说明
d_model64 / 128数据量小且显存紧时用 64
num_heads4 / 8d_model 必须能被 head 数整除
encoder_layers2 / 3超过 4 层收益开始递减
dropout0.1 / 0.2大数据集用 0.1,小数据集用 0.2
feedforward_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 min
PEMS043072018-01 ~ 025 min
PEMS081702016-07 ~ 085 min

动手写模型前,第一件事是确认数据是否已经做了缺失值插补。公开版 METR-LA 已过滤一部分无效记录,但仍有约 2% 的 NaN 或 inf。处理缺失值的顺序是:先做最近邻插补,再做标准化,最后进模型。反过来操作会把缺失值信息泄露给归一化层的统计量,导致测试集指标偏乐观。

切分方式沿用交通流领域的惯例:按时间顺序切为 70% 训练、10% 验证、20% 测试。与图像分类不同,交通流样本之间强相关,随机洗牌会让同一时段出现在训练集和测试集里,严格时间顺序切分才能保证「模型预测的是没见过的时间段」。这组比例在 METR-LA 上的测试集大约覆盖最后一个半月,难度比前段更大。

3.2 滑窗采样与归一化

给定一个固定窗口长度 T_in=12(过去一小时)和预测步长 T_out=12(未来一小时),滑窗会把长为样本数的时序矩阵切成无数个重叠样本。重叠是交通流预测中刻意设计的,它让模型能在有限数据上见到更多模式;缺点是相邻样本高度相似,如果不用 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_in=12, T_out=12, step=1): 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)

这里step=1意味着原始 207 个传感器约 34000 个时间步会被切成三万个重叠样本,训练随机打乱;换成step=6后样本数直接降到原来的六分之一,训练快 6 倍,但指标通常略差。优先推荐 step=1 配 batch_size 64,若显存或时间受限再增大 step。X 的形状是(样本数, 12, 207),Y 的形状是(样本数, 12, 207),模型要对 12 个未来时刻、所有传感器输出预测值。

3.3 构造带时间特征的 DataLoader

sliding_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_in=12): 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_size=64, shuffle=True, num_workers=4),训练循环就能稳定跑满 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_model=64, num_heads=4, num_layers=3, dropout=0.1): super().__init__() self.embed = SpatioTemporalEmbedding(num_nodes, d_model) encoder_layer = nn.TransformerEncoderLayer( d_model=d_model, nhead=num_heads, dim_feedforward=d_model*4, dropout=dropout, activation='gelu', batch_first=True ) 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, mask=mask) 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,训练若只优化 MAE,RMSE 会偏高;我一般用 Huber Loss,delta=1.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, delta=1.0) opt.zero_grad() loss.backward() opt.step() total_loss += loss.item() * len(x) return total_loss / len(loader.dataset)

predy的 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 之后的版本对TransformerEncoderLayermask参数处理一致,代码可以直接迁移到 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,再动模型,否则很难判断改动是提升还是引入抵消效应。改完记得固定测试集随机种子跑三遍取均值,交通流数据的天数规律会带来波动,只跑一次作结论很容易踩到高、低峰分布不均匀的坑。

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

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

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

立即咨询