简介:这是一份基于PyTorch的时空图卷积网络(STGCN)实现代码,源自IJCAI 2018论文官方工程,面向人体行为分析、动作识别等方向的研究者与开发者,解决骨骼序列数据中空间拓扑关系与时间动态规律的联合建模问题。压缩包共12个文件,以Python源码、Markdown文档和备份文件为主,整体约28.71MB,目前已有121人学习下载。资源提供了标准化的数据处理流程,包括关节坐标归一化与坐标系转换,并采用模块化方式搭建图卷积层、时序卷积层及全连接分类器,训练过程使用交叉熵损失和Adam优化器,辅以动态学习率调整。通过阅读源码与运行实验,可深入理解时空特征并行提取的实现技巧,掌握从数据预处理到模型评估的完整PyTorch项目开发范式,为智能监控、人机交互等应用场景提供可直接扩展的算法基础。
1. STGCN 是什么:用图卷积吃透路网拓扑,不是又一种“CNN 换皮”
有人说 STGCN 就是把 CNN 换了个卷积核,这句话误导了不少人。STGCN 时空图卷积网络(Spatio-Temporal Graph Convolutional Network)解决的是路网交通流预测这类自带拓扑结构的回归问题:每个监测点是一条时间序列,点与点之间又隔着一条真实路网。普通卷积神经网络把规则网格当输入,天然吃不下这种“节点 + 边”的数据。STGCN 用图卷积吃掉空间依赖,用门控时间卷积吃掉时序依赖,两条线叠成 ST-Conv Block,在 PyTorch 里几十行代码就能搭出可用版本。这篇笔记写给要自己写实现、调参数、复现基线的人,默认你已经会跑 PyTorch,但不一定懂图卷积。读完你能照着代码跑通一个最小模型,也知道真正上线前该在哪里较劲。
2. 先把网络拆开看:STGCN 的图卷积与时间卷积在 PyTorch 里各是什么形状
STGCN 的代码实现难不在 PyTorch 语法,难在把“图上的卷积”和“时序上的卷积”的维度对齐。我见过太多人把两套卷积直接拼在一起,维度报错后就开始瞎 permute。先把两条支线拆清楚,后面写模型就是按图索骥。
2.1 图卷积支线:邻接矩阵、归一化拉普拉斯与 Chebyshev 展开
普通卷积依赖平移不变性,但路网没有这个概念:每个节点的邻居数量不一样,邻居是谁也不一样。图卷积的常见做法是用邻接矩阵 A 描述节点连接关系,然后通过归一化控制消息传递的尺度。
最常用的形式是先把 A 加上自环得到 \tilde A,再算对称归一化拉普拉斯:
L = I - D^{-1/2} \tilde A D^{-1/2}
D 是 \tilde A 的度矩阵。这一步把节点特征从“绝对数值”变成“与邻居的差异”,避免度数大的节点把消息放大。STGCN 原文用的是 Chebyshev 多项式近似:把卷积核 g_θ 近似成 K 阶 Chebyshev 多项式,避免对拉普拉斯做昂贵的特征分解。展开式是 T_0(x)=x、T_1(x)=\tilde L x、T_k(x)=2\tilde L T_{k-1}(x)-T_{k-2}(x),这里的 \tilde L 是特征值缩放到 [-1, 1] 的拉普拉斯矩阵。
在 PyTorch 里,图卷积的输入特征通常长成 (B, C, T, N):B 是批次,C 是通道数,T 是时间步,N 是节点数。图卷积只该作用于 N 这一维,所以要把张量换到 (B, T, N, C) 或者 (B·T, N, C),用矩阵乘法把 L 和节点特征乘起来,时间维和通道维全程不动。K 阶里 K=2 到 3 就够用,继续加大很容易过平滑——信号在图上传来传去,最后所有节点趋向同一个值。
2.2 时间卷积支线:门控一维卷积为什么比 LSTM 更适合短序列
STGCN 的时间支线不是 RNN,而是一维卷积加门控。代码实现里通常写成Conv2d,卷积核是 (Kt, 1),也就是只在时间维上滑,节点维不动。门控机制采用 GLU:卷积输出拆成两半 P 和 Q,最终输出 P ⊙ σ(Q)。σ 是 sigmoid,它的取值范围给模型一个“选择保留多少信息”的能力,比单纯线性卷积表达能力更强。
有人问为什么不直接用 LSTM。短序列场景下(比如用过去 12 步预测未来 12 步),LSTM 的迭代式推理训练慢、梯度路径长,还不容易复现。卷积的优点是并行度高、行为可预期。跑实验时你会发现同一套超参在不同 seed 下 LSTM 结果能差好几个点,STGCN 的波动小得多。做基线和上线都更省心。
这里有一个容易忽略的关键点:普通 Conv1d 默认是“非因果”的,填充会让卷积核看到未来时刻。做短程预测时 padding 设为 Kt // 2,序列长度不变,影响不大。但要改成自回归式预测,就得把 padding 换成左填充,或者手工 mask 掉未来位置,否则推理时信息泄漏,线上表现会翻车。
2.3 ST-Conv Block 的拼装逻辑:残差、瓶颈与张量形状流转
STGCN 主体由几个 ST-Conv Block 堆叠。每个 Block 的结构是“时间卷积 → 图卷积 → 时间卷积”,中间夹残差连接。图卷积放在两个时间卷积中间,出发点很朴素:先让每个节点在时间上把自己理清楚,再沿图结构融合邻居信息,最后再在时间维度上把融合后的特征还原。
用我常用的参数排列说明张量流转。输入 (B, C_in, T, N),第一层时间卷积输出 (B, C_hidden, T, N),图卷积保持 (B, C_hidden, T, N),第二层时间卷积输出 (B, C_out, T, N)。残差连接从输入直接接到输出,如果 C_in 不等于 C_out,就用一个 1x1 卷积先对齐通道。这里可以做一个瓶颈设计:第一层时间卷积把通道压到 hidden,最后再放大到 out,让中间的图卷积在较低维度下工作。
| 参数 | 含义 | 常见取值 |
|---|---|---|
| K | Chebyshev 阶数,决定空间感受野 | 2 ~ 3 |
| Kt | 时间卷积核宽度 | 3 ~ 5 |
| C_hidden | 图卷积输入通道 | 32 ~ 64 |
| num_pred | 预测步数 | 6 / 12 |
这个结构在 PyTorch 里没有魔法,只有一组被反复验证过的 Conv2d 和 matmul。把形状对齐了,别的都好说。
3. 搭一个最小可运行的 STGCN 模型:PyTorch 实现代码与逐段说明
这一章直接给可复现代码。我按“输入形状约定 → 邻接矩阵处理 → 模型定义 → 前向验证”的顺序写。先把形状约定说死,后面所有代码都以这个为准。
3.1 输入输出形状约定:先把 (B, C, T, N) 这四维对齐
我习惯统一成 (B, C, T, N),其中 B 是 batch,C 是每个时刻的输入通道(单特征流量就是 1),T 是历史时间窗口长度,N 是监测点数量。模型输出是 (B, num_pred, N),表示每个节点未来 num_pred 步的预测值。
如果你的原始数据是 (T_total, N, 1),构造样本时要做两步:先用滑动窗口切出 X 和 Y,再把维度从 (B, T, N, F) 转成 (B, F, T, N)。这一步很多新手漏掉,PyTorch 的 Conv2d 默认第四维是空间宽度,直接把 (B, T, N, F) 丢进去,卷积会在时间和节点两个维度上同时滑动,结果完全不对。
def to_samples(raw, hist_steps=12, pred_steps=12): # raw: (T_total, N, F) 按时间顺序排列 total = raw.shape[0] x, y = [], [] for i in range(total - hist_steps - pred_steps + 1): x.append(raw[i: i + hist_steps]) # (hist_steps, N, F) y.append(raw[i + hist_steps: i + hist_steps + pred_steps]) x = torch.tensor(np.array(x)) # (B, T, N, F) y = torch.tensor(np.array(y)) # (B, pred_steps, N, F) x = x.permute(0, 3, 1, 2) # (B, F, T, N) y = y.permute(0, 3, 1, 2) # (B, F, pred_steps, N) return x, y这段代码关键点是最后的 permute。原始切片里每个样本是 (T, N, F),堆叠后是 (B, T, N, F),转成 (B, F, T, N) 才能喂给后面的 Conv2d。如果你的输入有多个特征(比如流量、速度、占有率),F > 1,就保留这一维;如果只有流量,F=1,后续取 y[:, 0] 即可。
3.2 邻接矩阵归一化代码:这一步错了后面全错
图卷积的质量基本由邻接矩阵决定。常见的数据集给的是检测点之间的距离矩阵,先按阈值转成 0/1 邻接矩阵,再算归一化拉普拉斯。不归一化直接喂给 Chebyshev 展开,特征值范围会超过 [-1,1],多项式递归几轮后数值直接爆炸,Loss 变成 NaN。
import numpy as np import torch def build_scaled_laplacian(adj): """ adj: (N, N) 的 0/1 邻接矩阵,先加自环再做对称归一化 返回: (N, N) 的 torch float32 张量,特征值范围落在 [-1, 1] 附近 """ n = adj.shape[0] adj = adj + np.eye(n) # 加自环 deg = adj.sum(axis=1) deg_inv_sqrt = np.power(deg, -0.5) deg_inv_sqrt[np.isinf(deg_inv_sqrt)] = 0.0 # 孤立点保护 d_mat = np.diag(deg_inv_sqrt) lap = np.eye(n) - d_mat @ adj @ d_mat # 对称归一化拉普拉斯 lam_max = np.linalg.eigvals(lap).real.max() # 求最大特征值 lap_scaled = 2.0 * lap / (lam_max + 1e-6) - np.eye(n) return torch.tensor(lap_scaled, dtype=torch.float32)两点说明。第一,自环必须加,否则一个节点的更新完全忽略自身信息,模型等于只学邻居置换。第二,用np.linalg.eigvals求最大特征值在小图上没问题,N 到几千时很慢,可以用幂迭代法替代,只求最大的那个特征值。实际数据里相邻检测点距离远大于阈值时,邻接矩阵会非常稀疏,这种稀疏性后面可以转成 PyTorch 稀疏张量省显存。
3.3 图卷积层、门控时间卷积层与网络主干的完整代码
先把两个基础模块写出来。ChebConv 实现 K 阶 Chebyshev 图卷积,TemporalConv 实现 GLU 门控时间卷积。
import torch.nn as nn class ChebConv(nn.Module): def __init__(self, in_channels, out_channels, K): super().__init__() self.K = K self.weight = nn.Parameter( torch.randn(K, in_channels, out_channels) * 0.05) self.bias = nn.Parameter(torch.zeros(out_channels)) def forward(self, x, lap): # x: (B, C, T, N) lap: (N, N) b, c, t, n = x.shape x = x.permute(0, 2, 3, 1) # (B, T, N, C) if self.K == 1: xs = [x] else: x0 = x x1 = torch.matmul(lap, x) # 一阶邻居消息 xs = [x0, x1] for _ in range(2, self.K): x2 = 2.0 * torch.matmul(lap, x1) - x0 xs.append(x2) x0, x1 = x1, x2 out = sum( torch.einsum('btnc,co->btno', xk, self.weight[k]) for k, xk in enumerate(xs) ) out = out + self.bias return out.permute(0, 3, 1, 2) # (B, C_out, T, N)这里torch.matmul(lap, x)的广播规则是:lap 是 (N, N),x 是 (B, T, N, C),PyTorch 会把前两维当成 batch 维,自动在每组 (N, C) 上做矩阵乘法。这一步把每个节点的一跳邻居特征聚合回来。Chebyshev 的递推直接照公式写,2.0 是多项式系数,K 越大感受野越广。
class TemporalConv(nn.Module): def __init__(self, in_channels, out_channels, Kt): super().__init__() self.conv = nn.Conv2d( in_channels, out_channels * 2, kernel_size=(Kt, 1), # 时间维卷积,节点维不动 padding=(Kt // 2, 0), ) def forward(self, x): y = self.conv(x) # (B, 2*out, T, N) p, q = torch.chunk(y, 2, dim=1) # 沿通道维拆成两半 return p * torch.sigmoid(q)GLU 的核心在两行:通道翻倍是给门控预留空间,torch.chunk把结果平均切成 P 和 Q,p * sigmoid(q)是门控输出。Kt 建议用奇数,比如 3 或 5,配合padding=Kt//2能保持时间长度不变。
最后组装 STGCN 主体。
class STConvBlock(nn.Module): def __init__(self, in_channels, hidden_channels, out_channels, K, Kt): super().__init__() self.tconv1 = TemporalConv(in_channels, hidden_channels, Kt) self.cheb = ChebConv(hidden_channels, hidden_channels, K) self.tconv2 = TemporalConv(hidden_channels, out_channels, Kt) self.residual = nn.Conv2d(in_channels, out_channels, 1) \ if in_channels != out_channels else nn.Identity() def forward(self, x, lap): res = self.residual(x) out = self.tconv1(x) out = self.cheb(out, lap) out = self.tconv2(out) return out + res # 残差连接 class STGCN(nn.Module): def __init__(self, num_nodes, in_channels, hidden_channels, out_channels, K, Kt, num_pred): super().__init__() self.block1 = STConvBlock(in_channels, hidden_channels, hidden_channels, K, Kt) self.block2 = STConvBlock(hidden_channels, hidden_channels, hidden_channels, K, Kt) self.output = nn.Conv2d(hidden_channels, num_pred, 1) def forward(self, x, lap): x = self.block1(x, lap) x = self.block2(x, lap) x = self.output(x) # (B, num_pred, T, N) return x[:, :, -1, :] # (B, num_pred, N)最后一个 1x1 卷积把通道扩成 num_pred,只取最后一个时刻的输出,直接映射到未来 num_pred 步。这种“一步到位”的多步预测叫 one-shot 策略,STGCN 原文就是这么做的。跑前向验证:
b, c_in, t, n = 64, 1, 12, 307 x = torch.randn(b, c_in, t, n) adj = (np.random.rand(n, n) < 0.01).astype(float) # 模拟稀疏邻接 lap = build_scaled_laplacian(adj) model = STGCN(num_nodes=n, in_channels=1, hidden_channels=32, out_channels=32, K=3, Kt=3, num_pred=12) y = model(x, lap) print(y.shape) # torch.Size([64, 12, 307])输出形状是 (64, 12, 307),含义是 64 个样本、每个节点未来 12 步的预测值。到这里模型已经通了,剩下的是训练和评估,这一步最容易让人产生“我在调参其实在碰运气”的错觉。
4. 训练与评估:STGCN 的几个必调参数和“看着收敛其实过拟合”的时刻
模型能跑和模型训练出有效结果之间隔着一条数据处理的河。图卷积类的模型对数据顺序和归一化极其敏感,训练阶段最常见的错误是:Loss 下降得漂亮,换到测试集立刻崩掉。
4.1 训练集 / 验证集 / 测试集怎么切,归一化参数在哪里拟合
时间序列数据绝对不能用随机切分。我见过有人用train_test_split(random_state=42)切交通数据,把 3 月某天的样本放进训练集,4 月的同班车样本放进测试集,评估结果虚高得离谱。正确做法是严格按时间顺序切。
total = raw.shape[0] train_end = int(total * 0.7) val_end = int(total * 0.8) train_raw = raw[:train_end] mean = train_raw.mean(axis=(0, 1), keepdims=True) std = train_raw.std(axis=(0, 1), keepdims=True) + 1e-6 def normalize(raw): return (raw - mean) / std train_X, train_Y = to_samples(normalize(raw[:train_end]), 12, 12) val_X, val_Y = to_samples(normalize(raw[train_end:val_end]), 12, 12) test_X, test_Y = to_samples(normalize(raw[val_end:]), 12, 12)注意mean和std只从训练段计算,验证和测试段复用。这要求训练段数据分布能代表整体。如果流量有很强的星期周期性,训练集最好覆盖完整的周一到周日,否则节假日的预测会整体漂移。
4.2 损失函数与评估指标:MAE、RMSE、MAPE 的实现细节
交通流数据大量存在缺失值和零值。很多公开数据集用 0 表示无车流,这和真正的“零流量”语义重叠。直接用nn.L1Loss()会让缺失位置参与梯度更新,把模型往平均值方向拉。我习惯先写一个带 mask 的损失函数。
def masked_mae(pred, true, null_val=0.0): mask = (true != null_val).float() m = 1e-4 + mask.sum() return (torch.abs(pred - true) * mask).sum() / m def masked_rmse(pred, true, null_val=0.0): mask = (true != null_val).float() m = 1e-4 + mask.sum() return torch.sqrt(((pred - true) ** 2 * mask).sum() / m) def masked_mape(pred, true, null_val=0.0): mask = (true != null_val).float() denom = torch.abs(true) + 1e-4 # 防止分母为 0 return ((torch.abs(pred - true) / denom) * mask).sum() / (1e-4 + mask.sum())MAPE 分母里的1e-4是关键。零流量时刻算出的百分比可能是百分之几千,直接把整个指标带歪。加常数会让 MAPE 不再“纯净”,但换来了可比性。报告中同时给出 MAE 和 RMSE,让读者能从绝对误差和粗差异两个角度判断。
训练循环里我一般用 Adam 初始学习率 0.001,配上ReduceLROnPlateau,patience 设 5。批量大小在显存允许范围内尽量取大,图卷积的einsum对 batch 扩张很敏感,同样数据量跑 32 batch 和 128 batch 的收敛速度差距明显。
optimizer = torch.optim.Adam(model.parameters(), lr=1e-3, weight_decay=1e-4) scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau( optimizer, patience=5, factor=0.5, verbose=True) for epoch in range(100): model.train() total_loss = 0.0 for xb, yb in train_loader: # xb: (B, F, T, N) yb: (B, F, pred_steps, N) pred = model(xb, lap) loss = masked_mae(pred, yb[:, 0]) optimizer.zero_grad() loss.backward() optimizer.step() total_loss += loss.item() val_loss = evaluate(model, val_loader, lap) scheduler.step(val_loss)4.3 参数调整优先级:K 阶、隐层通道、时间核宽度的顺序
STGCN 的可调参数不少,新手最容易乱调。我调参时的固定顺序是:
第一,K 阶。K 从 2 提到 3,观察验证集 MAE 变化。如果没变化,说明路网空间依赖只有一跳有效,继续加大 K 只会增加过平滑风险。如果明显下降,可以试 K=4,但不要超过 5。
第二,Kt 时间核宽度。Kt=3 是起步,如果历史窗口长(比如 24 步),Kt=5 甚至 7 有效。这里有个容易忽略的交互:Kt 变大后,padding 也要跟着变,否则时间维长度缩水,后面的残差连接形状对不上。
第三,隐层通道数。从 32 起步,显存允许就试 64 和 128。这个参数与数据集节点数强相关,N=307 的 METR-LA 和 N=1700 的 PEMS-BAY 最优值差别很大。节点越多,单个通道能表达的模式越有限,通常需要更大通道数。
最后调学习率和 dropout。STGCN 这类模型对 dropout 不敏感,0 到 0.3 之间影响不大,建议先不动。
5. STGCN 避坑指南:PyTorch 实现里最常见的 5 个翻车现场
下面的坑我都踩过,按“现象 → 原因 → 解决”写,每条都能对应到具体报错或异常指标。
5.1 Loss 从一开始就不降,甚至第一个 epoch 就 NaN
现象:训练第一步打印的 loss 是 inf 或 nan,或者前几个 epoch 稳定下降后突然爆掉。原因:拉普拉斯矩阵没有做缩放,特征值超出 Chebyshev 多项式预期的 [-1, 1] 范围,递推项数值指数膨胀。解决:用build_scaled_laplacian里的lam_max做缩放;或者放弃 Chebyshev,改用 GCN 风格的归一化D^{-1/2} \tilde A D^{-1/2},让最大特征值天然落在 2 附近,稳定性更好。
5.2 在验证集上 MAPE 突然变成 300%
现象:MAE 看着正常,MAPE 却高得离谱。原因:真实流量中有很多接近 0 的值,MAPE 的分母是真实值,微小分母会把误差放大到无法解读。解决:评估时单独统计“非零流量区间”的 MAPE,或者把分母改成max(true, threshold),threshold 按数据集分布取 10 或 20。报告里别只给一个 MAPE,把 MAE 放旁边一起看。
5.3 节点数一上千,训练 OOM
现象:N=1700 的数据集在 2080Ti 上 16G 显存跑不满一个 batch。原因:einsum('btnc,co->btno', ...)生成中间张量 (B, T, N, C_out),N 翻一倍显存就翻一倍,而且 Chebyshev 的每一阶都会保留一份。解决:先把 batch 降到 16,再用混合精度;如果还不够,把 lap 转成torch.sparse_coo_tensor,matmul在稀疏张量上会省很多内存和计算。
5.4 换一个数据集,预测结果变成一条平线
现象:训练 loss 正常,但所有节点的预测值都收敛到整体均值附近。原因:节点特征顺序和邻接矩阵顺序没有对齐。公开交通数据集的检测器编号和矩阵行列顺序经常不一致,直接读进来等于随机打乱了图结构。解决:在数据处理脚本里加一个assert adj.shape[0] == node_order.shape[0],并且打印前 5 个节点的邻居编号做人工核验。这个检查 30 秒能做完,能省一整天排错时间。
5.5 PyTorch 环境装了好几天,代码始终跑不到 GPU 上
现象:torch.cuda.is_available()返回 False,模型一直吃 CPU。原因:PyTorch 和 CUDA 版本不匹配,或者驱动太旧。解决:装 GPU 版 PyTorch 前先查nvidia-smi支持的 CUDA 版本,再按对应版本安装。装完后用一段小矩阵乘法验证,不要一上来就跑完整模型。环境稳定是 STGCN 复现的前提,环境问题不解决,后面所有调参都是白费。
6. 把 STGCN 从“能跑”做到“可信”:多步预测、消融验证与结果可视化
模型跑通只是第一步。我会再补两个验证动作,第一个是可视化单节点预测对比,第二个是消融。
可视化代码很简单,但能暴露很多指标掩盖的问题:
import matplotlib.pyplot as plt with torch.no_grad(): pred = model(test_X, lap).cpu().numpy() # (B, num_pred, N) true = test_Y.numpy()[:, 0] # (B, num_pred, N) node_idx = 42 plt.figure(figsize=(8, 4)) plt.plot(true[0, :, node_idx], label='ground truth') plt.plot(pred[0, :, node_idx], label='STGCN') plt.legend() plt.savefig('stgcn_pred_vs_truth.png', dpi=150)看这张图不是看曲线贴得有多紧,而是看预测是不是滞后一拍——如果真实曲线在拐点处总是提前或滞后一个时间步,说明模型没有学到趋势,只是在做平滑复制。这种情况下一百行调参代码都救不回来,需要回看 Kt 和门控机制。
消融实验我一般做三组:去掉图卷积,只用时间卷积;把 GLU 的 sigmoid 固定成 1,退化成普通卷积;把 K 从 3 降到 1。三组对照能直接回答“空间信息到底贡献了多少”。这种验证比调参更值得投入,因为它决定了 STGCN 在你这个数据集上是不是正确的选择——如果去掉图卷积性能不变,那说明路网结构对预测没有帮助,你需要的只是一个时序模型。
我现在拿到新的时空预测任务,第一件事不是调 K,而是先画一个节点 24 小时的真实流量曲线,再拿验证集预测叠上去看形状。形状对不上,参数调得再漂亮也是自欺欺人。这个习惯帮我少掉过很多自己骗自己的实验。希望帮到你。
本文还有配套的精品资源,点击获取