简介基于PyTorch的STGCN时空图卷积网络实现代码源自IJCAI 2018论文官方实现面向从事人体行为分析、骨骼动作识别等方向的研究者与开发者可用于视频监控、人机交互、医疗康复等场景的时空特征建模。压缩包共12个文件以Python脚本为主包含模型定义、训练与工具模块同时提供Markdown文档说明、配置文件与备份文件整体约28.71MB结构清晰便于直接阅读与二次开发。目前已有122人学习过。内容涵盖核心模型文件、训练主程序、数据处理辅助脚本及文档说明将图卷积网络与时间卷积网络有机结合实现了空间关节点拓扑关系与时间序列动态演变的联合建模。通过研读该代码读者能够掌握STGCN的完整流程包括数据预处理、模型构建、损失函数设置、优化器配置及动态学习率调整有助于深入理解时空图卷积的理论基础与工程实践方法同时可在此基础上进行改进实验如引入注意力机制或融合LSTM等提升模型性能。1. STGCN是什么、能做什么一个PyTorch版时空图卷积网络的价值“基于PyTorch的STGCN时空图卷积网络实现代码”——这句话压缩了三件事用PyTorch基础框架写代码、实现STGCN模型、并且是能跑通的实现而不是论文公式复述。STGCNSpatio-Temporal Graph Convolutional Network针对的是传感器网络上的预测问题交通流、城市人流量、电网负载、空气质量。这些数据不同于普通表格数据传感器节点之间有连接关系道路连通、站点相邻每个节点自己又持续产出时间序列两种结构必须同时建模。纯LSTM/Transformer只沿时间轴建模看不到节点间的连接纯GCN只能做空间消息传递抓不住趋势。STGCN把图卷积和时间卷积组合在同一个块里用一份代码同时处理两个维度。这篇文章按我实际复现STGCN的思路展开从结构原理、数据准备、模型代码、训练评估到常见坑给出一份可以直接拿去改用的PyTorch实现。2. 拆解STGCN结构时间卷积、图卷积与残差连接的配合2.1 为什么序列模型在图结构数据上失效先明确问题定义。假设有N个传感器节点每个节点每5分钟采集一个值历史窗口为T个时间步要预测未来h个时间步。整个输入可以组织成张量[B, C, N, T]B是batch sizeC是每个节点的特征通道数N是节点数T是历史步长。对交通流预测来说C通常等于1就是车速或流量但也可能是多维特征例如车速加占有率。把每个节点的历史序列接上位置编码丢给Transformer或者直接上LSTM这是很多人的第一反应。但这样做会丢掉一个关键信息节点之间存在显式的连接结构。道路上两个相邻传感器之间的相关性远大于两个相距很远的传感器这种关系用图邻接矩阵表示最自然。Transformer想学这种关系当然也能学到前提是数据量大到能支撑注意力矩阵去拟合N×N的关系但交通数据往往只有几万到几十万个样本显式结构不用反而指望模型自己发现结果是训练慢、数据效率低。STGCN的出发点恰好相反把邻接矩阵当成先验知识直接参与计算。图卷积在每个时间步上做一次消息传递让相邻节点的特征互相加权融合时间卷积再沿着时间轴提取趋势和周期性变化。两个操作交替比分开建模更容易训练参数量也比Transformer小一个量级。2.2 S-T块内部门控时间卷积与图卷积的先后逻辑STGCN的基本单元叫S-T块Spatio-Temporal Block论文里的经典结构是门控时间卷积 → 空间图卷积 → 门控时间卷积。两个时间卷积夹住一个图卷积形成“先提取时间特征再做空间扩散再提取时间特征”的路径。为什么把时间卷积放在两侧图卷积的输入如果还是原始时间序列节点之间会直接混合原始观测值而相邻节点的原始值又不一定同步容易引入噪声。先过一次时间卷积等于先对每个节点的历史窗口做平滑和特征提取再做空间消息传递时传给邻居的信息已经不是原始读数而是局部趋势特征语义更干净。第二个时间卷积负责在空间融合之后再聚合时间上下文相当于在邻居信息和本节点信息的基础上重新生成时序表示。时间卷积这里实际是一维卷积加GLU门控。一维卷积的卷积核只沿着时间轴滑动kernel_size设为3或5刚好覆盖短期的局部趋势。GLU把卷积输出沿通道维劈成两半一半经sigmoid后作为门控另一半作为内容逐元素相乘相当于让网络决定每个通道有多少信息往下传。一个有趣的现象是GLU和LSTM的输入门在功能上类似但门控值由卷积产生能并行计算训练速度比LSTM快不少。S-T块内部还有残差连接和LayerNorm。残差解决深层堆叠时的梯度衰减问题LayerNorm归一化的是节点维和通道维也就是对每个节点在每个时间步上的特征向量做归一化不归一化时间维这样网络对输入序列长度有一定容忍度预测步长变化时不需要重新训练归一化参数。2.3 用随机张量验证shape流动的最小代码在看完整实现之前先搭一个只含forward逻辑的最小版本用随机张量验证每一步shape。写模型最怕的是卷积padding算错导致时间维对不上最后在某个矩阵乘法上报出维度不匹配。下面这段代码就是用来提前暴露这些问题的初始化模型打印每一层的输出shape不等训练就知道结构有没有接错。import torch import torch.nn as nn import torch.nn.functional as F class DummyTemporalConv(nn.Module): 随机时间卷积只用于验证shape def __init__(self, in_ch, out_ch, kernel_size3): super().__init__() self.conv nn.Conv2d(in_ch, out_ch * 2, kernel_size(1, kernel_size), padding(0, (kernel_size - 1) // 2)) def forward(self, x): return self.conv(x) class DummyGraphConv(nn.Module): 随机图卷积只用于验证shape def __init__(self, in_ch, out_ch): super().__init__() self.linear nn.Linear(in_ch, out_ch) def forward(self, x, adj): B, C, N, T x.shape xt x.permute(0, 3, 1, 2) # [B, T, C, N] xt torch.matmul(adj, xt) # 空间消息传递 xt self.linear(xt) # 特征变换 return xt.permute(0, 2, 3, 1) # [B, C, N, T] B, C, N, T 8, 1, 207, 64 x torch.randn(B, C, N, T) adj torch.randn(N, N) tc DummyTemporalConv(C, 64) out tc(x) print(after temporal conv:, out.shape) # [8, 128, 207, 64] gc_out DummyGraphConv(128, 64)(F.glu(out, dim1), adj) print(after graph conv:, gc_out.shape) # [8, 64, 207, 64]这段代码的关键在F.glu(out, dim1)。时间卷积输出的通道数是64×2GLU按通道维把128个通道劈成两组各64逐元素相乘后输出通道数回到64。adj是[N, N]与xt做矩阵乘法时广播规则会先让xt变成[B, T, C, N]然后在N维上执行adj xt如果adj的N与输入节点数不一致这一行立刻报错。验证通过之后才算把STGCN的空间维度逻辑搞清楚后面再加残差和LayerNorm只是锦上添花。3. 数据准备与邻接矩阵归一化拉普拉斯矩阵与滑窗样本3.1 数据格式与train/val/test划分STGCN需要的原始数据是两张表一张观测矩阵X形状为[T_total, N]或[T_total, N, F]记录每个传感器在每个时刻的读数另一张距离矩阵或邻接矩阵W形状为[N, N]描述传感器之间的空间关系。在动手之前要先确认PyTorch环境是好的。不管用conda还是pip装好的环境只要import torch能通过版本不低于1.9就行。STGCN本身没有特殊的算子需求不需要很新的版本如果在Ubuntu上手动配过环境版本管理上多花点心思训练时反而省事。数据划分有个容易犯的错时间序列数据不能随机打乱后划分必须按时间顺序切。我的习惯是前70%训练中间10%验证最后20%测试。原因很简单模型训练时见过测试区间附近的样本验证和测试都会虚高按时间切分才能模拟真实的“拿过去预测未来”场景。如果手头没有现成数据最常用的两个公开数据集是METR-LA和PEMS-BAY。METR-LA是洛杉矶207个传感器采集的车速数据5分钟粒度PEMS-BAY是旧金山湾区的325个传感器数据。它们的原始观测矩阵都是[T_total, N]距离矩阵可以直接从传感器经纬度计算。3.2 从距离矩阵构造邻接矩阵并加自环构造邻接矩阵有阈值法和KNN法两种常见做法。阈值法逻辑直观两个传感器之间距离小于某个阈值就认为有连接连接权重的计算方式不同最常用的形式是高斯核W[i, j] exp(-dist[i, j]^2 / sigma^2)其中sigma是距离分布的尺度参数通常取所有传感器距离的方差。想省事就取sigma^2 0.1但更稳妥的做法是先统计节点间真实距离的分布取中位数或均值作为sigma。阈值法有个问题距离矩阵很大时双循环构建速度极慢207个节点还好超过1000个节点就建议用向量化实现。邻接矩阵构造完了必须加自环。自环的意思是每个节点到自己的连接权重要保留因为图卷积在聚合邻居信息时如果丢掉自己当前节点的原始特征就完全被邻居覆盖预测结果会震荡。最简单的方式就是W I然后再做归一化。3.3 邻接矩阵归一化与样本生成代码归一化采用对称归一化拉普拉斯形式这是论文里STGCN的标准做法。公式是D^(-1/2) * W * D^(-1/2)作用是让节点的度数差异不影响聚合尺度。一个度很大的枢纽节点聚合邻居特征后数值不至于爆炸一个孤立节点归一化后也不会被稀释。实现代码如下import torch def build_adjacency(dist_matrix, sigma20.1, thresholdNone): 从距离矩阵构造加权邻接矩阵。 dist_matrix: [N, N] 欧氏距离矩阵 sigma2: 高斯核方差 threshold: 距离阈值超过则不建边None则全连接 n dist_matrix.shape[0] adj torch.zeros(n, n) mask torch.ones(n, n) if threshold is None else (dist_matrix threshold).float() adj torch.exp(-dist_matrix ** 2 / sigma2) * mask adj.fill_diagonal_(1.0) # 自环每个节点保留自身信息 return adj def normalize_adj(adj): 对称归一化D^(-1/2) * A * D^(-1/2) adj adj torch.eye(adj.shape[0]) degree adj.sum(dim1) d_inv_sqrt torch.pow(degree, -0.5) d_inv_sqrt[torch.isinf(d_inv_sqrt)] 0.0 D_inv_sqrt torch.diag(d_inv_sqrt) return torch.mm(torch.mm(D_inv_sqrt, adj), D_inv_sqrt)fill_diagonal_把对角线置为1等价于加自环但这样做了之后normalize_adj里又加了一次单位阵两处看起来重复。实际顺序是build_adjacency里保留自环让邻接矩阵更完整normalize_adj里再加一次单位阵是保险写法避免某些实现把对角线清零。如果你明确知道build_adjacency已经处理过对角线normalize_adj里的torch.eye那一行可以去掉只保留归一化部分。滑窗样本生成是STGCN数据准备的最后一步。输入是历史in_steps个时间步预测未来out_steps个时间步。所谓滑窗就是每次把窗口整体向后移动一步。注意这一步很费内存METR-LA共约3.4万个时间步窗口宽度12、预测3步时能生成3万个样本每个样本是[1, N, 12]的浮点矩阵全部加载到内存完全没问题但换成更大规模数据时要用生成器逐个产出。import numpy as np import torch from torch.utils.data import Dataset class STGCNDataset(Dataset): STGCN样本数据集 features: [T_total, N] 原始观测矩阵 in_steps: 历史窗口长度 out_steps: 预测长度 def __init__(self, features, in_steps12, out_steps3): self.in_steps in_steps self.out_steps out_steps self.x, self.y [], [] T features.shape[0] for i in range(T - in_steps - out_steps 1): window features[i: i in_steps] # [in_steps, N] target features[i in_steps: i in_steps out_steps] # [out_steps, N] self.x.append(window.T[np.newaxis, :, :]) # [1, N, in_steps] self.y.append(target.T[np.newaxis, :, :]) # [1, N, out_steps] self.x torch.FloatTensor(np.array(self.x)) self.y torch.FloatTensor(np.array(self.y)) def __len__(self): return len(self.x) def __getitem__(self, idx): return self.x[idx], self.y[idx]样本shape这里要特别说明模型内部约定输入是[B, C, N, T]而window.T[np.newaxis, :, :]生成的形状是[1, N, in_steps]在Dataset里缺少C维训练时还需要unsqueeze(1)扩一个通道维。我倾向于在Dataset的__getitem__里不处理维度把维度变换全部放在训练循环或模型forward里做这样Dataset保持通用换个模型也能复用。真到训练时x.unsqueeze(1)一行就能补齐。4. PyTorch逐层实现STGCN从S-T块到完整网络4.1 输入张量约定与Glorot初始化正式写模型前先把约定定死之后所有代码都按这个来。输入张量x的形状是[B, C, N, T]其中B是batch sizeC是输入特征通道数N是节点数T是历史时间步。C在大多数交通数据集上等于1但模型结构不应该写死后续如果要加天气、节假日特征直接改C即可。邻接矩阵adj的形状是[N, N]在前向过程中参与广播计算。输出的预测形状是[B, N, out_steps]即未来每个时间步每个节点的预测值。所有线性层和卷积层用Glorot初始化这是论文里提到的一个细节实际影响也确实存在。PyTorch的nn.Linear默认初始化是Kaiming均匀分布换成xavier_uniform_后训练初期loss下降更平稳不会有第一个epoch就震荡的情况。4.2 门控时间卷积与图卷积的PyTorch代码时间卷积的核心是Conv2d加GLU。为什么用Conv2d而不是Conv1d因为输入是[B, C, N, T]节点维N在中间Conv1d默认作用在最后一维T上但这样卷积核就无法设置成(1, kernel_size)这样的二维形状。用Conv2d并把卷积核第一维设为1就实现了“只在时间维滑动、不在节点维滑动”的效果这是STGCN实现里最常见的细节。import torch import torch.nn as nn import torch.nn.functional as F class TemporalConv(nn.Module): 门控时间卷积Conv2d GLU in_channels: 输入通道数 out_channels: 输出通道数GLU压缩后 kernel_size: 时间维卷积核大小 def __init__(self, in_channels, out_channels, kernel_size3): super().__init__() self.out_channels out_channels self.conv nn.Conv2d( in_channels, out_channels * 2, kernel_size(1, kernel_size), padding(0, (kernel_size - 1) // 2) ) def forward(self, x): # x: [B, C, N, T] out self.conv(x) # [B, 2*C, N, T] return F.glu(out, dim1) # [B, C, N, T]out_channels * 2是因为GLU需要两个等宽的通道组一组做门控sigmoid一组做内容逐元素相乘。padding设为(kernel_size - 1) // 2kernel_size是奇数时时间维长度保持不变。这里有个容易踩的坑padding不能设成kernel_size // 2偶数卷积核会让时间维变化后面残差相加时shape对不上。图卷积的实现比较灵活。一种写法是用torch.matmul(adj, x)把邻接矩阵和特征张量相乘然后在节点维上做全连接变换。另一种写法是用nn.Conv2d直接对节点维做1×1卷积但邻接矩阵没有真正参与计算等于把邻接关系写死在优化以外不推荐。更接近论文原意的是乘法形式。class GraphConv(nn.Module): 图卷积层邻接矩阵广播矩阵乘法 特征线性变换 in_features: 输入特征维也是上一层的通道数 out_features: 输出特征维 def __init__(self, in_features, out_features): super().__init__() self.linear nn.Linear(in_features, out_features) def forward(self, x, adj): # x: [B, C, N, T], adj: [N, N] B, C, N, T x.shape xt x.permute(0, 3, 1, 2) # [B, T, C, N] xt torch.matmul(adj, xt) # 空间消息传递 [B, T, C, N] xt self.linear(xt) # 特征变换 [B, T, C, N] return xt.permute(0, 2, 3, 1) # [B, C, N, T]permute把时间维从第3位挪到第1位是因为torch.matmul(adj, xt)要求adj的最后一维和xt的N维对齐。矩阵乘法之后用nn.Linear对每个节点的特征向量做变换等价于每个节点独立经过一个全连接层权重在所有节点间共享。adj在这里起到了选择“哪些节点要聚合”的作用归一化后的值就是聚合权重。整个过程不会改变[B, C, N, T]的骨架只是把C维换成了新的out_features。4.3 两层STGCN组装、输出头与参数规模把时间卷积和图卷积组合成S-T块。S-T块包含两个时间卷积和一个图卷积顺序是时间→图→时间对应2.2里讲的结构。残差连接如果输入输出通道数不一致用一个1x1卷积把输入映射到输出通道数再相加。LayerNorm参数要放在[out_channels, N]两个维上。class STConvBlock(nn.Module): STGCN时空卷积块 in_channels: 输入通道数 out_channels: 输出通道数同时也是隐藏维 n_nodes: 节点数N用于LayerNorm kernel_size: 时间卷积核大小 def __init__(self, in_channels, out_channels, n_nodes, kernel_size3): super().__init__() self.tconv1 TemporalConv(in_channels, out_channels, kernel_size) self.gconv GraphConv(out_channels, out_channels) self.tconv2 TemporalConv(out_channels, out_channels, kernel_size) if in_channels ! out_channels: self.residual_conv nn.Conv2d(in_channels, out_channels, kernel_size1) else: self.residual_conv nn.Identity() self.norm nn.LayerNorm([out_channels, n_nodes]) def forward(self, x, adj): residual self.residual_conv(x) out self.tconv1(x) out self.gconv(out, adj) out self.tconv2(out) out out residual return self.norm(out)完整STGCN用两个S-T块堆叠第二个块的输出接一个1x1卷积作为输出层。n_pred是预测步数output卷积核大小为1作用在最后一个时间步上把特征通道映射成预测值。class STGCN(nn.Module): 完整STGCN网络 n_nodes: 传感器节点数 in_channels: 输入通道数通常为1 hidden_channels: 第一层隐藏维 out_channels: 第二层隐藏维 kernel_size: 时间卷积核大小 n_pred: 预测时间步数 def __init__(self, n_nodes, in_channels1, hidden_channels64, out_channels128, kernel_size3, n_pred3): super().__init__() self.block1 STConvBlock(in_channels, hidden_channels, n_nodes, kernel_size) self.block2 STConvBlock(hidden_channels, out_channels, n_nodes, kernel_size) self.output nn.Conv2d(out_channels, n_pred, kernel_size1) for module in self.modules(): if isinstance(module, (nn.Linear, nn.Conv2d)): nn.init.xavier_uniform_(module.weight) def forward(self, x, adj): B, C, N, T x.shape x x.unsqueeze(1) if C 1 else x # 补通道维 out self.block1(x, adj) # [B, 64, N, T] out self.block2(out, adj) # [B, 128, N, T] out self.output(out) # [B, n_pred, N, T] return out[:, :, :, -1] # 取最后一个时间步 [B, n_pred, N]输出层只取最后一个时间步是常见做法论文里原实现是在第二个S-T块后再加一个时间卷积层把时间维压缩到预测长度。用1x1卷积加取最后一个时间步更简单效果差异很小我一般会把这个差异归到超参调试里。参数规模上一个207节点的两层STGCN约50万参数占用显存约200MBbatch size为64和同等规模的Transformer相比小很多这也是STGCN在实际项目中仍然有生命力的原因之一。5. 训练配置与常见问题排查指标计算与5个踩坑记录5.1 Z-Score标准化与逆标准化归一化是时空预测模型最容易出问题的地方。交通数据的取值范围通常在0到100之间车速但不同传感器所在路段限速不同均值和方差差异很大。Z-Score标准化把每个节点拉成均值为0、方差为1的标准分布是STGCN论文及绝大多数复现采用的方式。def zscore_fit(data): data: [T_total, N]按节点计算mean/std mean data.mean(dim0, keepdimTrue) # [1, N] std data.std(dim0, keepdimTrue) 1e-6 return mean, std def zscore_transform(data, mean, std): return (data - mean) / std def zscore_inverse(data, mean, std): return data * std mean标准化必须在训练集上完成再把同一组mean/std应用到验证集和测试集。千万不能在整个数据集上先标准化再切分这样测试集的均值信息会泄漏到训练中。1e-6是为了防止某些传感器全天没车std为0导致除以0。逆标准化在评估指标时用让MAE/RMSE恢复到原始量纲否则汇报的指标没有实际物理意义。5.2 训练循环、学习率与梯度裁剪训练循环不复杂但有几个参数值得固定下来。损失函数用MAEL1损失比MSE效果好原因是交通数据存在尖峰事故、拥堵MSE会放大这些离群点的影响让模型过度拟合异常事件。优化器用Adam初始学习率0.001每5个epoch衰减0.7batch size 64梯度裁剪最大值5.0。这套参数是STGCN复现中最通用的组合不保证最优但保证不会翻车。import torch.optim as optim model STGCN(n_nodesN, in_channels1, hidden_channels64, out_channels128, kernel_size3, n_predout_steps) optimizer optim.Adam(model.parameters(), lr0.001) scheduler optim.lr_scheduler.StepLR(optimizer, step_size5, gamma0.7) criterion nn.L1Loss() for epoch in range(max_epochs): model.train() train_loss 0.0 for X, y in train_loader: # X: [B, in_steps, N]先转成 [B, 1, N, in_steps] X X.permute(0, 2, 1).unsqueeze(1) y y.permute(0, 2, 1) # [B, out_steps, N] optimizer.zero_grad() pred model(X, adj) loss criterion(pred, y) loss.backward() nn.utils.clip_grad_norm_(model.parameters(), 5.0) optimizer.step() train_loss loss.item() scheduler.step()X.permute(0, 2, 1)把[B, T, N]变成[B, N, T]unsqueeze(1)补通道维符合第4章约定的输入格式。预测输出在最后一个时间步上取但因为输出层直接在整个时间维上卷积模型实际上会用到整个历史窗口的信息来生成预测而不是只依赖最后一个时刻。梯度裁剪这里有实际意义。图卷积对邻接矩阵做矩阵乘法一旦输入数据中有几个异常大值梯度经过矩阵乘法后会指数级放大5.0这个阈值能掐住大部分NaN的源头。5.3 MAE、RMSE、MAPE的计算细节评估指标统一在逆标准化后的原始量纲上计算。每次val或test结束把预测值和真实值都逆标准化再算指标不要在标准化空间里汇报。def calc_metrics(y_true, y_pred, eps1e-6): y_true/y_pred: [B, out_steps, N] 原始量纲 mae torch.abs(y_true - y_pred).mean().item() rmse torch.sqrt(((y_true - y_pred) ** 2).mean()).item() mape (torch.abs(y_true - y_pred) / (torch.abs(y_true) eps)).mean().item() * 100 return mae, rmse, mapeMAPE有个经典坑车速为0时除以0会得到无穷大。eps能挡住数值问题但要从业务角度考虑交通预测里车速为0对应长时间堵死这类样本要么在预处理时剔除要么单独统计。我一般会在汇报MAPE时把真实值小于1的样本直接过滤掉否则MAPE永远偏大。5.4 踩坑记录现象、原因、处理办法坑一第一个epoch loss极低之后纹丝不动。现象训练第一轮MAE就降到0.1以下之后几十个epoch几乎没有变化。原因模型根本没学直接把输入的平均值或上一时间步的值当成预测输出。这种情况多发生在输出层只取最后一个时间步时模型偷懒学会“复制最后一个观测值”。处理办法打印第一个batch的预测值和真实值如果预测值出现大量重复检查输出层的卷积核是不是1×1并且在最后一个时间步上取数确认你的预测目标不是“上一帧复制”。坑二图卷积输出NaN。现象训练第2或第3个epochloss变成NaN回退到上一个epoch也无法恢复。原因梯度爆炸邻接矩阵不是归一化形式或者数据里有NaN。处理办法先用torch.isnan(features).any()检查数据再打印邻接矩阵的谱半径torch.linalg.eigvals(adj)最大特征值应小于1.5左右超过这个值说明归一化不对最后加上梯度裁剪。这三个步骤按顺序做90%的NaN都能解决。坑三时间维对不上残差相加报错。现象运行时提示size mismatch通常是[B, 64, N, T]和[B, 64, N, T1]无法相加。原因时间卷积padding设置错误或者kernel_size是偶数。处理办法kernel_size固定为奇数且padding取(kernel_size - 1) // 2。时间卷积有两个图卷积不改变时间维所以整条路径的T一直保持不变。坑四换数据集后指标全面变差。现象在METR-LA上复现还行换成自己数据集后MAE翻了一倍。原因邻接矩阵没重新构造。METR-LA是207个节点你的数据集可能是50个节点距离分布完全不同threshold和sigma2必须重新标定。处理办法先画邻接矩阵的度数分布确认每个节点平均连接数为5到10如果全连接或全不连接模型都学不出空间结构。坑五模型复现不出论文指标。现象同样的模型结构loss降得很慢20个epoch才到论文水平的80%。原因数据划分方式、标准化方式、预测目标都有影响。论文经常是用前70%训练直接预测后30%没有独立的验证集你加了验证集指标自然会略低。处理办法不要追求完全复现论文数字稳定复现论文量级且自己数据上相对传统方法有明显提升就算达标。我在实际项目中STGCN相比LSTM平均能降低10%到15%的MAE这已经值得上线。6. 验证模型是否真的在学随机标签实验与模型导出技巧6.1 随机标签实验判断模型有没有在死记硬背训练跑通不等于模型有效。我最常用的一套验证方法是随机标签实验把训练集的目标y整体随机打乱重新训练同样结构的模型。如果打乱后模型仍能在一两个epoch内把loss降到很低说明模型根本没有依赖输入和标签的关系只是在记忆数据集。打乱后loss应该在正常模型的1.5到2倍以上并且很难下降才说明正常模型确实学到了可泛化的规律。这个小实验中损失函数、学习率、epoch数完全不变只改一个随机种子结果非常直观。随机标签实验之外还要顺带检查样本间的泄漏把测试集里与训练集相邻的时间窗去掉比如时间上间距小于4个步长的样本剔除后再测指标。STGCN的滑窗机制天然会让训练集最后一个样本的输入覆盖测试集第一个样本的部分时间范围这种泄漏会把测试指标抬高好几个点。6.2 ONNX导出与shape下界模型稳定之后把PyTorch模型部署成服务是个常见需求。torch.onnx.export是标准路线但对STGCN有两个注意点一是时间维T必须固定因为模型在最后一个时间步取预测值动态输入会让adj广播的维度推断出错二是输入顺序[B, C, N, T]在导出和推理时要保持一致。以下代码是我常用的导出模板model.eval() dummy_x torch.randn(1, 1, N, 12) # 固定T为12 torch.onnx.export( model, (dummy_x, adj), stgcn.onnx, input_names[x, adj], output_names[pred], dynamic_axes{x: {0: batch_size}}, opset_version12 )导出后建议用onnxruntime加载做一次推理和PyTorch结果比对误差在1e-4以内就算通过。opset_version太低会缺一些算子支持12是安全选项。我现在的习惯是拿到任何时空数据集第一件事不是调模型而是先做三件事——检查邻接矩阵对角线有没有1确认train/val/test按时间切分以及跑一次随机标签实验验证数据管道没写错。这三件事做完STGCN的调试过程会省掉大半时间。希望帮到你。本文还有配套的精品资源点击获取