PyTorch复现STGCN:时空预测模型从数据处理到调参实战
2026/9/7 2:06:49 网站建设 项目流程

简介:时空预测模型 PyTorch 复现代码库,面向希望深入理解时空预测算法实现细节的研究生与开发者,也适合需要搭建模型训练框架的 PyTorch 使用者。资源共 65 个文件,压缩包仅 57KB,以 53 个 Python 源码为核心,涵盖模型定义、训练测试模板与数据处理工具,另有工程配置文件、说明文档及许可证文件。模型目录按结构分文件夹组织,严格参照论文公式和图示复现,统一约定输入为批大小、序列长度、通道数、高度、宽度的五维张量;每个模型拆分为可复用的小模块再组合,便于逐段对照论文理解前向传播逻辑。工具目录则提供大尺寸数据的分块处理方法、可继承重写的训练与测试模板类,以及自动生成目录树的脚本,可辅助自定义训练流程、整理实验代码。整体代码以教学和可读性为目标,未刻意追求运行效率,适合论文复现与算法学习。目前已有 579 人学习下载,值得时空预测方向的初学者和进阶者参考。 时空预测模型这两年是真的火,不管是交通流量预测、气象预报,还是人群密度分析、电网负荷预估,背后都离不开这套东西。我前阵子因为项目需要,把时空预测这块从头到尾用PyTorch老老实实复现了一遍,踩了不少坑,也积累了不少心得。这篇就把我完整的过程记录下来,从模型选型、数据处理、代码实现到训练调参,全都摊开来讲,希望能帮到正在折腾类似任务的朋友。

这篇东西适合什么读者?一种是刚接触时空预测、想找个成熟模型练手的研究生或工程师,另一种是已经跑通基础CNN/RNN、想往图神经网络方向深入的同学。我会以STGCN(时空图卷积网络)为主线来拆解,因为它结构清晰、训练成本适中、效果也够稳定,非常适合作为复现的样板模型。文章里所有代码都是PyTorch实现,完整可跑。

1. 时空预测到底在预测什么——模型选型前的思路拆解

1.1 时空问题的两个核心维度

先捋清楚问题本身。时空预测和普通时序预测最大的区别在于:数据不仅随时间变化,还同时在空间上相互影响。最简单的例子就是城市路网中的交通流量——某个路口的车流量暴涨,十几分钟后相邻路口的流量也会跟着变。如果只对每个路口单独做时间序列预测,等于把路口之间天然的依赖关系扔掉了,效果肯定打折扣。

所以时空预测要同时建模两个维度:

  • 时间维度:数据自身的周期性、趋势性、突发性。交通数据有早高峰晚高峰的规律,气象数据有昼夜温差的变化,这些靠时序模型(LSTM、TCN、Transformer)来捕捉。
  • 空间维度:不同节点之间的相互影响和依赖关系。路口A堵车会影响路口B,传感器C的数据异常可能源于传感器D的事件。这种关系需要用图结构来描述,靠图神经网络(GCN、GAT)来建模。

把这两者结合起来,就是“时空图神经网络”这条技术路线要解决的核心问题。

1.2 主流模型家族对比

目前做时空预测的主流模型大概分三类:

  • 卷积类:ConvLSTM把LSTM的状态转移换成卷积操作,适合规则网格数据(比如气象雷达图),但处理路网这种非欧几里得结构就很别扭。
  • 图卷积类:STGCN、Graph WaveNet这类模型把路网抽象成图,节点是传感器,边是道路连接关系,用图卷积做空间建模,用1D卷积或GRU做时间建模。这类模型最贴合交通、电网、通信网络等场景。
  • 注意力类:时空注意力Transformer(如STAR、GMAN)用self-attention捕捉长距离依赖,效果好但训练成本高,数据量小的时候容易过拟合。

我最终选了STGCN(Spatio-Temporal Graph Convolutional Network)作为复现对象,原因是它结构足够经典,很多后续模型都拿它当baseline,复现一遍等于打了一次基础。而且它不像Transformer那样吃数据量,中等规模的数据集就能训练出不错的效果。

1.3 复现前必须确认的三件事

动手写代码之前,先把几个问题想清楚,否则后面全是坑:

  1. 预测什么?是预测下一个时间步的值,还是未来12个步长的序列?这个决定模型输出层的设计。我这次做的是多步预测,一次输出未来12个时间步。
  2. 数据长什么样?每个样本需要包含时间窗口的历史数据(比如过去12个步长)和对应的未来值。输入张量的形状是(batch_size, 输入步长, 节点数, 特征维度)
  3. 用什么指标评价?交通预测领域最常用的是MAE、RMSE和MAPE。MAPE对真实值为0的情况很敏感,数据预处理时需要注意过滤或平滑。

2. 环境搭建与数据预处理:复现路上最容易翻车的两个环节

2.1 PyTorch环境配置要点

先说环境。复现时空预测模型对PyTorch版本没有特别苛刻的要求,我用的是PyTorch 2.0+,CUDA 11.8,实测稳定。如果你用的是新版显卡,直接装CUDA 12.1对应的PyTorch也没问题。关键点是:GPU驱动版本和PyTorch的CUDA版本要匹配,否则会出现"CUDA driver version is insufficient"这类报错。

建议用conda建独立环境,避免污染其他项目:

conda create -n stgcn python=3.9 conda activate stgcn pip install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu118 pip install numpy pandas scipy scikit-learn matplotlib tqdm

装好后用下面这段代码验证GPU是否可用:

import torch print(torch.__version__) print(torch.cuda.is_available()) print(torch.cuda.get_device_name(0))

如果最后一行输出了你的显卡型号,环境就没问题。

2.2 数据加载与时序窗口切分

数据我用的是公开的交通流量数据集PEMS04,包含307个传感器节点,每5分钟采集一次流量数据。当然手头没有这个数据集的也可以用自己的时序数据,关键是要有“多个节点同步记录”的结构。

数据处理的核心是构造训练样本。原始数据是一大张二维表(时间步长 × 节点数),需要切成一个个时间窗口:

def create_samples(data, input_len=12, output_len=12): """ data: (T, N) T为时间步数,N为节点数 input_len: 历史窗口长度 output_len: 预测窗口长度 """ samples_x, samples_y = [], [] T = data.shape[0] for i in range(T - input_len - output_len): x = data[i : i + input_len] # (input_len, N) y = data[i + input_len : i + input_len + output_len] # (output_len, N) samples_x.append(x) samples_y.append(y) return np.stack(samples_x), np.stack(samples_y)

然后按时间顺序切分成训练集、验证集、测试集。这里有个容易踩的坑:不能随机打乱!时间序列数据一旦随机shuffle,就相当于把未来信息泄漏到了训练集里,验证集和测试集就失去了意义。正确做法是按时间先后顺序切分,比如前70%训练、中间10%验证、最后20%测试。

2.3 邻接矩阵构建:图卷积的“图”从哪来

STGCN的空间建模依赖一个关键输入——邻接矩阵。它描述的是节点之间的连接关系。

常用的构建方法是基于距离的阈值高斯核:

def compute_adjacency_matrix(distances, sigma=0.1, threshold=0.5): """ distances: (N, N) 节点间距离矩阵 threshold: 距离超过该阈值视为不连通 """ N = distances.shape[0] adj = np.zeros((N, N)) for i in range(N): for j in range(N): if distances[i, j] <= threshold: adj[i, j] = np.exp(-(distances[i, j] ** 2) / (sigma ** 2)) return adj

这里高斯核的作用是:距离越近的节点,权重越大,代表空间上影响越强。阈值的作用是强制稀疏化——现实路网中两个相隔很远的传感器理论上没太大关系,不用建立连接,这也能减少计算量。

构建完邻接矩阵后,还需要做归一化处理。STGCN原作者用的是对称归一化:

def normalized_adj(adj): # D^{-1/2} * A * D^{-1/2} degree = np.sum(adj, axis=1) degree_inv_sqrt = np.power(degree, -0.5) degree_inv_sqrt[np.isinf(degree_inv_sqrt)] = 0.0 degree_inv_sqrt_mat = np.diag(degree_inv_sqrt) return np.dot(np.dot(degree_inv_sqrt_mat, adj), degree_inv_sqrt_mat)

这个归一化的意义在于:不同的节点度数(邻居数量)差异很大,如果不归一化,有的节点聚合的信息量会比别人大几个数量级,训练会非常不稳定。

3. 核心模型实现:STGCN逐层拆解

3.1 STGCN整体结构

STGCN的总体结构是“时空卷积块”堆叠加输出层。每个时空卷积块内部包含两个时间卷积模块和一个空间卷积模块,排列方式是:时间卷积 → 空间卷积 → 时间卷积,中间加残差连接。

这样的排列逻辑很直观:先沿时间维度提取局部时序特征,再做空间信息聚合,最后再融合一次时间维度。两个时间卷积夹一个空间卷积的“沙漏”设计,是为了让空间卷积能在更高层的时间特征上操作,信息提取更充分。

我用PyTorch搭出来的核心代码结构是这样的:

class STGCN(nn.Module): def __init__(self, in_channels, hidden_dim, out_channels, num_nodes, num_layers=2): super().__init__() self.blocks = nn.ModuleList() for i in range(num_layers): in_ch = in_channels if i == 0 else hidden_dim self.blocks.append(STConvBlock(in_ch, hidden_dim, num_nodes)) self.output_layer = nn.Conv2d(hidden_dim, out_channels, kernel_size=(1, 1)) def forward(self, x, adj): # x: (B, input_len, N, C_in) x = x.permute(0, 3, 1, 2) # (B, C_in, input_len, N) for block in self.blocks: x = block(x, adj) # 输出层 x = self.output_layer(x) # x: (B, out_channels, input_len, N) return x

注意这里我把输入张量重排成了(B, C, T, N)的格式,这是PyTorch卷积操作的标准布局。很多新手复现时维度没理清楚,写着写着就乱了,强烈建议在纸上把每个tensor的维度标出来再动手。

3.2 时间卷积模块:1D卷积加门控

时间卷积的作用是沿时间轴提取特征,同时压缩时间维度。STGCN用的是带门控机制的1D卷积:

class TemporalConv(nn.Module): def __init__(self, in_channels, out_channels, kernel_size=3): super().__init__() self.conv = nn.Conv2d( in_channels=in_channels, out_channels=2 * out_channels, # 一半用于门控 kernel_size=(kernel_size, 1), padding=(kernel_size // 2, 0) ) self.glu = nn.GLU(dim=1) # 沿通道维做门控 def forward(self, x): # x: (B, C_in, T, N) out = self.conv(x) # (B, 2*C_out, T, N) return self.glu(out) # (B, C_out, T, N)

GLU(门控线性单元)的公式是GLU(a, b) = a * sigmoid(b),其中a是“保留信息”,b是“控制门”。这种机制让模型可以学习哪些时间步的信息该保留、哪些该抑制,比纯粹堆卷积层灵活得多。

kernel_size我设的3,代表每个时间步的分类会参考前后各一个步长的信息,对于5分钟粒度的交通数据来说够用了。如果数据粒度更细,可以适当调大kernel。

3.3 空间卷积模块:图卷积的切比雪夫近似

空间卷积是STGCN的核心。它用的是一阶切比雪夫近似的图卷积,公式上就是:

Z = D^{-1/2} A' D^{-1/2} X W

其中A'是加了自环的邻接矩阵(自己和自己也算邻居),D是度矩阵,X是输入特征,W是可学习的权重矩阵。这个公式的本质是:每个节点把自己的特征和邻居节点的特征加权求和,再用可学习权重做线性变换。一层图卷积下来,每个节点就聚合了一步邻居的信息;堆多层,就能捕捉多跳关系。

PyTorch实现如下:

class SpatialConv(nn.Module): def __init__(self, in_channels, out_channels, num_nodes): super().__init__() self.num_nodes = num_nodes self.theta = nn.Parameter(torch.FloatTensor(in_channels, out_channels)) nn.init.xavier_uniform_(self.theta) def forward(self, x, adj): # x: (B, C_in, T, N) B, C_in, T, N = x.shape x = x.permute(0, 2, 3, 1) # (B, T, N, C_in) x = x.reshape(-1, N, C_in) # (B*T, N, C_in) # 图卷积核心:AXW ax = torch.matmul(adj, x) # (B*T, N, C_in) 空间聚合 out = torch.matmul(ax, self.theta) # 线性变换 out = out.reshape(B, T, N, -1) out = out.permute(0, 3, 1, 2) # (B, C_out, T, N) return out

这段代码看着简单,但有几个关键点:

  • 邻接矩阵adj在前向传播前就要做归一化,不能拿原始邻接矩阵直接用。
  • torch.matmul(adj, x)这一步是图卷积的灵魂,它让每个节点拿到邻居的加权信息。矩阵乘法的顺序别搞反了,adj在前、特征在后,相当于按行聚合。
  • 我把B*T合成了同一个维度,好处是节点数是固定的,可以用矩阵乘法批量处理,效率高。

3.4 时空卷积块的完整组装

把时间卷积和空间卷积串起来,加上残差连接和层归一化,就是完整的一个时空卷积块:

class STConvBlock(nn.Module): def __init__(self, in_channels, hidden_dim, num_nodes): super().__init__() self.temporal1 = TemporalConv(in_channels, hidden_dim) self.spatial = SpatialConv(hidden_dim, hidden_dim, num_nodes) self.temporal2 = TemporalConv(hidden_dim, hidden_dim) self.bn = nn.BatchNorm2d(hidden_dim) self.residual = nn.Conv2d(in_channels, hidden_dim, kernel_size=1) if in_channels != hidden_dim else None def forward(self, x, adj): out = self.temporal1(x) out = self.spatial(out, adj) out = self.temporal2(out) out = self.bn(out) if self.residual is not None: x = self.residual(x) return out + x # 残差连接

残差连接的作用是缓解深层网络的梯度消失问题。前一个块的输出直接加到后一个块的输出上,让梯度能有一条“高速公路”直接回传。我实测发现没有残差连接的STGCN在PEMS04上训练特别容易震荡,加了之后稳定很多。

4. 训练配置与调参实录

4.1 训练参数一览

我最终采用的训练配置如下,供你参考:

参数说明
batch_size64显存够用可以调大,加快训练
learning_rate0.001初始学习率,配合余弦退火衰减
epochs100配合早停,实测40~60轮左右收敛
optimizerAdambetas默认(0.9, 0.999)
lossMSE回归任务默认选择
input_len12用过去1小时数据(5分钟粒度)
output_len12预测未来1小时
hidden_dim64隐藏层维度,越大拟合能力越强
num_layers2时空块数量

4.2 为什么用MSE而不是MAE

时空预测的损失函数我首选MSE(均方误差),原因有两点:

  • MSE对误差大的样本惩罚更重,训练初期可以帮助模型快速学会“大趋势”。
  • MAE梯度恒定为±1,收敛后期容易在最优值附近来回震荡,MSE的梯度随误差减小而减小,收敛更平稳。

当然MSE也有缺点:对离群点非常敏感。如果你的数据里有明显的异常值,建议先做清洗或截断。我有个朋友直接拿原始数据跑,结果MAPE炸到30%以上,排了半天发现是几个极端值在捣乱。

训练循环的代码很常规,但有个小细节值得提——学习率调度器。我用的是余弦退火:

scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=epochs)

余弦退火的特点是前期学习率下降慢、后期下降快,既能快速到达最优区域,又能在后期精细搜索。比起固定学习率或StepLR,这种方式在我这个任务上明显更好收敛。

4.3 复现的“科学素养”:随机种子和模型保存

很多人复现论文效果不好,第一反应是模型代码有问题,其实很多情况下是没有固定随机种子

def set_seed(seed=42): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) torch.backends.cudnn.deterministic = True torch.backends.cudnn.benchmark = False

固定种子的意义在于让实验结果可重复。你调参时判断“这个改动是否有效”,前提是两次实验的随机性一致,否则你都不知道效果变好是因为改对了还是因为运气好。

模型保存也要注意,PyTorch里推荐只保存state_dict而不保存整个模型:

torch.save(model.state_dict(), 'stgcn_best.pth')

加载的时候需要先实例化模型再load,这样和PyTorch版本耦合度低,换环境也能加载。

4.4 训练过程中需要盯住的几个信号

训练时不要傻等100个epoch结束。我在训练过程中会同时监控训练损失和验证损失,特别关注两个信号:

  • 验证损失下降减缓:说明模型快收敛了,此时如果结合早停策略,可以省下大量时间。
  • 训练损失和验证损失的gap拉大:这是过拟合的前兆。STGCN的隐藏层维度如果设得太大、又没有足够的Dropout,在PEMS04这种中等规模数据上很容易过拟合。

我在代码里加了一个简单的早停机制:连续15个epoch验证损失不下降,就回退到历史最优模型并停止训练。实际跑下来大概在50轮左右触发早停,效果比硬train100轮好不少。

5. 常见问题与排查技巧实录

复现过程中我踩了不少坑,也帮朋友排查过类似问题,整理成速查表:

现象可能原因解决办法
训练时loss直接nan学习率过大 / 输入数据有NaN降低学习率;检查数据是否有缺失值未处理
维度不匹配报错tensor布局不对统一用(B, C, T, N)布局,每步打印shape验证
验证集loss远大于训练集过拟合减小hidden_dim / 加Dropout / 加大训练数据量
预测结果几乎全是一个常数模型未学到空间依赖检查邻接矩阵是否构建正确,尤其是归一化步骤
训练很慢,GPU利用率低batch_size太小 / DataLoader线程不足调大batch_size;设置num_workers>0
复现结果和论文差距大数据预处理细节不同确认归一化方式(z-score vs min-max)、切分比例

5.1 维度不匹配:最常见的新手杀手

我在踩过坑后发现,维度问题有个很实用的排查方法——在forward的每个关键步骤后打印tensor的shape,把维度写在注释里。比如:

x = x.permute(0, 3, 1, 2) # (64, 1, 12, 307) out = self.temporal1(x) # (64, 64, 12, 307) out = self.spatial(out, adj) # (64, 64, 12, 307)

这样一旦报错,扫一眼就知道是哪个环节的形状不对。这个方法听起来很原始,但真的比盯着堆栈找半天高效。

另外要注意,nn.Conv2d默认的输入布局是(B, C, H, W),如果你想对时间维做卷积,把时间维放在H的位置就行。很多人不习惯这种处理方式,总想用Conv1d,结果被迫reshape,反而更容易出bug。

5.2 梯度爆炸和loss不下降的处理

时空图卷积网络因为涉及矩阵乘法叠加,梯度爆炸的情况不算少见。我常用的三板斧:

  1. 梯度裁剪torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5),在backward之后、optimizer.step()之前执行。
  2. 降低学习率:从0.001直接调到0.0005,往往就能避免训练早期的不稳定。
  3. 检查输入数据:对数据做z-score标准化,把数值范围压到0附近,能显著提升训练稳定性。

如果loss长时间不下降,先检查是不是代码里有逻辑错误——比如在某个分支把梯度detach掉了,或者归一化层的参数被冻结了。这些都是我真实遇到过的坑。

5.3 复现结果对不上论文怎么办

这是最让人崩溃的情况。代码没问题,数据也一样,结果就是有差距。我的经验是排查顺序如下:

  • 数据预处理是否完全一致:归一化时机的微小差异,可能导致结果差2%~5%的MAPE。
  • 训练配置是否一致:batch_size、学习率衰减策略、早停设置都会影响最终结果。
  • 评估方式是否一致:有些论文用的是“反归一化后的预测值和真实值”计算指标,有些用归一化后的数据算,两种方式结果差异巨大。
  • 随机种子:不同种子跑出来的结果,方差可能超过1%的MAPE。多跑几个seed取平均,会比较接近论文结果。

6. 写在最后的一些体会

这次把STGCN用PyTorch完整复现一遍,最大的感受是:时空预测模型的门槛不在模型代码,而在数据处理和实验设计的严谨程度。STGCN本身的代码就一两百行,但它对输入数据的格式、邻接矩阵的构建、评估方式的要求都非常明确,任何一环出错,最终结果都会偏离预期。

我个人的建议是:复现任何模型之前,先花时间把数据管道搭好,尤其是邻接矩阵的计算和归一化,这部分占了整个复现工作将近一半的工作量。模型结构反倒是最能“抄作业”的部分——PyTorch的文档和开源实现都够多,对着源码改一改基本不会出大问题。

最后再分享一个心态上的建议:先在小数据集上跑通全流程,再上全量数据。我一开始直接在PEMS04上训练,报错之后每次调试要等三五分钟才知道结果,效率极低。后来改成先从数据里抽30个节点、训练10个epoch验证流程正确性,确认代码逻辑没问题再上全量数据,整体复现速度快了一倍以上。

如果你也在做时空预测相关的复现工作,希望这篇能帮你少走一些弯路。模型选型上建议从STGCN这类经典模型入手,跑通之后再往Graph WaveNet、PDFormer这些进阶模型迁移,会顺很多。

本文还有配套的精品资源,点击获取

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询