资讯动态

STGCN交通流预测复现:图卷积与时间卷积原理及PyTorch调参实践

发布时间:2026/10/2 2:40:15 来源:尧图企业网站定制
简介STGCN_IJCAI-18-master压缩包提供IJCAI 2018发表的时空图卷积网络论文实现代码是针对城市交通流量预测的完整深度学习解决方案适合智能交通、图神经网络与时间序列分析方向的研究者、学生或工程师参考学习。包内共19个文件主要包括11个Python脚本模型构建、训练、测试及数据加载与预处理工具、6张结果图预测值与实际值对比、训练损失曲线及路网结构示意、1份README说明文档和1个PeMS交通数据集压缩包压缩后大小约7.05MB文件组织规范便于按需取用。目前已有2307人学习下载。通过研读源码与图表可以深入理解STGCN如何利用图卷积捕获不同监测点间的空间依赖性并结合时间卷积建模流量演化规律同时掌握从原始数据清洗、模型训练到性能评估、可视化呈现的完整科研工作流对开展交通流预测实验或扩展图神经网络应用都有实际价值。1. 交通流预测的长期痛点STGCN 为什么在 IJCAI-18 之后成为复现热门做交通流预测的人多半经历过这样的夜晚用 LSTM 堆了十几层验证集 loss 还是像心电图一样上下抖动把历史一段时间的流量拼成特征喂给全连接网络预测结果永远比真实值滞后一个周期。STGCN 在 IJCAI-18 上改变这个局面——它把路网建模成一张真正的图让每个路段的预测能同时感知上下游的空间依赖而不是像图像那样把路网硬拉成二维矩阵。这篇笔记不重复论文实验而是把 STGCN_IJCAI-18 这个 Python 实现拆开讲明白它解决什么问题、在本地怎么跑通、参数怎么调、遇到报错怎么定位。适合刚入门的算法工程师也适合想用图神经网络做时序预测但还没想清楚边界的人。2. 先把原理吃透STGCN 里的图卷积和时间卷积分别解决什么问题2.1 交通路网为什么不能当作普通图像或普通序列交通流预测的输入不是图像也不是一组独立的序列。想象一座城市里的几百个卡口或地磁传感器它们按地理位置铺在路网上每个传感器只和相邻几个路段有强相关。上游堵车下游流量 20 分钟后才开始涨两条平行道路之间是竞争关系一条堵了另一条会分流。这种关系没法用普通二维卷积表达——普通卷积默认相邻像素是等间距、同性质的但交通路网中两个传感器相隔 5 公里可能是邻居相隔 200 米的两个传感器反而被河流隔开。用 LSTM 逐路口建模更可惜每个传感器的时间序列被单独处理路口之间的通信只能靠后端的全连接层去隐式学习数据量不够时根本学不到。STGCN 的做法是把每一个传感器当作图上的一个节点节点之间有边边的权重代表道路连通性、距离或历史流量相关性。这样空间依赖被显式编码进邻接矩阵模型每一层都在做“消息传递”一个节点的特征更新时会聚合它一跳或多跳邻居的特征。工程上的好处是只要把传感器编号和邻接矩阵准备好模型结构可以做到与路网规模解耦——换一个城市不需要改模型只需要换邻接矩阵。邻接矩阵怎么构造常见做法是用传感器距离的高斯核两个传感器之间距离越近边的权值越大。原论文里通常用 A_ij exp(-d_ij^2 / sigma^2)其中 d_ij 是两个传感器之间的距离sigma 是距离的带宽。这个公式要自己在数据预处理里算STGCN_IJCAI-18 的源码里通常有一个脚本会根据传感器经纬度生成邻接矩阵。这里有个细节邻接矩阵要不要归一化。STGCN 用的是拉普拉斯矩阵归一化也就是 D^{-1/2} A D^{-1/2}这一步直接决定图卷积能不能稳定收敛后面会专门讲。2.2 空间维谱域图卷积与切比雪夫近似的取舍空间图卷积听起来玄学但实现起来其实是一条固定路线。图卷积最早从谱域定义把图信号做图傅里叶变换在频域做滤波再变回来。计算代价很高因为要对拉普拉斯矩阵做特征分解。STGCN 弃用了完整谱域卷积改用切比雪夫多项式近似K 阶切比雪夫展开只需要计算 K 次拉普拉斯矩阵乘特征不需要特征分解。对每一个卷积层来说实际做的是这么一件事x_out sum_{k0}^{K-1} T_k(L_hat) * x * theta_k其中 T_k 是切比雪夫多项式L_hat 是归一化拉普拉斯矩阵theta_k 是对应第 k 阶多项式的可学习参数。这一段如果不好理解可以换个直观说法K 阶切比雪夫近似相当于把感受野限制在 K 跳邻居内。K1 时只聚合直接邻居K3 时能聚合到三跳以外的信息。交通数据里一个路段往往要参考上游三个路段以外的情况所以 STGCN 默认把 K 设成 3这在论文和大部分复现里都是一致的。为什么不用 GCN 常用的二阶近似GCN 的一阶近似写法简单但表达能力有限做节点分类够用做回归任务会显得“钝”。交通流预测是连续值回归需要更细腻的空间聚合所以 STGCN 保留了切比雪夫的 K 阶展开。实际使用中 K 超过 5 之后节点间的信息会过度平滑每个节点的特征都趋向于邻居的平均预测反而变差。这一点在调参时要特别注意。切比雪夫近似还带来一个工程优点计算只需要稀疏矩阵乘法。PyTorch 里可以用 torch.sparse 存储邻接矩阵大大减少显存占用。但稀疏矩阵乘法在 CPU 上并不总是比稠密矩阵快尤其当节点数只有几百个的时候这一点在避坑章节里重点说。2.3 时间维门控一维卷积为什么比 LSTM 更适合这个任务空间维度用图卷积之后时间维度上不能再用“一条序列一个模型”的思路。STGCN 在每个节点上做一维卷积直接用 CNN 捕捉时间依赖。看到这里有人会问不是都说时序预测用 LSTM 吗为什么这里用卷积两个原因。第一是并行性。一维卷积可以一次处理整个时间窗口LSTM 只能沿着时间步逐步计算。STGCN 的输入维度是 [B, N, C, T]其中 B 是 batchN 是节点数C 是通道数T 是时间步。时间卷积维度的卷积核沿着 T 方向滑动整个图在所有节点上共享同一套卷积核参数少、训练快。第二是梯度稳定性。LSTM 的长序列梯度问题虽然被门控缓解但训练时仍然容易振荡时间卷积的感受野是有限的反向传播路径短梯度更容易保持稳定。STGCN 在时间卷积之后加了一个门控线性单元输出变成经过 sigmoid 的门控值与原始卷积结果的逐元素乘积。这个门控让模型能决定时间窗口内哪些时刻的信息更重要比如早高峰的突变要保留深夜的平坦段可以忽略。原论文里时间维用的是标准一维卷积加 GLU有些复现版本改成因果卷积保证当前时刻只依赖过去时刻不依赖未来。如果你用的是 STGCN_IJCAI-18 这个项目需要确认它有没有做 padding 对齐如果做预测时泄露了未来信息测试指标会好看但落地失效。后文会在排查部分讲怎么检查这一点。3. 把 STGCN 跑通Python 环境准备与最小复现命令3.1 从 Python 安装到依赖清单我踩过的版本坑一个 STGCN 项目环境问题往往比模型问题更浪费时间。我的建议是不要用系统自带 Python直接按 Python 安装教程装一个干净版本然后创建独立虚拟环境。常见做法是用 conda它能同时管理 python 版本和 CUDA 相关依赖。如果你习惯 VSCode记得在 VSCode 的 python 环境配置里选择 conda 环境否则终端能跑训练但编辑器里 import 报错。下面这组命令我几乎在每个时序项目里都用conda create -n stgcn python3.8 conda activate stgcn python -m pip install --upgrade pip pip install torch1.8.0 torchvision0.9.0 --index-url https://download.pytorch.org/whl/cu111 pip install numpy pandas scipy scikit-learn matplotlib tqdm解释Python 3.8 是 PyG 和旧版 PyTorch 兼容最好的版本没必要追新torch 1.8.0 对应 CUDA 11.1如果你机器上 CUDA 版本更高可以去掉后面的 index-url直接装最新稳定版。numpy 不要装 2.xSTGCN 的很多源码是几年前的风格numpy 2.x 会把 np.float 等别名删掉import 直接崩。scipy 是计算拉普拉斯矩阵和稀疏矩阵需要tqdm 用来显示训练进度。这里有一个血泪经验如果你在 Windows 上装 PyTorch不会遇到 Linux 下编译的问题但要注意 PATH 里的 CUDA 版本和 PyTorch 的预期版本不一致会导致 torch.cuda.is_available() 返回 False。建议先用 python -c import torch;print(torch.cuda.is_available()) 验证再跑训练脚本。如果你用的是 Linux 系统直接 python 命令有时指向 Python 2需要确认 python3 --version。别小看这一步STGCN 源码里大量使用 f-string 语法Python 3.6 以下直接语法报错。所以与其纠结具体环境配置不如一开始就用 conda 锁死 Python 3.8省掉后面所有莫名其妙的兼容问题。3.2 下载预训练配置并启动训练三步命令详解STGCN_IJCAI-18-master 这个项目在 GitHub 上能直接搜到下载方式就是 git clone或者直接下载 zip 解压。我的建议是 clone方便后续拉更新。注意项目名里的 master 是分支名clone 下来后直接进目录git clone 仓库地址 STGCN_IJCAI-18-master cd STGCN_IJCAI-18-master python -m pip install -r requirements.txt用你搜索到的仓库地址替换 仓库地址。如果 clone 速度慢下载 zip 然后解压也一样目录名保持 STGCN_IJCAI-18-master 即可。然后看一下目录结构一般会有 data 或 datasets 文件夹MetrLA 和 PEMS-BAY 的数据集需要自己下载很多仓库只提供处理脚本不提供原始数据。确认数据放好后启动训练python train.py --model_name STGCN --dataset METR-LA --max_epochs 200 --batch_size 64 --device cuda:0参数说明--model_name 指定使用 STGCN 模型--dataset 指定 METR-LA 数据集--max_epochs 控制最大训练轮数。如果你只有 CPU把 --device 改成 cpu同时把 batch_size 调小到 16否则一个 epoch 可能跑半小时。第一次跑通建议把 max_epochs 直接设成 5先确认流程没问题再放开训练。一些仓库把参数放在 config 文件里而不是命令行如果你看到的是 parser.add_argument说明它支持命令行覆盖。如果训练脚本里硬编码了数据路径那就需要直接改源码里的 dataset 路径这种项目结构比较老但逻辑简单改一个路径而已。跑训练之前还有一个隐藏步骤检查项目的掩码矩阵。STGCN 的训练只对有效时间步算损失所以代码里会有一个 mask 标记哪些时间步可用。如果 mask 生成逻辑依赖日期解析而你的数据里时间列格式不一样mask 会全为 0训练时 loss 一动不动。这个坑非常隐蔽后面排查章节会再提到。3.3 第一次训练的预期输出怎么判断代码已经活了训练跑起来以后不要刷手机盯着前几个 batch 的输出。正常情况应该看到 loss通常是 MAE从较大的值快速下降比如从 30 降到 10然后在几个 epoch 内缓慢下降到个位数。下面是典型的训练日志epoch: 1, batch: 20, loss(mae): 25.36, lr: 0.001 epoch: 1, batch: 40, loss(mae): 18.77, lr: 0.001 epoch: 2, batch: 20, loss(mae): 9.42, lr: 0.001 epoch: 3, batch: 20, loss(mae): 6.21, lr: 0.001看到 loss 在降说明模型在学接下来等着每个 epoch 结束后的验证指标。如果 loss 一直卡在 30 不动或前几个 batch 就出现 NaN不要急着调网络结构大概率是数据预处理或学习率的问题详见第 5 章。第一次跑通后我强烈建议你做一件事用固定的随机种子重新跑一次。很多交通流预测的结果复现不出来就是因为没有固定 seed。在训练脚本开头加上import random import numpy as np import torch def set_seed(seed): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed(seed) torch.backends.cudnn.deterministic True然后调用 set_seed(0)。这样后面对比参数才公平否则你换一个超参数结果的差异可能只是随机波动。日志里还有一个关键字段是 lr。很多实现会用学习率衰减你需要确认 lr 是否按照预定的 schedule 在下降。如果 lr 始终不变后期 loss 可能会在一个平台期停住怎么训练都下不去。看到 lr 在正常衰减说明训练循环的调度器已经生效代码是真的活了。4. 看懂模型参数把论文里的数字翻译成可调的超参数4.1 STGCN 的网络结构堆叠与通道数设计STGCN 整体是一个“时空卷积块 输出层”的结构论文里常见的是两个时空卷积块堆叠再接全连接输出预测长度的值。每个时空卷积块内部先做时间卷积再做空间卷积然后残差连接。具体到代码一个块的核心逻辑大概是class STConvBlock(nn.Module): def __init__(self, in_channels, out_channels, K, T): super().__init__() self.temporal_conv1 TemporalConv(in_channels, out_channels, kernel_size(K, 1)) self.spatial_conv SpatialConv(out_channels, out_channels, K) self.temporal_conv2 TemporalConv(out_channels, out_channels, kernel_size(K, 1)) def forward(self, x): out self.temporal_conv1(x) out self.spatial_conv(out) out self.temporal_conv2(out) return x out这段代码是示意不是项目原稿但逻辑和大多数复现一致。第一个时间卷积先提炼时序特征中间的空间卷积聚合邻居信息第二个时间卷积再做一次时序混合。残差连接保证了深层的梯度不消失。通道数设计上一般把输入流量数据提升到 64 或 128 通道再逐层缩到预测需要的维度。通道数翻倍会显著提高模型容量但交通流信号本身没有图像那么丰富64 到 128 之间通常就够。超过 256 之后训练时间翻倍指标提升非常有限甚至过拟合。把关键参数整理成一张表调参时先对着它查参数常见值含义影响num_nodes207 / 325传感器个数决定邻接矩阵尺寸in_channels1每个节点每个时刻的特征维度一般固定为1hidden_channels64时空卷积隐藏通道越大拟合能力越强K3切比雪夫多项式的阶数空间感受野T12输入时间步数决定看多长的历史H12预测时间步数预测未来多久4.2 6 个必调的时空参数及其影响第一个是输入时间步 T。原论文常用 12对应过去一小时5 分钟一个采样点。T 太短模型看不全早晚高峰的完整过程T 太长会引入和当前预测不相关的历史信息训练成本也高。第二个是预测步长 HH1 是单步预测H12 是多步预测。多步预测有两种做法直接输出全部 H 个时间步或者自回归式逐步预测。直接输出更稳定自回归会累积误差。第三个是 K 阶数K2 和 K3 的差距在 METR-LA 上大约是 0.3 个 MAE看似不大但折线图里能看出来K3 对突发拥堵捕捉更准。第四个是 batch size如果节点数 207batch size 64 占用大约 5GB 显存32 是更稳妥的起点。第五个是学习率初始学习率 0.001 配合 Adam 是安全默认千万不要一上来就试 0.01。第六个是 dropout时空卷积里常用 0.1 到 0.3 之间节点数少的时候 dropout 太低容易过拟合。这六个参数里T 和 K 是 STGCN 区别于普通时序模型的核心。调 T 的时候注意输入时间窗口覆盖的周期要完整。如果数据是 5 分钟采样T12 只覆盖一小时对于超过一小时的拥堵传播不够用想覆盖早高峰 3 小时T 至少要 36。但这会让模型输入的通道数翻三倍训练时间也翻倍。折中方案是先做实验看 12、24、36 三组的验证曲线再选择拐点。学习率的调度也不可忽视。STGCN_IJCAI-18 的常见实现会在训练中降低学习率用 StepLR 每 20 个 epoch 乘以 0.7。很多复现者忘记这一步后期 loss 会一直震荡。如果看到训练集 loss 还在下降但验证集 loss 开始反弹先不要急着加正则先看是不是学习率太大导致越过最优点。把学习率降到原来的五分之一往往立刻见效。4.3 数据集划分与评价指标MAE/MAPE 的算法级解释交通流预测的标准做法是按时间顺序划分不能随机打乱。用前 70% 的数据训练后 30% 测试。如果打乱了训练集里混入未来信息测试指标会虚高但这种模型部署后马上露馅。评价指标里MAE 是绝对误差的平均值MAPE 是相对误差的平均百分比。MAE 受流量量纲影响METR-LA 上 MAE3 左右已经算不错MAPE 对凌晨接近零的流量非常敏感流量为 0 时误差会被放大成无穷。常见处理办法是在计算 MAPE 时给真实值加一个平滑项def masked_mape(y_true, y_pred): epsilon 1e-5 return np.mean(np.abs((y_true - y_pred) / (y_true epsilon)))这样分母不为零。另一个要注意的是多步预测的评价指标要逐时间步分别计算看第 12 步的误差是不是明显高于第 1 步——如果第 12 步 MAE 翻了三倍说明模型并没有学到长期依赖只是学到了“复制最近时刻”。数据划分还有一个容易忽略的细节标准化参数必须在训练集上计算再应用到验证集和测试集。如果你对整个数据集先标准化再划分测试集的均值方差会泄露进训练过程导致结果虚高。这个错误在交通流预测项目里出现频率很高很多人换了数据集后指标对不上其实就是这里出了问题。5. 复现 STGCN 的常见问题排查5 个让人想放弃的坑5.1 问题一数据加载后维度对不上reshape 报错现象运行 train.py 后立即报错类似 “shape [64, 12, 207, 1] is invalid for input of size ...”或者 reshape 维度算不对。原因不同仓库对输入数据的维度定义不一致。有的按 [B, T, N, F] 存储有的按 [B, N, T, F]。STGCN 原实现一般期望 [B, N, T, F]如果你用 MetrLA.h5 原始接口读出来的是 [T, N, F]需要转置。解决在数据加载后打印 x.shape明确每一维的含义再统一成 [B, N, T, F]B 是批量N 是节点数T 是时间步F 是特征数。可以用一句代码转置x x.transpose(0, 1) # 从 [B,T,N,F] 变为 [B,N,T,F] 的例子转置前想清楚哪维是节点哪维是时间。最笨但有效的办法是每次 reshape 后都加一行 assert x.shape (batch, num_nodes, time_steps, features)让维度错误尽早暴露。5.2 问题二loss 一直不降甚至 NaN现象训练刚开始 loss 正常几十个 batch 后突然变成 NaN或者从第一个 epoch 开始 loss 就卡在初始值附近不动。原因NaN 的来源一般是学习率过大导致梯度爆炸或者数据里有缺失值 NaN。loss 不降则可能是归一化范围不对流量数据没有压缩到合理区间或者邻接矩阵没有归一化图卷积的特征值范围过大导致梯度异常。解决先把学习率调到 0.0005 重跑再把输入数据做 z-score 归一化即减均值除标准差。检查邻接矩阵是否做了 D^{-1/2} A D^{-1/2}。用 torch.isnan(x).any() 在 forward 之后检查。最快的定位方法是把 batch size 设为 1如果一个样本就能跑通问题大概率在 batch 拼凑或 padding 逻辑上。5.3 问题三换数据集后结果骤降问题不在模型而是归一化现象在 METR-LA 上 MAE 正常换到 PEMS-BAY 后 MAE 上升 50%怎么调参都回不去。原因两个数据集的量纲不一样。METR-LA 的车速范围是 0-65 mphPEMS-BAY 的车速范围是 0-70 mph但均值方差差异较大。如果你沿用 METR-LA 的归一化参数预测结果会整体偏移。这也是为什么换数据集时不能只换文件。解决为每个数据集单独计算训练集的均值和标准差并保存到文件里。测试集和验证集也用训练集的统计量归一化不能混入测试集的统计量。另外邻接矩阵的计算方式要统一用经纬度算距离的带宽参数 sigma对每个数据集要重新标定否则空间依赖被错误缩放。5.4 问题四GPU 显存占用过高BatchSize 一调就崩现象batch size 设 64epoch 跑到一半 OOMout of memory减小到 32 还是 OOM。原因STGCN 图卷积在计算切比雪夫多项式时如果邻接矩阵用的是稠密矩阵显存会随节点数二次增长。207 个节点时还能撑325 个节点时立刻爆发。解决把邻接矩阵转换为稀疏格式代码里用 torch.sparse_coo_tensor乘法用 torch.sparse.mm。还有一个容易被忽略的检查你是否把全部历史时间步同时送入模型。有些实现会在每个 batch 里展开所有时刻的中间特征这会造成显存爆炸。可以等模型前向传播后释放中间变量或者在训练循环里用 with torch.no_grad() 处理验证集。真要临时续命就把 batch size 降到 8并把 PyTorch 的 gradient checkpointing 打开。5.5 问题五使用稀疏矩阵后训练变慢反而不如原始实现现象为了省显存把邻接矩阵改成稀疏格式训练速度慢了一倍每个 epoch 的时间从 2 分钟变 4 分钟。原因PyTorch 的稀疏矩阵乘法在 CPU 上实现并不成熟尤其是对很多零散元素的小图稠密矩阵乘法的底层 BLAS 优化反而更快。只有当节点数上千且邻接矩阵稀疏度很高时稀疏格式才会体现速度优势。解决不要盲目追求稀疏。207 个节点、大约 2% 的邻居密度时稠密邻接矩阵的显存占用也才 2072074 字节约 0.17MB完全吃得消。先跑通再谈优化。如果你必须用大规模路网考虑用消息传递框架比如 PyG 的 MessagePassing而不是手写稀疏矩阵。提示上面这五个坑是按出现频率排的维度问题和归一化问题占了大头。如果你严格按照第 3 章的环境和命令来至少能跳过前两个。6. 更高级的用法用早停法和可视化确认模型是否真学到交通规律6.1 多做一步用早停法保存最优模型训练 STGCN 时很多人只看最终的测试指标但如果 epoch 数太多模型会过拟合测试指标不升反降。我习惯在训练循环里根据验证集 MAE 保存最优权重if val_mae best_mae: best_mae val_mae torch.save(model.state_dict(), best_stgcn.pt)这样你在后面的实验里始终有一个最优版本不会被最终 epoch 的次优结果误导。6.2 画预测值和真实值的对比曲线一眼看出滞后问题训练结束后我建议不要只看指标把测试集第 1 个节点的预测曲线和真实曲线画在一起。用 matplotlib 就能做这本身就是 python 数据分析与可视化里最常用的一步import matplotlib.pyplot as plt plt.plot(y_true[:96], labelground truth) plt.plot(y_pred[:96], labelprediction) plt.legend() plt.show()如果两条曲线形状一致但在时间轴上整体向右平移说明模型实际上在复制最近时刻的流量变化而不是真正学到了交通传播规律。解决思路是适当增大输入时间步 T或者调整时间卷积的卷积核大小让模型有更多历史上下文。我自己曾经在 PEMS-BAY 上遇到过 MAE 很漂亮但曲线明显滞后的情况后来发现是测试时把所有时间步都做成了滑动窗口模型把最近一个时刻复制就直接赢了。从那以后我每个项目都至少留出半天做这个可视化检查因为指标会骗人曲线不会。希望帮到你。本文还有配套的精品资源点击获取

读完文章,也想定制专属网站?

尧图设计师 24 小时内与您沟通定制方案

免费获取报价 →
↑