news 2026/9/14 2:22:41

STGCN的PyTorch实现:从图卷积到时间卷积的完整代码解析

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
STGCN的PyTorch实现:从图卷积到时间卷积的完整代码解析

简介:STGCN-PyTorch-master.zip是一套基于PyTorch实现的STGCN(时空图卷积网络)代码包,面向从事人体动作识别、时序数据建模的深度学习开发者与研究者。该模型来自IJCAI 2018论文,采用空间图卷积与时间卷积联合建模,可有效捕捉人体关节拓扑关系及动作动态。压缩包共8个文件,包括3个Python脚本(主程序、工具函数、模型定义)、2个Markdown说明文档、LICENSE和.gitignore,并附带METR-LA数据集压缩包,整体约14.36MB,目录结构清晰,适合入门学习与二次开发。目前已有2343人学习下载,资源涵盖数据加载、模型构建、训练评估等完整流程,还提供了预测演示思路,可直接对照源码理解STGCN核心原理,并迁移至交通流量预测等相关场景。

1. STGCN代码分析:先看懂数据流,再碰模型

给定一个STGCN-PyTorch项目,常见的压缩包命名如STGCN-PyTorch-master.zip,里面通常包含数据预处理脚本、模型定义和训练入口。很多人在做完STGCN的代码分析后都会卡在同一处:单看每个torch.nn.Module都认识,但组合起来维度总是对不上。这个问题的根源在于STGCN不是简单把图卷积和时间卷积串起来,而是有着严格的维度转换约定——(通道, 时间, 节点) 和 (节点, 特征) 之间的排列组合。这里不打算逐行念源码,而是按一条可复现的路径,把STGCN的PyTorch实现拆成数据流、核心模块、训练循环和调试技巧四层,并给出可以直接改着用的代码片段。

2. STGCN核心模块的PyTorch实现:从邻接矩阵到时空卷积

在PyTorch基础框架中,STGCN的实现并不算复杂,但容易把人绕晕的是三个张量:输入特征X的形状、邻接矩阵A的形状、以及中间隐藏状态的形状。典型的STGCN输入是一个形状为(N, F, T)的张量,N是节点数,F是每个节点的特征维度(例如流量、速度、占用率),T是时间窗口长度。而PyTorch的Conv1d期望输入为(B, C, T),所以如果你直接传入(N, F, T)就会报错。常见做法是在模型前先做一次permute,把节点维度放到batch位置,或者把节点和通道合并。下面的分析基于最常见的STGCN-PyTorch实现:将输入视作形状(B, T, N, F),经过一个reshape变成(B, F, T, N),再依次穿过时空卷积块。

2.1 图卷积层:用邻接矩阵实现节点信息聚合

图卷积层的作用是对每个时间步上的节点特征做空间信息传播。STGCN采用的切比雪夫多项式一阶近似,可以写成以下形式:

import torch import torch.nn as nn class GraphConv(nn.Module): def __init__(self, in_features, out_features, bias=True): super().__init__() self.linear = nn.Linear(in_features, out_features, bias=bias) self.sigma = nn.ReLU() def forward(self, x, adj): # x: (B, N, F_in), 每个时间步的节点特征 # adj: (N, N) 归一化邻接矩阵 # 一阶近似公式: sigma(adj @ x @ W) h = torch.matmul(adj, x) # 聚合邻居特征,得到(B, N, F_in) out = self.linear(h) # 特征线性变换,得到(B, N, F_out) return self.sigma(out)

这个实现的关键在于adj @ x是在节点维度上的矩阵乘法。adj的形状是(N, N),x的形状是(B, N, F_in),torch.matmul会把前两维做常规矩阵乘,即对每个batch和特征通道,计算邻居特征的加权和。参数方面,in_features就是输入的通道数,out_features是图卷积输出的通道数。注意这里省略了切比雪夫多项式的scaled Laplacian构造,因为很多实现直接在数据预处理阶段通过D^(-0.5) * A * D^(-0.5)得到了adj,图卷积层只做一次矩阵乘。

如果要对每个时间步分别做图卷积,可以先把维度展开为(B*T, N, F),一次性送入这个模块,再reshape回来。STGCN最早版本的代码就是这么干的:将时间维度与batch维度合并,让一个Linear层同时处理所有时间步。这样做的另一个好处是可以直接调用GPU矩阵乘法,不需要循环。

2.2 时间卷积层:用Conv1d完成因果时序特征提取

时间卷积层用来捕捉时间依赖,STGCN中使用的是带空洞的一维因果卷积。因果卷积要求t时刻的输出只依赖于t以及之前的输入,这在PyTorch中可以通过左侧padding来实现。

import torch.nn.functional as F class TemporalConv(nn.Module): def __init__(self, channels_in, channels_out, kernel_size=3, dilation=1): super().__init__() self.dilation = dilation self.padding = (kernel_size - 1) * dilation self.conv = nn.Conv1d(channels_in, channels_out, kernel_size, padding=self.padding, dilation=dilation) self.gate = nn.Conv1d(channels_in, channels_out, kernel_size, padding=self.padding, dilation=dilation) self.sigma = nn.Sigmoid() def forward(self, x): # x: (B, C, T) if self.padding > 0: conv_out = self.conv(x)[:, :, :-self.padding] gate_out = self.gate(x)[:, :, :-self.padding] else: conv_out = self.conv(x) gate_out = self.gate(x) return conv_out * self.sigma(gate_out)

这段代码使用门控线性单元(Gated Linear Unit, GLU)作为时间卷积的激活机制。conv产生主分支,gate产生门控分支,两者逐元素相乘后作为输出。关键在于padding和右移截断:F.conv1d默认是两端填充,而因果卷积只保留左侧信息,所以要在最后把右侧多出的self.padding列切掉。如果dilation为1,这就是普通因果卷积;当dilation大于1时,卷积核的感受野会指数扩大,适合捕捉长时间跨度。

在STGCN的PyTorch实现中,时间卷积层一般在维度交换后使用。例如ST-Conv Block的输入先被改为(B, C, T, N),再用permute将通道和时间调整到合适的顺序,确保Conv1d作用在时间维上。这里channels_in对应传输到该层时的特征维,channels_out可以设为当前隐藏维度。

2.3 时空卷积块ST-Conv Block与残差连接

单个ST-Conv Block的拓扑是“时间卷积 -> 空间图卷积 -> 时间卷积”,并在两端各接一次批归一化,最后加上残差连接。下面是使用Conv2d实现的一个版本,维度变化都保持在四维张量上,便于阅读:

class TemporalConvLayer(nn.Module): def __init__(self, kt, c_in, c_out): super().__init__() # 卷积核的第二个尺寸为1,表示只沿时间维滑动 self.conv = nn.Conv2d(c_in, c_out, kernel_size=(kt, 1), padding=((kt - 1) // 2, 0)) self.gate = nn.Conv2d(c_in, c_out, kernel_size=(kt, 1), padding=((kt - 1) // 2, 0)) def forward(self, x): # x: (B, C, T, N) return self.conv(x) * torch.sigmoid(self.gate(x)) class SpatialConvLayer(nn.Module): def __init__(self, c_in, c_out, num_nodes): super().__init__() self.theta = nn.Linear(c_in, c_out, bias=False) self.num_nodes = num_nodes def forward(self, x, adj): # x: (B, C, T, N) B, C, T, N = x.shape x = x.permute(0, 2, 3, 1) # (B, T, N, C) x = x.reshape(B * T, N, C) # 邻接矩阵聚合:adj(N,N) @ x(N,F) -> (B*T, N, C) x = torch.matmul(adj, x) x = self.theta(x) # (B*T, N, C_out) x = x.reshape(B, T, N, -1).permute(0, 3, 1, 2) return x class STConvBlock(nn.Module): def __init__(self, c_in, c_out, num_nodes, kt=3): super().__init__() self.tconv1 = TemporalConvLayer(kt, c_in, c_out) self.sconv = SpatialConvLayer(c_out, c_out, num_nodes) self.tconv2 = TemporalConvLayer(kt, c_out, c_out) self.bn = nn.BatchNorm2d(c_out) self.residual = nn.Conv2d(c_in, c_out, 1) if c_in != c_out else nn.Identity() def forward(self, x, adj): # x: (B, c_in, T, N) res = self.residual(x) x = self.tconv1(x) x = self.sconv(x, adj) x = self.tconv2(x) x = self.bn(x) return x + res

上面的SpatialConvLayertorch.matmul(adj, x)做空间聚合,然后用nn.Linear做特征变换。理论上先聚合再线性与先线性再聚合是等价的,但聚合在前能减少线性层的输入规模,便于调试。需要注意的是,adj必须是已经归一化的稠密矩阵或torch.sparse.FloatTensor。当adj为稀疏矩阵时,torch.matmul也能处理,但性能更好的是torch.spmm。如果节点数不大(如200以内),稠密矩阵乘足够快。

下面汇总ST-Conv Block内各层的输入输出形状:

模块输入形状输出形状说明
TemporalConvLayer(B, C_in, T, N)(B, C_out, T, N)时间维卷积+门控
SpatialConvLayer(B, C_out, T, N)(B, C_out, T, N)节点维聚合+特征变换
BN + 残差(B, C_out, T, N)(B, C_out, T, N)稳定训练

实际写代码时,你不需要把图卷积单独拆成一个文件。很多STGCN-PyTorch项目会把temporal_conv_layerspatial_conv_layerst_conv_block放在同一个model.py里,因为三者的参数耦合度很高。在做代码分析时,我习惯先把这个文件读透,再去看main.py里的训练逻辑。

3. 训练数据准备:把路网快照变成图信号序列

3.1 数据切片的维度约定:B, N, F, T

在STGCN代码分析的第一步是理解数据的组织方式。METR-LA和PEMS-BAY这类交通数据集通常给出的是二维矩阵 (T_N, N),T_N是时间步总数,N是传感器节点数。为了生成训练样本,需要用滑动窗口在时间轴上切片,每个样本包含历史T个时间步的流量数据,预测未来T_pred个时间步。输入样本的组织方式在不同开源实现里有差异:有的用(B, T, N, F),有的用(B, N, F, T)。PyTorch中的STGCN实现为了配合Conv2d,常把输入整理成(B, F, T, N),其中B是批次大小,F是特征维度(单一流量传感器时F=1),T是历史窗口长度,N是节点数。转换过程通常使用np.expand_dimstranspose完成。

import numpy as np def create_sequences(data, input_len, pred_len, step=1): # data: (time_steps, num_nodes) samples_x, samples_y = [], [] for i in range(0, len(data) - input_len - pred_len + 1, step): x = data[i : i + input_len] # (input_len, num_nodes) y = data[i + input_len : i + input_len + pred_len] # 转成 (num_nodes, 1, input_len),F固定为1 samples_x.append(x.T[:, np.newaxis, :]) # (N, 1, T) samples_y.append(y.T[:, np.newaxis, :]) # (N, 1, T_pred) return np.array(samples_x), np.array(samples_y)

这里的切片步长step决定了样本重叠程度,也直接决定训练集大小。如果step=1,相邻样本仅滑动一个时间步,样本量最大,但会产生极强的时序相关性;如果step=12(对应5分钟的流量数据,12步恰好1小时),则样本独立性更好,但总量减少。测试时一般用不同的起点做多个预测,因此step可以适当调大以降低计算量。

x.T把时间维转置到最后一维,得到(N, T),然后插入一个新轴到第1维,得到(N, 1, T)。在后续放进DataLoader时,多个样本会堆叠成(B, N, 1, T),再通过x.permute(0, 2, 3, 1)变成(B, 1, T, N)。注意特征维度F固定为1,如果你的数据有多传感器类型(如速度+占有率),切片的第二维就不是1,而是特征数F,此时x.T后要reshape成(N, F, T)

下面这个表格总结了不同阶段的张量布局,方便在代码分析时对照:

数据形态形状用途
原始传感器矩阵(time_steps, num_nodes)直接来自CSV或HDF5
单个输入样本(num_nodes, 1, input_len)Dataset里的x
批处理样本(B, num_nodes, 1, input_len)DataLoader默认输出
模型输入(B, 1, input_len, num_nodes)执行permute(0, 2, 3, 1)

3.2 用Dataset构建STGCN输入样本

在PyTorch基础框架中,Dataset负责将numpy数组包装成可迭代样本。下面是一个极简实现,适合STGCN的输入形态。

from torch.utils.data import Dataset class STGCNDataset(Dataset): def __init__(self, x, y): self.x = x # (num_samples, num_nodes, 1, input_len) self.y = y # (num_samples, num_nodes, 1, pred_len) def __len__(self): return len(self.x) def __getitem__(self, idx): return self.x[idx], self.y[idx]

默认的collate_fn会对第0维进行堆叠,因此一个batch的形状是(B, N, 1, T)(B, N, 1, T_pred)。如果你的模型定义里forward使用的是(B, F, T, N),只需要在训练循环开头加一行x = x.permute(0, 2, 3, 1)。有些实现直接在__getitem__里做permute,比如return self.x[idx].transpose(2, 3),这样会改变输出顺序,容易和标签不一致。我建议保持Dataset返回原始维度,把所有维度变换集中在调用模型的地方,出错时更好排查。

3.3 邻接矩阵归一化与掩码处理

STGCN的图卷积依赖邻接矩阵,常见做法是计算对称归一化的拉普拉斯矩阵:

import numpy as np def normalized_adj(adj): adj = adj + np.eye(adj.shape[0]) # 加自环 d = np.sum(adj, axis=1) d_inv_sqrt = np.power(d, -0.5) d_inv_sqrt[np.isinf(d_inv_sqrt)] = 0.0 d_inv_sqrt_mat = np.diag(d_inv_sqrt) return d_inv_sqrt_mat @ adj @ d_inv_sqrt_mat

adj原始矩阵中adj[i][j]表示节点i和j之间的道路连接或距离倒数。加自环是为了让节点在聚合时保留自身特征,否则第一步聚合会丢失中心节点信息。d_inv_sqrt是度矩阵的负二分之一次方,公式等价于D^{-1/2} A D^{-1/2}。在PyTorch中,如果直接把这个矩阵转为torch.FloatTensor,作为SpatialConvLayeradj参数即可。

还有一个容易被忽视的问题:交通数据中部分时间段的传感器可能无记录,通常用NaN表示。在构造滑动窗口之前,一定要先做一个掩码或者用前后均值填充。否则那些NaN会通过多个时间步扩散,变成模型里的一大片NaN,且很难从损失函数值中发现。代码分析时,我通常会在create_sequences之前加入一个data = np.nan_to_num(data, nan=np.nanmean(data))的兜底处理,并在训练时观察验证集loss是否突然变成nan。

4. STGCN的PyTorch训练循环与超参数调节

4.1 损失函数、优化器与评估指标

STGCN的常见损失函数是平均绝对误差(MAE)或均方误差(MSE)。交通预测任务里MAE更常用,因为它的梯度更平稳,对离群点不敏感。PyTorch中可以直接用torch.nn.L1Loss(),也可以自定义带掩码的版本,用来跳过无效节点。评估指标通常还有MAPE和RMSE,但它们在训练中不作为loss,只做验证参考。

import torch import torch.nn as nn criterion = nn.L1Loss() optimizer = torch.optim.Adam(model.parameters(), lr=0.001, weight_decay=1e-5) scheduler = torch.optim.lr_scheduler.StepLR(optimizer, step_size=10, gamma=0.7)

lr=0.001是STGCN比较常见的起点,weight_decay设为1e-5防止过拟合。由于图卷积层参数较少,主参量在时间卷积层,所以学习率可以按层拆分:用param_groups给不同层设置不同学习率,例如图卷积层用0.0005,时间卷积层用0.001。这是调优时的一个有效手段。

4.2 训练循环:前向传播、反向传播与梯度裁剪

下面给出一个集成了维度变换和验证逻辑的训练循环模板。它可以直接套在STGCN上。

def train_one_epoch(model, dataloader, optimizer, criterion, device, adj): model.train() total_loss = 0.0 for x, y in dataloader: # x: (B, N, 1, T), y: (B, N, 1, T_pred) x = x.permute(0, 2, 3, 1).to(device) # (B, 1, T, N) y = y.squeeze(2).to(device) # (B, N, T_pred) optimizer.zero_grad() out = model(x, adj) # 输出形状需要和y对齐 # 如果模型输出为(B, 1, T_pred, N),转成(B, N, T_pred) out = out.squeeze(1).permute(0, 2, 1) loss = criterion(out, y) loss.backward() torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=3.0) optimizer.step() total_loss += loss.item() * len(x) return total_loss / len(dataloader.dataset)

x.permute(0, 2, 3, 1)(B, N, 1, T)变成(B, 1, T, N),正好喂给STConvBlock。y.squeeze(2)把特征维度去掉,变成(B, N, T_pred),因为预测目标通常只关心数值。clip_grad_norm_max_norm取值范围在1.0到5.0之间,STGCN在训练初期梯度容易出现尖峰,裁剪后能显著减少NaN。注意outy的对齐取决于你模型的输出排列,我一般会在模型定义里让输出形状与输入一致(即(B, F, T_pred, N)),再在训练循环显式转换,这样模型内部不会越改越乱。

4.3 超参数调节的关键点与影响

STGCN的超参数并不算多,但每个参数的连锁反应很大。下面这张表是代码分析时最常需要调整的几项:

超参数典型范围对结果的影响
输入时间窗口T6~24决定感受野,过长会引入噪声,过短则预测不准
ST块数量1~3增加深度但显著提高显存占用和训练时间
时间卷积核大小kt3~7控制时间局部关联,奇数通常配合padding
隐藏单元数32~128每增加一倍,参数量约增加四倍(主要在图卷积后的线性层)
batch_size16~64影响BN统计量和收敛速度,太小容易震荡
dropout0.0~0.3通常放在每个ST块之后的残差连接前,防止过拟合

一个常见的调法:先固定T=12、隐藏单元=64、batch_size=32,跑20轮看loss曲线,再逐步调整kt和ST块数量。如果验证loss在训练后期震荡,先调低学习率或增大weight_decay;如果模型输出全是均值(即流量预测成常数),大概率是图卷积层的学习率过低,或邻接矩阵没有加自环。

5. 调试STGCN代码的几个落地技巧

调试STGCN和调试普通CNN不太一样,因为问题往往出在维度排列和邻接矩阵上,而不是网络不收敛。

5.1 先用小数据跑通forward,验证张量形状

拿到STGCN-PyTorch-master.zip后,不要直接跑全量数据。我一般会构造一个最小样例:N=5个节点,T=12步,batch_size=2,只用1个ST块,一次性打印每一层的输出形状。这一步能把80%的维度错误暴露出来。调试代码片段如下:

model = STGCN(...) x = torch.randn(2, 1, 12, 5) adj = torch.randn(5, 5) # 测试用,实际要归一化 with torch.no_grad(): out = model(x, adj) print(out.shape)

如果模型内部使用SpatialConvLayer,注意adj必须能参与torch.matmul,如果用稀疏格式,要先确保节点数和稀疏索引匹配。最稳妥的做法是在模型第一个ST块前加一个断言:

assert x.shape[2] > model.kt, "输入时间长度必须大于时间卷积核"

5.2 梯度与激活值异常定位

训练中途遇到loss不降,建议观察图卷积输出和梯度范数。

model.sconv.theta.weight.register_hook(lambda grad: print("grad norm:", grad.norm().item()))

这个hook会打印图卷积线性层权重的梯度L2范数。如果范数小于1e-5,说明这个层没有学到信息,可能是邻接矩阵行全部为0,或输入节点特征差异过小。如果梯度为nan,说明前向已经出现了nan,可以用torch.autograd.set_detect_anomaly(True)定位到具体生成的张量。注意这个开关会拖慢训练,只在调试时使用。

5.3 用torch.profiler定位性能瓶颈

当模型能跑通但训练很慢时,用PyTorch自带的torch.profiler看每个模块的时间分布。

from torch.profiler import profile, ProfilerActivity with profile(activities=[ProfilerActivity.CPU, ProfilerActivity.CUDA]) as prof: out = model(x, adj) loss = criterion(out, y) loss.backward() print(prof.key_averages().table(sort_by="cuda_time_total", row_limit=15))

常见的性能瓶颈是torch.matmul(adj, x),特别是当adj是稠密(N, N)而N很大时,计算量是BTN^2*C。如果节点数超过1000,建议把adj转成稀疏矩阵,或者改用DGL/PyG里的spmm算子。另一个瓶颈是频繁的permutereshape产生的显存拷贝,可以通过在输入阶段一次性把布局固定为(B, F, T, N),内部不再交换节点维来避免。

这些调试技巧不依赖特定项目版本,适用于大多数STGCN-PyTorch实现。按照“小数据验证形状 -> 观察梯度 -> profile时间”的顺序走一遍,代码分析才算真的落地。

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

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/14 2:21:15

YOLOV5交通标志识别:从数据集构建到模型部署全流程解析

简介:YOLOv5交通标志识别检测项目是一套面向毕业设计、课程设计与期末大作业的完整资源,包含数据集、源码和预训练模型,帮助开发者快速掌握目标检测项目全流程,避免从零搭建环境的繁琐。资源共266个文件,以Python源码、…

作者头像 李华
网站建设 2026/9/14 2:20:49

YOLOv5+PyQt5人脸表情识别系统实战指南

简介:本资源是一套基于YOLOv5 v7.0实现的人脸表情识别完整工程,面向计算机视觉初学者、深度学习实践者及PyQt界面开发学习者,解决从模型部署到交互式应用落地的关键问题。压缩包共2000个文件,含39个核心Python脚本(如t…

作者头像 李华
网站建设 2026/9/14 2:20:37

STM32F103+FatFs文件系统管理:配置、日志与排错

简介:面向STM32F103单片机开发者的FATFS文件系统管理实验例程,基于HAL库和KEIL环境编写,适合正在学习嵌入式文件系统应用、或需要快速搭建存储管理功能的读者。压缩包内含226个文件,以C源文件与H头文件为主,覆盖FATFS核…

作者头像 李华
网站建设 2026/9/14 2:20:31

WordPress主题CoreNext免授权版安装配置与安全检查指南

简介:CoreNext 1.7.1.1免授权开心版是一套由果核出品的WordPress轻量化主题模板,面向需要快速搭建个人站点或深入研究主题开发的站长、开发者,主打界面简洁、运行高效,且代码全开源,解决了付费主题授权成本高、黑盒不安…

作者头像 李华
网站建设 2026/9/14 2:20:16

手搓教程:工程师的确定性防线与AI协同方法论

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/14 2:20:04

Arnis 世界生成工具:把真实城市搬进 Minecraft 的完整指南

Arnis 世界生成工具:把真实城市搬进 Minecraft 的完整指南 【免费下载链接】arnis Generate any location from the real world in Minecraft with a high level of detail. 项目地址: https://gitcode.com/GitHub_Trending/ar/arnis Arnis 是一款免费开源的…

作者头像 李华