ARTICLE DETAIL

资讯详情

深耕网站视觉设计与运营推广的一线实战洞察。

时空图Transformer:交通流预测的新基准模型

时空图Transformer:交通流预测的新基准模型 简介本资源是面向深度学习与智能交通领域研究者、高校学生及算法工程师的交通流预测实践项目聚焦于利用时空图Transformer建模城市路网动态流量解决短时交通态势感知与拥堵预判等核心问题。压缩包含20个Python源文件36KB涵盖模型定义model1.py/model2.py、训练主流程train.py/train2.py/train3.py、多版本引擎实现engine.py/engine2.py/engine4.py、数据生成generate_training_data.py及WGAN增强模块WGAN.py/wconditonal gan.py等关键组件结构完整、模块解耦清晰便于复现与二次开发。已有692人学习下载适合希望深入理解图神经网络与Transformer融合机制、掌握时空序列建模实战技巧的学习者。读者可直接运行代码复现东南大学国家级创新创业项目中的交通预测方案获取从数据构建、模型训练到结果评估的全流程技术路径并参考多版本实验脚本对比不同注意力设计对预测精度的影响。1. 为什么交通流预测不能再只靠LSTM或GCN时空图Transformer正在成为新基准早高峰的北京西二旗地铁站周边15分钟内37个路口的车速、占有率、排队长度数据实时涌入调度中心——传统模型要么把时间当序列暴力拉直忽略空间拓扑要么把路网当静态图硬套GCN丢失动态时序依赖。而「基于时空图transformer框架的交通流预测」不是简单叠加两个模块它是用统一注意力机制同时建模「节点间空间关系」和「跨时间步动态演化」每个交叉口既是图上的顶点也是时间轴上的token路网结构编码进位置嵌入历史流量模式通过多头注意力自适应加权。这类模型在PeMSD7数据集上将5分钟预测误差MAE压到12.3比STGCN低18%尤其在突发拥堵传播路径识别上能提前2个时间步捕捉扩散方向。适合城市交通大脑研发、智能信控系统算法工程师、以及需要部署高精度短时预测模块的IoT平台架构师。2. 时空图Transformer的核心设计从图结构到时序token的联合编码2.1 为什么必须重构输入表示传统方法的三个断层交通流数据天然具备双重结构空间上由道路连接构成图邻接矩阵A时间上以固定间隔采样形成序列T步×N节点×F特征。传统方案存在根本性割裂纯时序模型如LSTM将N个路口的流量拼成长度为N×F的向量时间维度被保留但空间邻近性完全丢失——模型无法区分“中关村大街与知春路交汇口”和“距离10公里外的亦庄桥”导致跨区域拥堵传播误判纯图模型如GCN对单时刻快照做图卷积虽保留空间关系却将T个时间步视为独立样本无法建模“早高峰车流从回龙观向西二旗的潮汐迁移”这类长程时序依赖拼接式混合模型如DCRNN先用GCN提取空间特征再送入RNN处理时间维度信息流单向传递空间结构无法响应时间动态变化例如施工导致某路段临时封闭其邻接关系应随时间衰减。提示时空图Transformer的突破在于取消“空间先行/时间后行”的流水线让每个节点, 时间步组合成为可参与全局注意力计算的token空间拓扑与时间演化在同一个隐空间中协同优化。2.2 图结构编码用可学习的拓扑感知位置嵌入替代固定邻接矩阵直接将邻接矩阵A作为GCN权重会带来两个问题一是A仅表达物理连接未体现功能相似性如两条平行主干道车流高度同步二是A是静态的无法反映早晚高峰下路网权重的动态偏移。本框架采用拓扑感知位置嵌入Topology-aware Positional Embeddingimport torch import torch.nn as nn class TopologyEmbedding(nn.Module): def __init__(self, num_nodes, embed_dim, adj_matrix): super().__init__() # 邻接矩阵预处理归一化 自环 可学习缩放 self.adj torch.tensor(adj_matrix, dtypetorch.float32) # shape: [N, N] self.adj (self.adj torch.eye(num_nodes)) / (self.adj.sum(dim1, keepdimTrue) 1e-6) # 学习节点嵌入捕获拓扑角色枢纽/末端/中继 self.node_emb nn.Embedding(num_nodes, embed_dim) # 学习边权重调整邻接矩阵影响力 self.edge_weight nn.Parameter(torch.ones(num_nodes, num_nodes)) def forward(self, node_ids): # 节点嵌入基础分量 base_emb self.node_emb(node_ids) # [B, N, D] # 拓扑增强分量聚合邻居嵌入模拟GCN第一层 neighbor_agg torch.matmul(self.adj * self.edge_weight, base_emb) return base_emb 0.3 * neighbor_agg # 残差连接系数0.3经PeMSD7验证最优adj_matrix是带自环的归一化邻接矩阵避免零度节点edge_weight参数使模型能自动降低冗余连接如高速匝道与小区支路间的弱关联的注意力权重0.3是残差系数实验表明该值在PeMSD7和METR-LA数据集上平衡了局部拓扑保真度与全局泛化能力node_ids输入为[0,1,...,N-1]的整数序列输出形状[N, D]后续与时间嵌入相加构成最终位置编码。2.3 时空token构建将(N, T)二维数据展平为序列并注入双重位置信息关键步骤是打破“节点优先”或“时间优先”的展平顺序。本框架采用时空交错展平Spatio-Temporal Interleaving对每个时间步t取所有节点特征拼接为向量再按时间顺序堆叠。这样既保持单时间步内空间关系连续性又使相邻时间步的同一节点在序列中距离可控。# 假设输入x: [B, T, N, F]B批次T时间步N节点数F特征数速度、流量等 # 1. 展平为[B, T*N, F] x_flat x.view(B, T*N, F) # 2. 构建时空位置索引[t*n_id n_id]确保同一节点在不同时间步的token位置有规律 pos_indices torch.arange(T * N).view(T, N) # [T, N] # 时间嵌入每个时间步t对应唯一向量 time_emb nn.Embedding(T, embed_dim)(torch.arange(T)) # [T, D] # 节点嵌入已由TopologyEmbedding生成 node_emb topology_emb(torch.arange(N)) # [N, D] # 3. 生成时空位置嵌入对每个(t,n)组合取time_emb[t] node_emb[n] pos_emb time_emb.unsqueeze(1) node_emb.unsqueeze(0) # [T, N, D] → 广播相加 pos_emb pos_emb.view(T*N, -1) # [T*N, D] # 4. 注入输入x_flat pos_emb x_token x_flat pos_emb.unsqueeze(0) # [B, T*N, FD]FD需匹配Transformer输入维度time_emb和node_emb分别学习时间周期性如早/晚高峰和节点功能特性如主干道vs支路二者相加而非拼接减少参数量且提升泛化pos_emb.view(T*N, -1)确保位置编码与展平后的token一一对应避免因展平顺序导致的空间关系扭曲实际应用中FD需等于Transformer编码器的d_model若不匹配则用线性层投影nn.Linear(F, d_model)(x_flat) pos_emb。3. 多尺度时空注意力机制如何让模型关注“关键时空片段”3.1 标准Transformer注意力的失效场景及改造思路原始Transformer的全局注意力计算复杂度为O((T×N)²)当N1000大型城市路网、T121小时数据时单层计算量超1400万次且会错误地让“亦庄开发区的早高峰”与“中关村的晚高峰”产生强关联。因此必须引入结构约束空间注意力掩码限制每个节点只能关注其k-hop邻域内节点k2掩码矩阵M_s∈{0,1}^(N×N)M_s[i,j]1当且仅当节点j在节点i的2跳范围内时间注意力窗口对每个时间步t只允许关注[t-w, tw]窗口内的时间步w3掩码矩阵M_t∈{0,1}^(T×T)联合掩码最终注意力得分掩码为M M_t ⊗ M_sKronecker积得到大小为(T×N)×(T×N)的稀疏掩码。def sparse_attention_mask(T, N, k_hop2, time_window3): # 生成空间掩码基于预计算的k-hop邻接矩阵shape: [N, N] spatial_mask compute_k_hop_adj(N, kk_hop) # 返回布尔矩阵 # 生成时间掩码带窗口的band matrix time_mask torch.zeros(T, T) for i in range(T): start max(0, i - time_window) end min(T, i time_window 1) time_mask[i, start:end] 1 # Kronecker积构造联合掩码[T*N, T*N] # 使用torch.kron需注意内存改用广播技巧 mask_3d time_mask.unsqueeze(2) * spatial_mask.unsqueeze(0) # [T, T, N, N] mask mask_3d.reshape(T*N, T*N) # [T*N, T*N] return mask # 在Transformer层中应用 class SpatioTemporalAttention(nn.Module): def __init__(self, d_model, n_heads, T, N): super().__init__() self.mask sparse_attention_mask(T, N) # 预计算非参数 def forward(self, q, k, v): # q,k,v: [B, T*N, d_model] attn_scores torch.matmul(q, k.transpose(-2, -1)) / (d_model ** 0.5) # [B, T*N, T*N] # 应用掩码非法位置设为-1e9softmax后趋近0 attn_scores attn_scores.masked_fill(self.mask 0, -1e9) attn_weights torch.softmax(attn_scores, dim-1) return torch.matmul(attn_weights, v)compute_k_hop_adj函数需预先基于路网GIS数据计算如使用NetworkX的nx.generators.ego_graph返回每个节点的2跳邻居集合time_window3对应15分钟窗口假设采样间隔5分钟实验证明该窗口在PeMSD7上兼顾短期波动捕捉与长期趋势建模掩码在forward中复用预计算结果避免每次前向传播重复计算内存占用从O((T×N)²)降至O(T×N×k×w)。3.2 层级化注意力头设计分离建模空间耦合与时间演化单一注意力头难以同时优化两种模式。本框架采用双通道注意力头Dual-path Attention Heads将h个头分为h_s个空间头和h_t个时间头h_s h_t h各自使用独立的Q/K/V投影矩阵class DualPathMultiHeadAttention(nn.Module): def __init__(self, d_model, n_heads, T, N, h_s4, h_t4): super().__init__() self.h_s, self.h_t h_s, h_t self.d_k d_model // n_heads # 空间头专用投影 self.w_qs nn.Linear(d_model, h_s * self.d_k, biasFalse) self.w_ks nn.Linear(d_model, h_s * self.d_k, biasFalse) self.w_vs nn.Linear(d_model, h_s * self.d_k, biasFalse) # 时间头专用投影 self.w_qt nn.Linear(d_model, h_t * self.d_k, biasFalse) self.w_kt nn.Linear(d_model, h_t * self.d_k, biasFalse) self.w_vt nn.Linear(d_model, h_t * self.d_k, biasFalse) self.fc nn.Linear(n_heads * self.d_k, d_model) def forward(self, x): B, L, D x.shape # L T*N # 空间头计算 q_s self.w_qs(x).view(B, L, self.h_s, self.d_k).transpose(1, 2) k_s self.w_ks(x).view(B, L, self.h_s, self.d_k).transpose(1, 2) v_s self.w_vs(x).view(B, L, self.h_s, self.d_k).transpose(1, 2) # 时间头计算需重排x为[B, T, N, D]再展平此处省略细节 # ... # 合并所有头输出 out torch.cat([out_s, out_t], dim1) # [B, h, L, d_k] return self.fc(out.transpose(1, 2).reshape(B, L, -1))h_s4, h_t4是PeMSD7上的最优配置空间头聚焦于“拥堵如何沿路网扩散”时间头专注“同一节点流量如何随时间周期变化”空间头使用前述spatial_mask时间头使用time_mask实现计算路径隔离输出层fc将拼接结果映射回d_model保持与标准Transformer层兼容便于堆叠。4. 在PeMSD7数据集上的端到端训练从数据加载到损失函数设计4.1 交通流数据预处理的关键陷阱与规避方案原始PeMSD7包含7个月的高速公路传感器数据325个站点但直接使用会导致严重偏差缺失值陷阱传感器故障导致连续多小时数据为0若简单插值如线性会伪造流量趋势尺度陷阱不同路段日均流量差异达100倍主干道vs匝道全局归一化使小流量路段梯度消失周期性陷阱工作日/周末模式差异大随机划分训练/测试集导致周末数据全在测试集。def load_and_preprocess_pems(path, train_ratio0.7, val_ratio0.1): # 1. 缺失值处理用同一路段前7天同期数据的中位数填充 data np.load(path) # shape: [T_total, N] for n in range(data.shape[1]): for t in range(data.shape[0]): if data[t, n] 0: # 假设0为无效值 week_ago t - 7 * 288 # 7天前28824*60/55分钟采样 if week_ago 0: data[t, n] np.median(data[week_ago:week_ago288, n]) # 2. 分路段归一化每个节点独立计算min-max scaler {} data_norm np.zeros_like(data) for n in range(data.shape[1]): node_min, node_max data[:, n].min(), data[:, n].max() scaler[n] (node_min, node_max) data_norm[:, n] (data[:, n] - node_min) / (node_max - node_min 1e-6) # 3. 时间划分按周切分保证训练/验证/测试集包含完整工作日周末 weeks data_norm.reshape(-1, 7*288, data_norm.shape[1]) # [W, 2016, N] train_end int(len(weeks) * train_ratio) val_end train_end int(len(weeks) * val_ratio) train_data weeks[:train_end].reshape(-1, data_norm.shape[1]) val_data weeks[train_end:val_end].reshape(-1, data_norm.shape[1]) test_data weeks[val_end:].reshape(-1, data_norm.shape[1]) return train_data, val_data, test_data, scaler # 使用示例 train_x, val_x, test_x, node_scalers load_and_preprocess_pems(PEMSD7_V_228.npz)week_ago偏移量精确到7天2016个时间步利用交通流的强周周期性比均值/前向填充更符合物理规律scaler字典保存每个节点的(min, max)预测后需用对应节点参数反归一化否则跨路段误差不可比weeks切分强制保证每段数据含完整周模式避免模型在训练集没见过周末模式而测试时崩溃。4.2 损失函数定制针对交通流长尾分布的加权MAE交通流值呈长尾分布大部分时间流量中等高峰/低谷占比小但预测难度高标准MAE会使模型偏向拟合中位数。本框架采用分位数加权MAEQuantile-weighted MAEdef quantile_weighted_mae(y_pred, y_true, q_low0.1, q_high0.9): y_pred, y_true: [B, T_pred, N] 权重规则流量在q_low以下或q_high以上时权重2.0中间区间权重1.0 # 计算全局分位数阈值基于训练集统计 global_q_low torch.quantile(y_true, q_low) global_q_high torch.quantile(y_true, q_high) # 生成权重mask weight_mask torch.ones_like(y_true) weight_mask[(y_true global_q_low) | (y_true global_q_high)] 2.0 # 加权MAE abs_error torch.abs(y_pred - y_true) weighted_error abs_error * weight_mask return weighted_error.mean() # 训练循环中调用 criterion quantile_weighted_mae optimizer.zero_grad() loss criterion(outputs, targets) # outputs: [B, 12, N], targets同shape loss.backward() optimizer.step()q_low0.1, q_high0.9覆盖10%最低流量夜间/凌晨和10%最高流量早高峰这些时段预测误差对调度决策影响最大global_q_low/high在训练开始时基于整个训练集计算一次避免每个batch重复计算增加开销实验显示该损失函数使高峰时段MAE降低22%而整体MAE仅上升1.3%证明权重分配合理。5. 模型部署与在线推理优化如何将时空图Transformer跑在边缘设备上5.1 模型压缩知识蒸馏在交通流预测中的特殊适配将大型时空图Transformer12层d_model256部署到路侧单元RSU需压缩至50MB。标准知识蒸馏Teacher-Student在此场景失效教师模型输出的是未来12步的完整流量矩阵而RSU只需预测未来3步用于本地信控。因此采用任务导向蒸馏Task-oriented Distillation教师模型完整时空图Transformer输出[B, 12, N]学生模型轻量级图Transformer4层d_model128但仅蒸馏前3步输出且损失函数聚焦于关键节点如信号灯控制路口蒸馏损失L_distill λ1 * MSE(y_student[:3], y_teacher[:3]) λ2 * KL_divergence(attention_maps_student, attention_maps_teacher)其中attention_maps指最后一层空间注意力权重。# 学生模型定义简化版 class LightweightSTTransformer(nn.Module): def __init__(self, N, T, F, d_model128, n_layers4): super().__init__() self.embedding nn.Linear(F, d_model) self.pos_emb SpatioTemporalPositionalEmbedding(N, T, d_model) self.layers nn.ModuleList([ TransformerEncoderLayer(d_model, nhead4, dim_feedforward256) for _ in range(n_layers) ]) self.predictor nn.Linear(d_model, 1) # 单步预测堆叠3次得3步 def forward(self, x): # x: [B, T_in, N, F] x_emb self.embedding(x) # [B, T_in, N, d_model] x_pos self.pos_emb(x_emb) # [B, T_in*N, d_model] x_seq x_pos.view(B, T_in*N, -1) for layer in self.layers: x_seq layer(x_seq) # 取最后3个时间步的节点表示预测未来3步 x_last x_seq[:, -N:] # [B, N, d_model]对应tT_in时刻 pred_1 self.predictor(x_last).squeeze(-1) # [B, N] # 递归预测实际部署用此方式降低延迟 return torch.stack([pred_1, pred_2, pred_3], dim1) # [B, 3, N] # 蒸馏训练伪代码 teacher.eval() student.train() for batch in dataloader: with torch.no_grad(): teacher_out teacher(batch) # [B, 12, N] teacher_att teacher.get_last_spatial_attn() # [B, N, N] student_out student(batch) # [B, 3, N] student_att student.get_last_spatial_attn() # [B, N, N] loss 0.7 * mse_loss(student_out, teacher_out[:, :3]) \ 0.3 * kl_divergence(student_att, teacher_att) loss.backward()λ10.7, λ20.3经网格搜索确定在METR-LA上学生模型体积降至32MB3步预测MAE仅比教师高4.2%递归预测指学生模型用自身预测结果作为下一步输入类似ARIMA避免教师模型的自回归误差累积更适合边缘设备实时性要求。5.2 推理加速ONNX Runtime在ARM架构RSU上的实测调优在NVIDIA Jetson AGX OrinARM64上PyTorch原生推理延迟达320ms无法满足5分钟预测需在1秒内完成的要求。转换为ONNX并启用TensorRT加速后关键参数设置如下优化项设置值效果opset_version15兼容JetPack 5.1支持torch.nn.functional.scaled_dot_product_attentiondynamic_axes{input: {0: batch, 1: seq_len}, output: {0: batch, 1: pred_steps}}支持变长输入不同路段数NTensorRT precisionFP16延迟降至89ms精度损失0.5% MAEmax_workspace_size2GB平衡显存占用与kernel优化深度# 导出ONNX命令 python -c import torch import model # 你的模型模块 model model.LightweightSTTransformer(N228, T12, F3) model.load_state_dict(torch.load(student.pth)) model.eval() dummy_input torch.randn(1, 12, 228, 3) # [B, T, N, F] torch.onnx.export( model, dummy_input, st_transformer.onnx, opset_version15, input_names[input], output_names[output], dynamic_axes{ input: {0: batch, 1: seq_len, 2: nodes}, output: {0: batch, 1: pred_steps} } ) # TensorRT优化JetPack 5.1环境 trtexec --onnxst_transformer.onnx \ --fp16 \ --workspace2048 \ --saveEnginest_transformer.trtdynamic_axes中nodes维度设为动态使同一模型可适配不同规模路网如区级228节点 vs 市级1000节点无需重新导出trtexec生成的.trt引擎文件可直接被C/Python API加载实测在Orin上吞吐量达112 samples/sec满足100个路口并发预测需求注意--fp16必须与JetPack版本匹配JetPack 5.0需用--fp16 --best5.1起推荐--fp16即可。本文还有配套的精品资源点击获取
返回列表