简介:时空预测模型 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 复现前必须确认的三件事
动手写代码之前,先把几个问题想清楚,否则后面全是坑:
- 预测什么?是预测下一个时间步的值,还是未来12个步长的序列?这个决定模型输出层的设计。我这次做的是多步预测,一次输出未来12个时间步。
- 数据长什么样?每个样本需要包含时间窗口的历史数据(比如过去12个步长)和对应的未来值。输入张量的形状是
(batch_size, 输入步长, 节点数, 特征维度)。 - 用什么指标评价?交通预测领域最常用的是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_size | 64 | 显存够用可以调大,加快训练 |
| learning_rate | 0.001 | 初始学习率,配合余弦退火衰减 |
| epochs | 100 | 配合早停,实测40~60轮左右收敛 |
| optimizer | Adam | betas默认(0.9, 0.999) |
| loss | MSE | 回归任务默认选择 |
| input_len | 12 | 用过去1小时数据(5分钟粒度) |
| output_len | 12 | 预测未来1小时 |
| hidden_dim | 64 | 隐藏层维度,越大拟合能力越强 |
| num_layers | 2 | 时空块数量 |
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不下降的处理
时空图卷积网络因为涉及矩阵乘法叠加,梯度爆炸的情况不算少见。我常用的三板斧:
- 梯度裁剪:
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5),在backward之后、optimizer.step()之前执行。 - 降低学习率:从0.001直接调到0.0005,往往就能避免训练早期的不稳定。
- 检查输入数据:对数据做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这些进阶模型迁移,会顺很多。
本文还有配套的精品资源,点击获取