共计 2140 个字符,预计需要花费 6 分钟才能阅读完成。
交通流预测的稀疏数据挑战
在城市路网中,约 60% 的道路检测器存在间歇性失效问题。当早高峰出现交通事故时,传统基于固定传感器的模型(如 ARIMA)常因关键路段数据缺失导致预测完全失效。某直辖市交管局案例显示,5% 的检测器离线就会使 ST-ResNet 的 MAE 指标上升 37%。
方法论对比
| 模型 | 输入维度 | RMSE(5min) | MAE(15min) |
|---|---|---|---|
| ARIMA | 单点时序 | 23.4 | 18.7 |
| ST-ResNet | 网格划分 | 17.2 | 14.3 |
| ASTNN | 动态路网 | 11.8 | 9.2 |
关键改进在于 ASTNN 的两种注意力机制:
1. 空间注意力(Spatial Attention):动态学习路段关联权重 $\alpha_{ij}=\sigma(W_a[h_i||h_j])$
2. 时间注意力(Temporal Attention):捕获跨时间步依赖 $\beta_t=\text{softmax}(W_b\cdot\text{tanh}(U_bH))$
核心实现
动态图构建
import torch
from torch_geometric.utils import dense_to_sparse
def build_dynamic_graph(speed_data, threshold=0.7):
"""
speed_data: [T, N_nodes] 历史速度数据
returns: edge_index [2, E], edge_attr [E]
"""
# 计算 Pearson 相关性
corr_matrix = torch.corrcoef(speed_data.T) # [N, N]
# 生成邻接矩阵
adj = (corr_matrix > threshold).float()
adj.fill_diagonal_(0) # 移除自环
# 转换为稀疏表示
edge_index, edge_attr = dense_to_sparse(adj)
return edge_index, edge_attr
时空注意力模块
import torch.nn as nn
import torch.nn.functional as F
class MultiHeadSpatialAttention(nn.Module):
def __init__(self, node_feats, heads=4):
super().__init__()
self.heads = heads
self.d_k = node_feats // heads
# 线性变换层
self.W_q = nn.Linear(node_feats, node_feats)
self.W_k = nn.Linear(node_feats, node_feats)
self.W_v = nn.Linear(node_feats, node_feats)
def forward(self, h, edge_index):
"""
h: [N, node_feats] 节点特征
edge_index: [2, E] 边连接关系
"""
# 多头切分
Q = self.W_q(h).view(-1, self.heads, self.d_k) # [N, heads, d_k]
K = self.W_k(h).view(-1, self.heads, self.d_k)
V = self.W_v(h).view(-1, self.heads, self.d_k)
# 计算注意力分数
attn_scores = (Q @ K.transpose(-2,-1)) / torch.sqrt(torch.tensor(self.d_k))
# 应用邻接矩阵掩码
row, col = edge_index
mask = torch.zeros(h.size(0), h.size(0), device=h.device)
mask[row, col] = 1
attn_scores = attn_scores.masked_fill(mask.unsqueeze(1)==0, -1e9)
# 归一化并聚合
attn_weights = F.softmax(attn_scores, dim=-1)
h_out = (attn_weights @ V).transpose(1,2).reshape(-1, self.heads*self.d_k)
return h_out
性能优化实践
DGL 加速技巧
import dgl
def build_dgl_graph(edge_index, edge_attr):
g = dgl.graph((edge_index[0], edge_index[1]))
g.edata['w'] = edge_attr # 边权重
# 启用 CUDA 图优化
if torch.cuda.is_available():
g = g.to('cuda:0')
return g
稀疏矩阵优化
- 使用 COO 格式存储邻接矩阵
- 对静态路网部分预计算并缓存
- 采用
torch.sparse_coo_tensor减少内存占用
关键调参经验
| 参数 | 推荐范围 | 影响分析 |
|---|---|---|
| 注意力头数 | 4-8 | 头数过多易导致过拟合 |
| 学习率 | 1e-4 ~ 3e-3 | 配合梯度裁剪效果更佳 |
| 历史时间步 | 12~24 | 过长序列会引入噪声 |
数据集与运行环境
- 推荐数据集:Los-loop (PeMS) Kaggle 链接
- 一键运行:

通过实测,在 RTX 3090 上训练 1 个 epoch 约需 45 秒(batch_size=32)。建议先在小规模路网(<100 节点)验证模型效果,再逐步扩展到大区域路网。
正文完

