简介:面向交通流预测与时空图神经网络研究者的ASTGCN算法Python实现资源包,包含PEMS04与PEMS08两组真实高速公路交通数据集的相关配置及数据,适合学习复现基于注意力机制的时空图卷积网络。资源围绕模型定义、训练、测试、数据预处理、评估指标等完整流程展开,源码中既有核心的ASTGCN网络结构,也有对应的MSTGCN变体,并配有数据处理、工具函数和指标计算等辅助脚本,另有两篇原始论文PDF便于对照理解。压缩包共有二十五个文件,大小约为二十点七一兆字节,以Python源码、说明文档和论文文件为主,目录结构清楚,可直接用于交通预测实验与二次开发。目前已有1982人学习使用,推荐给具备一定深度学习基础、希望深入理解或改进图卷积交通流预测模型的研究者。
1. ASTGCN 解决的不只是“预测”,而是路网级时空依赖建模
把某个检测器的历史数据丢给 LSTM 做交通流预测,跑一阵子你会发现精度很快碰到天花板——序列模型看得到时间维度,但路网本质是图结构,单点模型感知不到上下游的联动关系。ASTGCN 这个基于注意力机制的时空图卷积网络,就是冲着这个问题去的:它把路网当成一张带拓扑结构的图,用图卷积消化空间依赖,用门控时间卷积消化时间依赖,再用注意力机制动态调整每个节点对邻居的关注强度。配合加州 PEMS 系统公开的 PEMS04(307 个节点)和 PEMS08(170 个节点)两份高速公路检测器数据,它在 Python 生态里成了交通流深度学习方向被复现最多的基线模型之一。这篇笔记适合正在复现论文、做课设或者刚切换到时空图神经网络方向的人,我把数据格式、模型拆解、训练配置和踩过的坑一次讲完。
2. 读懂 PEMS04 与 PEMS08 数据,先把图结构建对
2.1 PEMS04 和 PEMS08 数据集的字段构成
交通流预测是深度学习在智能交通里落地最早的场景之一,而 PEMS04 与 PEMS08 几乎成了国内做这个方向的标准数据集。PEMS04 覆盖旧金山湾区 307 个检测器,PEMS08 覆盖圣贝纳迪诺地区 170 个检测器,采样间隔都是 5 分钟,一天 288 个点位,数据维度通常是节点数、时间步、特征数三维。
数据包一般包含两个关键文件:一个是 data.npz,里面是检测器采集到的连续数值;另一个是 adj.npz,存的是节点间的关系矩阵。PEMS04 的邻接矩阵通常是一个 (307, 307, 3) 的张量,第三维分别对应距离、交通相似性等不同语义的权重;PEMS08 的结构类似,只是节点数变成 170。第一次拿到数据包时,我建议先不要急着建模,而是把两个 npz 的 shape 和数据类型完整打印一遍:
import numpy as np data = np.load("data.npz") adj = np.load("adj.npz") print(data.files) print(data["data"].shape, data["data"].dtype) print(adj.files) print(adj["adj"].shape, adj["adj"].dtype)加载后立刻能看到两个最容易踩的坑:一是 data 的组织顺序可能是 (T, N, D) 而不是 (N, T, D),二是 adj 的第三维多通道结构往往被人忽略,直接用单通道距离矩阵代替会丢失相似度信息。这一步确认清楚,后面的滑动窗口切片和模型输入维度才不会写错。
2.2 邻接矩阵为什么要重新归一化
拿到原始距离矩阵后,大多数复现版本会把它转换成 0 到 1 之间的权重,而不是直接用距离值。常见做法是用高斯核函数做一次非线性变换:距离越近的节点权重越接近 1,距离越远的权重指数衰减。我这里用 σ 取 0.1 做一个参考实现,你也可以按数据集中距离矩阵的中位数来估计 σ 的取值:
def distance_to_weight(dist_matrix, sigma=0.1): dist_matrix = dist_matrix.astype(np.float32) # 高斯核:距离越近权重越大 weight = np.exp(-(dist_matrix ** 2) / (sigma ** 2)) # 自身到自身的权重置 0,避免自环干扰图卷积 np.fill_diagonal(weight, 0.0) return weight这里有个容易被忽略的参数语义:sigma 控制着空间影响的衰减速度。sigma 太小,只有紧邻的检测器之间有联系,图几乎退化成每条道路独立的线;sigma 太大,整个路网所有节点互相都有弱连接,图卷积的局部性就被稀释了。我一般先用距离矩阵的分位数扫描几个 sigma,然后看训练集上验证 loss 的变化趋势来确定,不要一上来就在完整训练流程里调。
2.3 滑动窗口切分:预测未来一小时需要多大窗口
PEMS 数据的标签构造方式和普通分类任务不同,本质上是一个序列到序列的预测任务。原始论文的设定是用过去一小时(12 个时间步)预测未来一小时,时间步长对应 5 分钟。切片原则是严格按时间顺序滑动,每个样本包含一个输入窗口和一个标签窗口。
def create_samples(data, in_steps=12, out_steps=12): # data 期望 shape: (N, T, D) N, T, D = data.shape inputs, labels = [], [] for start in range(T - in_steps - out_steps + 1): x_end = start + in_steps y_start = x_end y_end = y_start + out_steps inputs.append(data[:, start:x_end, :]) # (N, in_steps, D) labels.append(data[:, y_start:y_end, :]) # (N, out_steps, D) return np.array(inputs), np.array(labels)切分后得到的输入 shape 是 (样本数, 307, 12, 3),输出 shape 是 (样本数, 307, 12, 3),其中最后一维是特征数。这里要特别强调的是,交通流预测的特征维通常包含流量、速度、占有率三个物理量,早期有些复现只取了流量单通道,效果差别非常大——速度信息对拥堵传播的建模有决定性作用。切片完成后,按时间顺序把样本切成训练集、验证集、测试集,常见比例是 6:2:2,并且禁止随机打乱后切分,否则会造成严重的数据泄漏。
3. ASTGCN 模型拆解:注意力机制、图卷积和时间卷积如何协作
3.1 为什么路网不能用普通 CNN 建模
图像是欧几里得空间里的网格结构,卷积核在像素点上做滑动窗口很自然。但路网里每个检测器的邻居数量不一样,分布也不规则,用固定大小的卷积核无法覆盖不同节点周围完全不同的拓扑结构。图的谱理论提供了一个替代方案:把图信号变换到频域,用拉普拉斯矩阵的特征分解定义卷积操作,这就是图卷积的由来。
ASTGCN 在这个基础上又加了两层设计:空间注意力机制让模型在计算每个节点特征时动态决定该看哪些邻居,时间注意力机制帮助模型在长时间依赖中找到关键的历史时刻。这一点对交通场景特别重要——早高峰时期相邻路段相关性极强,而凌晨时段的模式完全不同,静态权重不足以描述这种动态变化。
3.2 空间注意力:节点之间动态计算关注权重
空间注意力的核心计算流程可以概括为:把输入特征先做两次线性变换,用矩阵乘法计算节点两两之间的相关性分数,再经过 softmax 归一化成注意力权重矩阵,最后用这个权重矩阵去加权原始的邻接矩阵。
class SpatialAttention(nn.Module): def __init__(self, in_channels, num_nodes, num_steps): super().__init__() self.W1 = nn.Parameter(torch.randn(num_steps)) self.W2 = nn.Parameter(torch.randn(in_channels, num_steps)) self.W3 = nn.Parameter(torch.randn(in_channels)) self.bias = nn.Parameter(torch.randn(num_nodes, num_nodes)) def forward(self, x): # x: (B, N, T, C) B, N, T, C = x.shape # 压缩时间维度后计算相关性 x1 = torch.einsum("bntc,t->bntc", x, self.W1) x2 = torch.einsum("bntc,ct->bnc", x, self.W2) score = torch.einsum("bntc,c->bnt", x, self.W3) score = torch.mean(score, dim=1, keepdim=True) # 对批次取平均 attention = torch.softmax(score + self.bias.unsqueeze(0), dim=-1) return attention这段代码里的 einsum 操作是关键,它把输入特征分别向时间维和通道维做投影,最后得到每个节点对所有其他节点的注意力分数。实际工程里,很多人简化这一步,直接用邻接矩阵做 softmax,模型精度会下降但也能跑通。我的建议是先把完整版跑通,确认精度符合预期再考虑简化,否则很难判断到底是注意力模块出了问题还是数据预处理出了问题。
3.3 图卷积网络与切比雪夫多项式的配合
谱域图卷积在工程上很少直接做拉普拉斯矩阵的特征分解,因为大规模路网下特征分解的计算代价太高。ASTGCN 采用切比雪夫多项式近似,用 K 阶多项式展开逼近卷积核,只需要做矩阵乘法,不需要做特征分解。K 的取值意味着一个节点最多能感知到 K 跳范围内的邻居信息。
def chebyshev_polynomials(L_tilde, K): N = L_tilde.shape[0] T_0 = torch.eye(N, device=L_tilde.device) T_1 = L_tilde.clone() polys = [T_0, T_1] for k in range(2, K): polys.append(2 * L_tilde @ polys[-1] - polys[-2]) return polys def normalize_laplacian(A): # A: (N, N) 邻接矩阵 D = torch.sum(A, dim=1) + 1e-5 # 加 epsilon 防止零除 D_inv_sqrt = torch.diag(1.0 / torch.sqrt(D)) L = torch.eye(A.shape[0], device=A.device) - D_inv_sqrt @ A @ D_inv_sqrt return LPEMS04 有 307 个节点,K 一般取 3,也就是每个节点聚合到三阶邻居的信息。K 太大会让特征过度平滑,所有节点收敛成相似表达,K 太小则空间信息吃不够。图卷积层的输出通道数一般跟随输入通道数,在 ASTGCN 原文里图卷积层输出的通道数和时间卷积层保持对齐,这样残差连接可以直接相加。
3.4 门控时间卷积:比 LSTM 更高效的时间依赖建模
原始 ASTGCN 使用门控时间卷积来处理时间维,本质上是一维卷积配合门控线性单元。一维卷积相比于 LSTM 有两个优势:训练时可以并行计算,不容易出现梯度消失。门控机制让模型学会决定历史信息中哪些部分要保留、哪些部分要过滤。
一个卷积核负责生成候选特征,另一个卷积核负责生成门控信号,两者逐元素相乘得到输出。在 PyTorch 里可以用两个 nn.Conv2d 实现,卷积核的第一个维度覆盖时间轴。
class TemporalConv(nn.Module): def __init__(self, in_channels, out_channels, kernel_size=3): super().__init__() self.conv1 = nn.Conv2d(in_channels, out_channels, (kernel_size, 1), padding=(kernel_size//2, 0)) self.conv2 = nn.Conv2d(in_channels, out_channels, (kernel_size, 1), padding=(kernel_size//2, 0)) def forward(self, x): # x: (B, C, N, T) gate = torch.sigmoid(self.conv2(x)) return self.conv1(x) * gatekernel_size 控制每个时间步能看到的过去范围。PEMS 数据 12 个时间步对应 1 小时,建议 kernel_size 取 3,两层堆叠之后感受野能覆盖到 5 个时间步,对小时内的短时突变有足够感知力。时间卷积通常配合残差连接,避免深层网络退化。
3.5 近期、日周期、周周期三通道融合
真实交通数据有明显的周期性,算法把输入切成了三个视角:预测时刻之前的一小时、昨天同一时刻的一小时、上周同一天同一时刻的一小时。三个视角共享同一套注意力加图卷积加时间卷积的网络结构,但各自独立训练权重,最后把三个分支的输出拼接或加权求和。
这个设计非常符合交通领域的先验知识——工作日早高峰和周末早高峰的拥堵模式不同,单纯用最近一小时数据无法刻画这种周期性差异。在实现时,日周期分支和周周期分支的时间步长度往往比近期分支更长,三个分支输入的时间维度不一样,但输出层必须对齐到相同的时间步数,才能拼接后送入输出层。
4. 把 ASTGCN 在本地跑起来:Python 实现与训练配置
4.1 用 PyTorch 搭建 ASTGCN 核心模型
完整的 ASTGCN 模型由三个结构相同的分支组成,每个分支内部包含两层时空注意力模块,最后接一个输出层把通道数映射到目标特征维度。下面给出一份精简版实现,包含空间注意力模块、图卷积模块、时间卷积模块和一个分支的组装方式:
import torch import torch.nn as nn import torch.nn.functional as F class ASTGCNBranch(nn.Module): def __init__(self, in_channels, num_nodes, num_steps, K=3, hidden=64): super().__init__() self.spatial_attention = SpatialAttention(in_channels, num_nodes, num_steps) self.temporal_conv1 = TemporalConv(in_channels, hidden, kernel_size=3) self.gcn_layer = GraphConv(hidden, hidden, K) self.temporal_conv2 = TemporalConv(hidden, hidden, kernel_size=3) self.batch_norm = nn.BatchNorm2d(hidden) def forward(self, x, A): # x: (B, N, T, C) att = self.spatial_attention(x) A_att = A * att.squeeze(1) x = self.temporal_conv1(x.permute(0, 3, 1, 2)) # -> (B, hidden, N, T) x = self.gcn_layer(x, A_att) x = self.temporal_conv2(x) x = self.batch_norm(x) return x.permute(0, 2, 3, 1) # -> (B, N, T, hidden)GraphConv 的内部实现直接调用上一章定义的切比雪夫多项式列表,将每一个多项式阶的图卷积结果做线性加权求和。代码中 A_att 是经过注意力加权的邻接矩阵,这正好体现了注意力机制和图卷积协作的方式:先让注意力模块计算动态权重,再在图卷积中按这个权重聚合邻居信息。
4.2 损失函数与评估指标的选择
交通流预测最常用的损失函数是均方误差(MSE),评估指标则看两个维度:平均绝对误差(MAE)衡量整体偏差,平均绝对百分比误差(MAPE)衡量相对偏差。流量场景下 MAPE 会遇到零值问题——凌晨某些检测器车流量为 0,直接算百分比会产生无穷大,所以工程上会给分母加一个很小的常数。
def masked_mape(pred, true, eps=1e-5): diff = torch.abs(pred - true) return torch.mean(diff / (torch.abs(true) + eps))这个掩码策略很重要,训练初期如果 MAPE 直接输出几十甚至几千,先检查分母是否出现了接近零的标签值,而不是急着调模型结构。测试集上的最终报告通常同时给出三组指标(MAE、RMSE、MAPE),和其他已发表论文对比时也要确保指标口径一致,有些论文报的是单步预测误差,有些报的是 12 步平均误差,混在一起比没有意义。
4.3 训练循环里的关键参数配置
训练 ASTGCN 我习惯用 Adam 优化器配合余弦退火学习率。PEMS04 和 PEMS08 的数据量不同,PEMS04 样本数更多,batch size 可以给到 64,PEMS08 建议降到 32,避免显存溢出。初始学习率设置在 0.001 到 0.002 之间,权重衰减取 0.0001,dropout 在 0.3 左右。
from torch.optim import Adam from torch.optim.lr_scheduler import CosineAnnealingLR model = ASTGCN(..., num_nodes=307, in_channels=3, hidden=64) optimizer = Adam(model.parameters(), lr=0.001, weight_decay=1e-4) scheduler = CosineAnnealingLR(optimizer, T_max=200, eta_min=1e-5) criterion = nn.MSELoss() for epoch in range(200): model.train() for x_batch, y_batch in train_loader: optimizer.zero_grad() output = model(x_batch, adj_matrix) loss = criterion(output, y_batch) loss.backward() nn.utils.clip_grad_norm_(model.parameters(), 5.0) optimizer.step() scheduler.step()梯度裁剪是个容易被忽略但关键时刻能救命的操作,交通数据偶尔会出现极端值,一个异常的梯度更新就可能把训练好的权重毁掉,裁剪到 5.0 能有效防止这种单点崩溃。训练过程中每个 epoch 结束后在验证集上计算 MAE,记录最优模型权重,后续测试全部基于这个最优权重来评估,不要用最后一个 epoch 的结果。
5. ASTGCN 训练的 5 个常见坑与排查方法
5.1 现象:loss 直接变成 NaN,训练没法继续
这个坑几乎每个从零复现的人都会遇到。原因有三类:第一是邻接矩阵中有孤立节点或全零行,归一化拉普拉斯矩阵时分母为 0;第二是数据里存在 NaN 值,前向传播之后梯度自然变成 NaN;第三是学习率过大导致梯度爆炸。
解决方法是先检查原始数据,用 np.isnan(data).any() 确认数据完整性;再在归一化拉普拉斯时分母加 1e-5 的 epsilon;最后给优化器加上梯度裁剪。排查顺序一定是从数据到模型到优化器,不要一上来就调学习率,否则会掩盖数据问题。
5.2 现象:训练 loss 持续下降,验证集 MAE 却停滞不动
这是典型的过拟合。ASTGCN 由三个分支组成,参数量并不小,PEMS 数据集时间维度只有几十天,模型很容易把训练集的时序模式背下来。一个反直觉的排查点:先看验证集曲线的下降趋势是否和训练集完全同步,如果训练 loss 一路向下、验证 loss 在某个 epoch 后开始回升,说明模型开始记忆训练集中特有的噪声模式。
解决方法是优先增强 dropout 的概率,把时间卷积层和图卷积层之间的 dropout 从 0.2 提高到 0.4;其次是增加权重衰减系数到 0.001;最后才考虑缩小 hidden 通道数。ASTGCN 这类模型容量很大,对于 PEMS08 这种 170 个节点的数据集,hidden 取 32 往往比取 64 更稳。
5.3 现象:显存不够,batch size 一调大就 OOM
图卷积层需要保存中间的多项式矩阵和注意力矩阵,307 个节点的邻接矩阵并不大,真正的显存瓶颈在注意力机制里的矩阵乘法。当 batch size 为 64、时间步为 12 时,空间注意力需要计算 (batch, 307, 307, 12) 的中间张量,多分支模型会把显存推到边缘。
解决方法是把 batch size 降到 32 或 16,同时确认验证和测试阶段不需要保存梯度,用 torch.no_grad() 包住前向传播。另一个有效手段是关闭图卷积层中的中间变量保留,不需要做梯度回传的层显式 detach,但注意这只适用于训练完毕后的测试阶段。在训练阶段,把切比雪夫多项式 K 从 3 降到 2 也能显著降低显存占用,代价是精度略微下降。
5.4 现象:PEMS04 上复现正常,换 PEMS08 后精度大幅下降
PEMS04 有 307 个节点,PEMS08 只有 170 个节点,两者数据分布差异很大。常见做法是把模型在 PEMS04 上训练好的权重直接迁移到 PEMS08 做微调,但这忽略了节点数量的变化——邻接矩阵的维度都变了,模型权重根本无法直接复制。
正确做法是把 PEMS08 当作独立任务重新初始化模型,只复用 PEMS04 训练过程中的超参数经验。如果是数据量不足,可以尝试用 PEMS04 预训练模型的特征提取层来初始化 PEMS08 模型的对应层,但需要把图卷积的权重矩阵维度重新映射,这一步很容易出错。我的经验是:除非刻意做迁移学习实验,否则老老实实分别训练,PEMS08 的数据虽然少,但独立训练 150 个 epoch 已经足够收敛。
5.5 现象:训练时验证集指标很好,测试集上效果暴跌
这是数据泄漏的典型症状。很多人习惯先对整个数据集做 min-max 归一化,再做训练测试切分,测试集的分布信息就这样泄露到了训练过程中。虽然交通流数据在时间上有连续性,泄露的影响不像图像分类那么致命,但会让模型对测试集的总体均值产生依赖,导致真实场景下误差偏大。
正确的顺序是先按 6:2:2 切分时间序列,再对训练集单独计算归一化参数,用训练集的 min 和 max 去变换验证集和测试集。我强烈建议把归一化逻辑写进一个自定义的 Dataset 类里,训练集持有一个 scaler,验证集和测试集各自持有同一个 scaler 的引用,从流程上杜绝误用。
6. 用最小实验验证模型真的学到了:从基线到结论
复现 ASTGCN 之后的下一步不是急着调参,而是先验证自己的实验闭环是否可信。我的习惯是先实现一个历史平均(HA)基线:用预测时刻之前同一时间段的历史均值作为预测值,这个基线跑在测试集上能给出一个 MAE 下限参考值。如果 ASTGCN 在验证集上的 MAE 不显著优于 HA,说明模型大概率没有学到有效模式,再怎么调注意力模块都白搭。
基线对比实验用固定随机种子跑三次取平均,每个 seed 对应一个完整的训练循环。深度学习训练有随机性,单次运行的结果可能波动很大,特别是 PEMS08 这种数据量较小的场景,某一次运气好的结果不能代表模型真实水平。我通常用 seed 等于 0、1、2 跑三组,报告均值和标准差。
实际验证过程中还有一个容易忽略的维度检查技巧:模型输出 shape 必须是 (batch, N, out_steps, D),很多人建模时把通道维和时间维顺序写反,训练 loss 一直降不下来,检查才发现维度错位。在第一次训练之前加一条断言语句可以省下一整天的时间:assert output.shape == y_batch.shape。
最后分享一个我的调参习惯:所有的超参数变更都必须绑定到验证集 MAE 的变化上,没有验证指标支撑的修改一律不做。ASTGCN 的调参空间很大,但真正对结果有决定性的因子只有三个——输入窗口长度、图卷积阶数 K、dropout 比例。窗口长度决定模型的短期记忆上限,K 决定空间聚合范围,dropout 直接控制拟合程度。把其他参数固定住,一次只动一个,记录每组参数下的验证曲线。愿这些经验能帮你少走弯路,希望帮到你。
本文还有配套的精品资源,点击获取