尧图网站设计 尧图网站设计YAOTU DESIGN
ARTICLE DETAIL

资讯详情

深耕网站设计与一线实操的经验洞察。

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

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, biasTrue): super().__init__() self.linear nn.Linear(in_features, out_features, biasbias) 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_size3, dilation1): super().__init__() self.dilation dilation self.padding (kernel_size - 1) * dilation self.conv nn.Conv1d(channels_in, channels_out, kernel_size, paddingself.padding, dilationdilation) self.gate nn.Conv1d(channels_in, channels_out, kernel_size, paddingself.padding, dilationdilation) 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, biasFalse) 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, kt3): 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上面的SpatialConvLayer用torch.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_layer、spatial_conv_layer和st_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是特征维度单一流量传感器时F1T是历史窗口长度N是节点数。转换过程通常使用np.expand_dims和transpose完成。import numpy as np def create_sequences(data, input_len, pred_len, step1): # 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决定了样本重叠程度也直接决定训练集大小。如果step1相邻样本仅滑动一个时间步样本量最大但会产生极强的时序相关性如果step12对应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, axis1) 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_matadj原始矩阵中adj[i][j]表示节点i和j之间的道路连接或距离倒数。加自环是为了让节点在聚合时保留自身特征否则第一步聚合会丢失中心节点信息。d_inv_sqrt是度矩阵的负二分之一次方公式等价于D^{-1/2} A D^{-1/2}。在PyTorch中如果直接把这个矩阵转为torch.FloatTensor作为SpatialConvLayer的adj参数即可。还有一个容易被忽视的问题交通数据中部分时间段的传感器可能无记录通常用NaN表示。在构造滑动窗口之前一定要先做一个掩码或者用前后均值填充。否则那些NaN会通过多个时间步扩散变成模型里的一大片NaN且很难从损失函数值中发现。代码分析时我通常会在create_sequences之前加入一个data np.nan_to_num(data, nannp.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(), lr0.001, weight_decay1e-5) scheduler torch.optim.lr_scheduler.StepLR(optimizer, step_size10, gamma0.7)lr0.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_norm3.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。注意out与y的对齐取决于你模型的输出排列我一般会在模型定义里让输出形状与输入一致即(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块之后的残差连接前防止过拟合一个常见的调法先固定T12、隐藏单元64、batch_size32跑20轮看loss曲线再逐步调整kt和ST块数量。如果验证loss在训练后期震荡先调低学习率或增大weight_decay如果模型输出全是均值即流量预测成常数大概率是图卷积层的学习率过低或邻接矩阵没有加自环。5. 调试STGCN代码的几个落地技巧调试STGCN和调试普通CNN不太一样因为问题往往出在维度排列和邻接矩阵上而不是网络不收敛。5.1 先用小数据跑通forward验证张量形状拿到STGCN-PyTorch-master.zip后不要直接跑全量数据。我一般会构造一个最小样例N5个节点T12步batch_size2只用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_bycuda_time_total, row_limit15))常见的性能瓶颈是torch.matmul(adj, x)特别是当adj是稠密(N, N)而N很大时计算量是BTN^2*C。如果节点数超过1000建议把adj转成稀疏矩阵或者改用DGL/PyG里的spmm算子。另一个瓶颈是频繁的permute和reshape产生的显存拷贝可以通过在输入阶段一次性把布局固定为(B, F, T, N)内部不再交换节点维来避免。这些调试技巧不依赖特定项目版本适用于大多数STGCN-PyTorch实现。按照“小数据验证形状 - 观察梯度 - profile时间”的顺序走一遍代码分析才算真的落地。本文还有配套的精品资源点击获取
返回列表