免费获取学习方案
ARTICLE DETAIL

资讯详情

深耕编程基础知识与建站技术分享的一线实战洞察。

图神经网络交通流量预测实战:GCN/GAT/ChebNet代码解析与避坑指南

图神经网络交通流量预测实战:GCN/GAT/ChebNet代码解析与避坑指南 简介基于图卷积网络GCN的交通流量预测项目源码适合高校学生作为机器学习期末大作业或课程设计参考也适合希望快速上手图神经网络的初级开发者。项目以Python为主要语言围绕PeMS04高速公路交通流数据集完成建模、训练与预测代码附有详细注释并覆盖数据预处理、模型构建、评估与可视化等关键环节。压缩包共16个文件包含7个Python脚本、模型权重h5文件、PeMS04数据文件、可视化结果图与项目说明文档整体约33.61MB脚本涵盖GCN、GAT、ChebNet等常见图神经网络变体模块划分清晰。已有245人学习下载适合需要在较短时间内搭建完整预测流程的学习者。资源提供了从数据解析到结果分析的全链路实现并附带预训练权重和对比图表可以直接加载运行、查看效果还能为课程设计报告提供实验素材。1. GCN图神经网络交通流量预测不是调包是能跑通的骨架交通流量预测这件事最难的不是LSTM那套时间序列建模而是路网本身的空间结构——相邻路口的车流互相影响这种依赖关系用普通全连接网络根本学不出来。图神经网络GCN就是干这个的把路网抽象成一张图路口是节点车流是节点特征路段是边然后让信息沿着边传播。整个项目源码里的三种模型GCN、GAT、ChebNet全是这个思路数据用的是加州PeMS04高速路网真实流量数据不是捏造的。这份资源本身是一套完整的Python工程不是零散的算法片段。它含数据加载、模型定义、训练脚本、预测脚本、可视化脚本还附一份已经训好的GAT权重文件。期末大作业、课程设计、GIS相关专业的毕设拿它做底子改改就行。新手从零跑通一个图神经网络项目用它最省事熟手可以直接拿来做对比实验换层、换参数、换数据集都有入口。2. 三个图网络模型逐个拆文件结构、参数与选型边界2.1 项目文件结构与三类图卷积的代码入口打开压缩包先判断这是什么。目录里的核心文件有gcnnet.py、gat.py、chebnet.py三个模型定义文件分别对应三种图神经网络实现。traffic_dataset.py负责把PeMS04原始数据切成训练样本traffic_prediction.py是主入口脚本dataView.py画预测结果对比图。剩下还有utils.py放公共函数、PeMS04.npz和PeMS04.csv是原始数据与预处理后的特征文件。项目根目录/ ├── utils.py # 公共工具函数主要是数据归一化与指标计算 ├── gcnnet.py # GCN模型定义 ├── gat.py # GAT模型定义 ├── chebnet.py # ChebNet模型定义 ├── traffic_dataset.py # 数据集加载与滑动窗口切分 ├── traffic_prediction.py # 主训练/预测脚本 ├── dataView.py # 预测结果可视化脚本 ├── GAT_result.h5 # 已训练好的GAT权重 ├── PeMS04.npz # 预处理后的数据矩阵 ├── PeMS04.csv # 原始流量记录 └── README.md # 说明文档文件划分很清楚模型定义和数据管线解耦。新手不要一上来就动traffic_prediction.py里的训练逻辑先跑通预测流程再回头看模型内部。PeMS04.npz是已经处理过的矩阵数据一般用numpy.load()加载即可比直接解析PeMS04.csv快得多这也是为什么项目把两份数据都留着。2.2 GCN与ChebNet的谱域实现图卷积的数学基底GCN的核心是把图卷积定义成谱域上的滤波操作用切比雪夫多项式的一阶近似简化计算这是 Kipf 那篇经典论文的路线。gcnnet.py里通常就是两层图卷积加激活函数import torch import torch.nn as nn import torch.nn.functional as F class GCNLayer(nn.Module): def __init__(self, in_dim, out_dim): super().__init__() self.linear nn.Linear(in_dim, out_dim) def forward(self, x, adj): # adj: 归一化后的邻接矩阵shape [N, N] # x: 节点特征shape [N, in_dim] support self.linear(x) # 先做线性变换 out torch.mm(adj, support) # 邻接矩阵聚合邻居信息 return out class GCNNet(nn.Module): def __init__(self, in_dim, hidden_dim, out_dim): super().__init__() self.layer1 GCNLayer(in_dim, hidden_dim) self.layer2 GCNLayer(hidden_dim, out_dim) def forward(self, x, adj): h F.relu(self.layer1(x, adj)) out self.layer2(h, adj) return out这段代码里最重要的就是torch.mm(adj, support)这步它是整个图卷积的核心。adj必须是归一化后的邻接矩阵常见做法是D^{-1/2} A D^{-1/2}这种对称归一化。如果直接用原始邻接矩阵节点度数大的路口特征会被放大模型训练会很不稳定。这个项目的数据预处理里这一步通常已经做掉了但如果换成自己的数据集这步必须自己补。ChebNet在chebnet.py里它用的是K阶切比雪夫多项式展开本质是GCN的推广版本。GCN只考虑一阶邻居ChebNet可以聚合K阶邻居的信息。参数K越大感受野越宽但计算量也在涨。实际用的时候K2或3效果比较稳往上加收益很小还容易过平滑——所有节点的特征趋于一致。2.3 GAT的空间注意力视角为什么不选GCNGCN的邻居聚合权重是固定的由邻接矩阵决定。GAT不一样它用注意力机制动态计算邻居权重每个邻居对中心节点的贡献不是预设好的而是靠网络自己学。这在大规模路网里很有用因为不同时段、不同路段之间的影响强度本身是变化的早晚高峰和深夜的关系强度完全不同。GAT_result.h5这个文件说明这个项目里最终跑得最好的是GAT模型。文件是h5格式就是存了一整个训练好的状态字典加载时要用torch.load()配合load_state_dict()恢复模型权重。加载前必须确认模型结构和训练时完全一致不然key对不上。3. 秒级跑通预训练模型GAT_result.h5加载与预测脚本走读3.1 环境准备与最小复现路径先别想训练的事第一步是把预测链路跑通亲眼看到模型吐出一串预测值再说。整个项目依赖不多PyTorch加numpy加pandas基本就够了。版本上PyTorch 1.8以上都行2.x也兼容。不用GPU跑预测也很快一条测试样本的前向传播几十毫秒就结束了。# 建议新建虚拟环境避免把系统Python搞乱 conda create -n gat python3.8 -y conda activate gat # 安装核心依赖 pip install torch numpy pandas matplotlib装完之后先看一眼traffic_prediction.py里有没有写死路径。很多项目默认当前工作目录就是项目根目录如果你在别的目录下启的Pythonload(PeMS04.npz)会直接报文件不存在。稳妥做法是在项目根目录下执行python traffic_prediction.py --mode predict或者打开脚本把数据路径改成绝对路径。这是我每次拿到新项目都先干的活——先把路径问题干掉再谈模型。3.2 主流程走读数据加载、模型恢复与预测输出打开traffic_prediction.py核心流程是三段式加载数据、恢复模型、生成预测。加载数据用的是PeMS04.npz这个文件里通常存的是预处理好的邻接矩阵和特征序列。模型加载走的是load_state_dict逻辑import numpy as np import torch from gat import GATNet # 1. 加载预处理数据 data np.load(PeMS04.npz) x_data data[x] # 输入特征序列 y_data data[y] # 真实流量标签 adj data[adj] # 归一化后的邻接矩阵 # 2. 恢复GAT模型权重 model GATNet(in_dim12, hidden_dim64, out_dim3) checkpoint torch.load(GAT_result.h5, map_locationcpu) model.load_state_dict(checkpoint[model_state_dict]) model.eval() # 3. 前向传播输出预测 with torch.no_grad(): pred model(torch.FloatTensor(x_data), torch.FloatTensor(adj)) pred_np pred.numpy()这里的in_dim12是输入时间步长也就是用过去12个时间点的流量预测未来out_dim3是预测的未来3个时间点。PeMS04的原始采集频率是5分钟一条记录12个时间步对应过去1小时。如果你拿到手的项目里这两个值不一样以PeMS04.npz里实际张量shape为准不要死守我写的这个值。3.3 输出指标与可视化图表解读训练好的模型预测完光看数字看不出好坏。dataView.py就是干这个的它把预测值和真实流量曲线画在一起左边是训练集效果右边是测试集效果。项目还专门保存了node_10_3.png和gat_node_120.png这类图表从文件名能看出是按节点维度画的——node_10_3大概率是第10个节点、第3个预测步的对比效果。看图的时候重点看两条曲线是否贴合尤其是峰值的相位有没有偏移。流量预测最容易出现的情况是预测曲线比真实曲线滞后一个步长看起来像整体向右平移。如果你复现时画出来的图也这样先别急着怀疑代码把时间步长缩小试试或者把预测步数从3改成1往往滞后感就没了。4. 数据管线吃透PeMS04.npz的加载、归一化与数据切分4.1 PeMS04数据集与npz内部结构PeMS04是加州高速公路传感器网络公开数据集每条记录是某个传感器站点每5分钟的交通流量均值。原始csv是按时间排的大宽表行是时间戳列是传感器ID。模型吃不了这种原始格式所以项目里预先处理成了PeMS04.npz。我拿到别人的npz文件第一件事永远是先看key和shape而不是直接开跑import numpy as np data np.load(PeMS04.npz) print(data.files) # 查看所有存储的键名 for key in data.files: arr data[key] print(f{key}: {arr.shape}, {arr.dtype})PeMS04.npz里一般会有几个固定的键x、y、adj可能还有归一化的均值方差。x的形状通常是[样本数, 节点数, 时间步长]y是[样本数, 节点数, 预测步长]。这个排列顺序直接决定了后面维度要不要做permute。4.2 滑动窗口切分逻辑与train/val/test划分traffic_dataset.py里最关键的一段就是怎么把连续流量序列切成一个个样本以及怎么防止数据泄漏。我见过不少新手在这里翻车——把相邻时间步的样本同时塞进训练集和验证集导致验证指标虚高最后上真实场景效果稀碎。def create_samples(data, input_steps12, pred_steps3): samples_x, samples_y [], [] for i in range(len(data) - input_steps - pred_steps 1): x data[i : i input_steps] y data[i input_steps : i input_steps pred_steps] samples_x.append(x) samples_y.append(y) return np.array(samples_x), np.array(samples_y)这段代码是最朴素也最常见的滑动窗口切法。input_steps是回看窗口长度pred_steps是预测跨度两个窗口中间没有重叠避免信息穿越。切好之后按比例划分训练集、验证集、测试集常见的做法是前70%训练、中间10%验证、最后20%测试严格按照时间顺序切不做随机打乱。交通数据是时间序列随机打乱等于把时间信息揉碎了模型的泛化能力会被严重高估。4.3 邻接矩阵的构建与归一化邻接矩阵描述的是路口之间的连接关系。PeMS04的原始数据里有每个传感器站点的经纬度通常的做法是计算站点两两之间的距离小于某个阈值的就视为相连或者直接用K近邻法每个节点连最近的K个节点。这个项目的npz里已经内置了构建好的邻接矩阵不需要你再算。但有一点必须确认矩阵是不是已经归一化过了。GCN/GAT对邻接矩阵的数值范围很敏感如果原始adj里全是0和1八成没归一化跑出来的训练曲线通常是乱跳的。判断方法很简单打印adj.sum(axis1)看每行的和是不是1左右如果是说明已经做了随机游走归一化如果忽大忽小就得自己补一步。补的时候我一般用对称归一化def normalize_adj(adj): row_sum adj.sum(axis1) d_inv_sqrt np.power(row_sum, -0.5) d_inv_sqrt[np.isinf(d_inv_sqrt)] 0.0 d_mat_inv_sqrt np.diag(d_inv_sqrt) adj_norm d_mat_inv_sqrt adj d_mat_inv_sqrt return adj_normnp.power(row_sum, -0.5)对0度节点会产生无穷大所以必须有np.isinf那一步兜底。现实中不会出现没有任何邻居的孤立传感器但数据预处理时少了一两个站点的连通关系这种脏数据还是可能遇到。5. 避坑清单训练翻车与指标不对的六个常见原因坑这东西踩过一次就记住了。以下六条全是我在实际跑图神经网络交通流量预测项目时遇到过的问题按“现象→原因→解决”写遇到类似情况直接对号入座。坑一加载模型时报key不匹配state_dict缺key或多key现象执行model.load_state_dict(checkpoint[model_state_dict])直接抛RuntimeError提示Missing key(s) in state_dict。原因训练时定义的模型结构跟现在加载时不一致。最常见的是hidden_dim或者注意力头数不同也有的是pytorch版本不同导致参数命名规则变化。解决先别急着改代码把checkpoint里的key打印出来跟当前模型的state_dict key做对比一个一个对应。确认差异之后以checkpoint的保存结构为准调整模型定义参数。我一般是写一段临时脚本把两侧key都打出来差哪个改哪个。坑二PeMS04.npz加载出来的维度跟自己预期不一致现象代码跑了几十行报维度错误x_data.shape写成了[N, 12, node]模型里实际期望[N, node, 12]。原因npz里的数据排列顺序和模型的输入约定不一致。有些项目特征维度排在最前面有些节点维度排在最前面。解决在数据加载之后、喂给模型之前统一加一行permute或transpose把维度转成模型期望的格式。不要嫌麻烦这行代码能省掉后面几小时的排错。坑三损失值从一开始就疯狂震荡或者直接变成NaN现象训练第一个epochloss像过山车一样反复横跳某个epoch直接变成NaN再也回不来。原因最常见的是学习率过大GCN/GAT这类图模型对学习率比普通MLP敏感得多。另一个原因是邻接矩阵没有归一化节点特征数值范围太大梯度爆炸。解决先把学习率降到1e-3甚至5e-4再用std归一化节点特征最后检查邻接矩阵的归一化。三件事按顺序排查大概率是其中一个环节没有做彻底。坑四验证集指标比训练集还好结果上线后预测一塌糊涂现象单看验证集loss和MAE都很漂亮但换到新路段或者第二天数据上预测曲线几乎是一条水平线。原因数据泄漏。滑动窗口切分时训练集和验证集的时间窗口重叠了——直接把原始序列按索引随机切没有按时间先后切。验证集里混着训练样本的“未来”所以模型相当于提前看到了答案。解决强制按时间顺序划分前70%训练、后30%测试。如果数据集里有日期标记按日期边界切而不是按索引随机切。这条是血泪经验血的教训做过一次以后每次都先检查分割边界。坑五预测曲线滞后峰值的相位总是偏移现象预测值和真实值趋势一致但永远慢半拍看起来预测曲线是被平移了。原因预测目标设置得太宽pred_steps设得太大。用过去1小时预测未来15分钟跟预测未来1小时难度完全不同。预测窗口越大模型越倾向于输出保守的平均值滞后感越明显。解决先跑pred_steps1验证代码能吃通再逐步加大。另外检查一下滑窗切分时y是不是错取成了x的错位副本——有时数据中心不全代码自动用前一个时刻填充也会造成滞后。坑六GAT训起来特别慢显存也不够用现象GCN跑起来很流畅换上GAT之后训练时间翻了三倍batch稍大直接OOM。原因GAT要在每条边上算注意力还要多头拼接复杂度比GCN高一个量级。GCN是稀疏邻接矩阵一次矩阵乘法GAT是多头注意力逐点计算。解决先把注意力头数从8降到4特征维度从64降到32。PeMS04这种中等规模路网4头、32维的配置已经能把效果跑到接近满配水平。另外训练时batch_size从32降到16一般就不炸显存了。6. 从零训练加可视化回放用一张对比图验证模型真实效果预训练模型跑通之后一定要自己动手训一遍。这一步把GAT换成GCNChebNet也换成自己的参数体会三种模型在同一数据集上的收敛差异。训练时的核心循环不复杂真正决定成败的是数据管线和模型定义之间每个维度是否对得上def train(model, optimizer, x, y, adj, epochs100): loss_fn nn.MSELoss() for epoch in range(epochs): model.train() optimizer.zero_grad() pred model(x, adj) # [batch, node, pred_steps] loss loss_fn(pred, y) # y shape 必须与 pred 完全一致 loss.backward() optimizer.step() if epoch % 10 0: print(fEpoch {epoch}, loss: {loss.item():.6f})训练完别急着结束把预测结果画出来。用dataView.py分别绘制GCN、GAT、ChebNet在同一个测试节点上的预测曲线三张图堆在一起比较。你会发现GAT在流量突变期更敏锐GCN在中低流量段更平稳ChebNet的高阶邻居信息在某些波动区间有明显优势。这些差异从数据表里看不出感觉画成图马上就有体感。画图时把数据先反归一化否则Y轴是0到1的小数看着指标漂亮但完全无法向别人解释实际的流量数值。找到utils.py里的反归一化函数传预测和真实值之前先处理一步。从那以后我每次换模型、换数据集都强制走一遍全流程先打印输入特征、邻接矩阵、标签三者的shape核对全部对上号再开训。训完必须把预测跟真实曲线画在同一张图里肉眼确认曲率对齐而不是单看loss数字好不好看。这样排查问题的速度快很多希望帮到你。本文还有配套的精品资源点击获取
返回列表