图神经网络归纳式学习实战:GraphSAGE与动态图新节点预测

图神经网络归纳式学习实战:GraphSAGE与动态图新节点预测
训练好的图神经网络上线第一个月表现完美第二个月新数据涌入效果突然开始断崖式下跌。这是很多 GNN 开发者都会撞上的“鬼故事”。问题往往不是模型选错了也不是超参数没调好而是从头就选错了评测范式。翻开网上大多数图神经网络教程跑的都是 Cora、Citeseer 这类固定图数据集模型训练时把整张图几乎都看了一遍测试时只是“蒙住某些节点的标签”。这种设定叫转导式学习Transductive Learning它和真实业务中“新节点不断出现”的场景根本不是一回事。真正支撑工业界动态场景的是另一种范式——归纳式学习Inductive Learning。它要求模型在训练阶段完全看不到测试节点等到推理时再对从没见过的节点“现学现卖”。理解并实现归纳式学习才是从 GNN 入门迈向 GNN 实战的关键一步。这篇文章是“图神经网络实战”系列的第 10 篇。我会把归纳式学习从概念到代码讲透转导式和归纳式到底差在哪为什么 GraphSAGE 天生适合归纳场景以及如何用 PyTorch GeometricPyG在 PPI 数据集上跑通一个完整的归纳学习实验。全程会给出可复制的代码并附上常见问题和工程建议。1. 这篇文章真正要解决的问题先看一个典型场景。假设你在做“大气污染图神经网络”相关项目把城市里的空气质量监测站当作图节点站点之间的地理邻近关系或气象影响关系作为边目标是预测某个站点在未来几小时的 PM2.5 浓度。项目初期你在一张 100 个站点的图上训练了一个 GNN效果不错。问题来了两个月后城市里新增了 10 个监测站点你需要预测新站点的污染浓度怎么办如果模型是转导式训练出来的它只“认识”原来 100 个节点对应的 embedding。新站点没有属于自己的 embedding模型根本无从下手。你只能把 110 个站点的数据拼在一起重新训练一遍。如果每隔一两个月就新增站点你的团队就会陷入“永远在重训模型”的循环。归纳式学习解决的就是这个问题。它的核心思路是模型学习的不是“某个节点长什么样”而是“如何根据一个节点的特征和它邻居的特征来生成预测”。只要新节点也有特征、也有邻居模型就能直接推理不需要重新训练。这篇文章适合三类读者已经跑通过 GCN 或 GAT 入门教程但不知道如何把模型用到动态图上的人。做论文复现时发现数据集的划分方式与“归纳式”相关却始终搞不清区别的人。需要在推荐系统、交通流量预测、传感器网络、知识图谱等场景中处理新增节点的工程师。读完本文你能完成三件事第一准确判断自己的任务到底该用转导式还是归纳式第二理解 GraphSAGE 和邻居采样为什么能支持归纳推理第三独立在 PyG 中用 PPI 数据集跑通一个完整的归纳学习实验。2. 基础概念与核心原理2.1 先花一分钟回忆 GNN 在做什么图神经网络的核心机制是消息传递Message Passing。每一层每个节点都会聚合自己邻居的特征生成新的节点表示h_v^(l1) UPDATE( h_v^(l), AGG( { h_u^(l) : u ∈ N(v) } ) )其中N(v)是节点 v 的邻居集合AGG是聚合函数如 sum、mean、maxUPDATE通常是线性变换加非线性激活。层数越多节点能看到的“视野”越广两层之后一个节点就能间接感知到二跳邻居的信息。这是 GNN 一切能力的起点。关键问题在于这个聚合过程所使用的图到底是“训练时已经全部暴露”的图还是“推理时仍然可能长出新节点”的图这正是转导式和归纳式的分野。2.2 转导式学习与归纳式学习的本质区别先给出一个直白的判断转导式学习训练时模型可以看到整张图的节点特征和边结构只是看不到测试节点的标签。归纳式学习训练时模型完全接触不到测试节点更看不到测试图。它只能从训练图或训练子图中学习一套通用聚合规则推理时把这套规则套用到新节点上。用一个类比来理解。转导式学习像是开卷考试题目里的主人公你全都见过只是不知道谁考了多少分你要根据大家平时表现推测分数。归纳式学习像是闭卷考试换了一批新学生你只能从以前教过的学生身上总结“好学生都有什么特征”然后在新学生进考场时凭他们的表现特征直接判断。这里有一个非常容易踩的误区很多人把 Cora 数据集的默认切分当成归纳式实验。实际上Cora 默认的 train/val/test mask 只是隐藏了标签整张图的节点特征和边结构在训练时都参与了消息传递。也就是说测试节点的特征在训练阶段就已经通过邻域聚合进入了模型参数。这是标准的半监督转导式节点分类并不是真正的归纳式评测。2.3 为什么 GCN 通常被归为转导式GCN 的传播公式是H^(l1) σ( Â · H^(l) · W^(l) )其中Â是对加了自环的邻接矩阵做对称归一化得到的矩阵。很多教材实现会预先在全图上计算这个Â然后直接做矩阵乘法。这意味着模型在训练开始之前就已经“看过”整张图的拓扑包括测试节点的连接关系。严格来说GCN 架构本身并不天然排斥归纳式使用。如果改用 PyG 中GCNConv这样的局部消息传递实现并配合正确的训练协议GCN 也能处理新节点。但标准的科研评测协议和大量教程实现让 GCN 成了转导式范式的典型代表。GraphSAGE 则从设计之初就放弃了“全图预计算”这条路。它在每一层只对节点的局部邻居做采样和聚合训练时用 mini-batch 完成推理时只要给新节点构造一个局部子图就能前向传播。因此GraphSAGE 成为讨论归纳式学习时绕不开的基准模型。2.4 两种范式对比对比维度转导式学习归纳式学习测试节点特征/结构训练时可见仅标签隐藏训练时完全不可见训练方式通常在全图上计算传播训练图上随机采样 mini-batch适用场景静态图上的节点分类动态图、跨图泛化、新节点预测推理方式直接读取节点的最终 embedding为新节点构造局部计算图再前向代表模型GCN、GAT教材版实现GraphSAGE、GIN、GraphSAINT典型数据集Cora、Citeseer、PubMedPPI、ogbn-arxiv时间切分3. 环境准备与前置条件本文的代码基于 PyTorch 和 PyTorch GeometricPyG。建议使用 Python 3.8 及以上版本PyTorch 与 PyG 的版本需要匹配。以下安装命令以 CPU 版为例如果你使用 GPU请先到 PyTorch 官网选择与本地 CUDA 版本匹配的安装命令再安装 PyG。# 创建并激活虚拟环境 conda create -n gnn python3.9 -y conda activate gnn # 安装 PyTorchCPU 版示例 pip install torch # 安装 PyTorch Geometric pip install torch-geometric安装完成后可以先运行下面这段代码验证环境python - PY import torch import torch_geometric from torch_geometric.nn import SAGEConv print(torch version:, torch.__version__) print(pyg version:, torch_geometric.__version__) print(SAGEConv imported successfully) PY如果一切正常你会看到 torch 和 torch_geometric 的版本号以及一行SAGEConv imported successfully。这里有一点需要提醒PyG 的某些扩展算子如 torch-scatter、torch-sparse在部分模型中是可选依赖可能需要从预编译 wheel 安装。本文用到的SAGEConv和DataLoader属于基础功能一般情况下上面的安装方式已经够用。若后续安装扩展算子请以 PyG 官方文档为准避免版本不兼容导致的编译错误。4. 数据准备归纳式实验的正确姿势4.1 为什么 Cora 默认切分不是归纳式Cora 是 GNN 入门最常见的数据集包含 2708 个节点、5429 条边、7 个类别每个节点有一个 1433 维的稀疏词袋特征。PyG 加载后数据对象自带train_mask、val_mask、test_mask看起来“训练集、验证集、测试集”都有了但这里的 mask 只控制标签是否可见不控制图结构是否可见。在标准 GCN 训练流程中模型每一轮都会用data.edge_index在这 2708 个节点的全图上做消息传递。测试节点的特征会通过边流入训练节点的邻居聚合过程从而被模型间接“记住”。这种设定下测试节点的特征分布已经进入了训练信号评测出来的精度自然比真正面对新节点时要乐观。因此如果你要做严格的归纳式实验第一件事就是重新审视数据划分方式。4.2 适合归纳式评测的公开数据集推荐两个最常用的归纳式数据集。第一个是 PPIProtein-Protein Interaction。它包含 24 张蛋白质相互作用图其中 20 张用于训练、2 张用于验证、2 张用于测试。每张图都是独立的一整张蛋白质网络测试图在训练阶段完全没有出现过。节点特征是 50 维标签是 121 维的多标签也就是说每个节点可以同时属于多个类别。PPI 是 PyG 内置的归纳式多标签分类数据集也是本文代码示例的主角。第二个是 ogbn-arxiv。它把论文按时间切分2017 年及之前的论文用于训练2018 年作为验证集2019 年作为测试集。这种方法模拟了“未来论文是未知的”这一真实场景比随机切分更能反映模型的实际泛化能力。4.3 如何在任意数据集上构造归纳式切分当你需要在业务数据上做归纳式训练时最稳妥的方法是按时间或社区划分节点集合然后构造“训练子图”。训练子图只保留训练节点内部的边删除所有连向验证/测试节点的边。下面是用 Cora 演示的构造代码# 文件路径build_inductive_split.py import torch from torch_geometric.datasets import Planetoid # 加载 Cora 全图 dataset Planetoid(root/tmp/Planetoid, nameCora) data dataset[0] num_nodes data.num_nodes perm torch.randperm(num_nodes) train_nodes perm[:1500] val_nodes perm[1500:2000] test_nodes perm[2000:] train_mask torch.zeros(num_nodes, dtypetorch.bool) val_mask torch.zeros(num_nodes, dtypetorch.bool) test_mask torch.zeros(num_nodes, dtypetorch.bool) train_mask[train_nodes] True val_mask[val_nodes] True test_mask[test_nodes] True # 训练子图只保留两端都在训练节点集合内的边 src, dst data.edge_index keep train_mask[src] train_mask[dst] train_edge_index data.edge_index[:, keep] # 训练特征也只看训练节点 train_x data.x[train_nodes] print(f总节点数: {num_nodes}) print(f训练节点: {len(train_nodes)}, 验证节点: {len(val_nodes)}, 测试节点: {len(test_nodes)}) print(f训练子图边数: {train_edge_index.shape[1]}原图: {data.edge_index.shape[1]})这段代码的关键逻辑是keep train_mask[src] train_mask[dst]只有一条边的两个端点都属于训练节点集合时这条边才会进入训练子图。这样训练时模型就接触不到任何测试节点的特征和连接关系。推理时再用同样的方式为每个新节点抽取它的局部邻居子图喂给模型计算输出。5. 核心方法拆解邻居采样与 GraphSAGE5.1 邻居采样解决什么问题在全图上做 GCN 消息传递每一层都要处理整张图的邻接矩阵。对于几亿节点的大规模图这几乎不可能落地。GraphSAGE 提出了一个更实用的思路既然一个节点的表示只依赖于它的 k 跳邻居那我训练时就没必要看整张图只要为每个训练节点随机采样固定数量的邻居形成一棵“局部计算树”即可。这种采样方式叫邻居采样Neighbor Sampling。常见的设置是用 fanout 数组控制每一层的采样数量例如[10, 25]表示第一层为每个节点采样 10 个邻居第二层为每个邻居再采样 25 个二阶邻居。采样后的子图大小不再随全图规模增长而是受 fanout 控制训练内存因此变得可控。更重要的是这棵局部计算树是模型推理时的通用模板。新节点出现时我们同样为它采样 k 跳邻居构造一棵新的计算树然后让模型在这棵树上做消息传递。模型见过的“计算树形态”是局部的所以它能处理任何位置的新节点。这是归纳式能力的直接来源。5.2 GraphSAGE 的聚合方式GraphSAGE 的节点更新公式可以写成h_v^(k) σ( W^(k) · CONCAT( h_v^(k-1), AGG( { h_u^(k-1) : u ∈ N(v) } ) ) )注意它把节点自身表示和邻居聚合结果做了拼接这意味着模型既保留了自己的信息又吸收了局部结构信息。论文中给出了三种聚合函数Mean 聚合器对邻居表示求平均。这是最基础、最稳定的方式PyG 的SAGEConv默认采用这种方式。LSTM 聚合器把邻居表示当作序列输入 LSTM。它对邻居顺序敏感需要随机打乱邻居顺序表达能力更强但计算量更大。Pool 聚合器先对每个邻居表示做一个 MLP 变换再做逐元素 max 或 mean 池化。它兼顾了表达能力和稳定性。在 PyG 的SAGEConv中你不需要自己实现这些细节。它内部已经按“拼接自身表示 聚合邻居表示 → 线性变换”的标准流程完成计算。你只需要关注层数、隐藏维度和 dropout 等超参数。5.3 归纳能力从哪来归纳式模型和转导式模型的根本区别在于转导式模型可能会隐式学习节点 ID 与输出的映射关系而归纳式模型只学习“如何聚合邻居特征”的函数。后者的参数在整个图中共享不依赖某个节点的位置也不依赖某个图的全局结构。因此只要新节点的特征分布和连接模式与训练数据一致模型就能给出合理的预测。这同时也给工程实践提了一个醒如果新上线节点的特征分布发生了明显偏移任何归纳式模型都会失效。归纳不等于万能它假设的是“世界没有本质变化只是来了新个体”。6. 完整示例代码实现与运行验证下面用一个完整的脚本在 PyG 的 PPI 数据集上训练一个多层 GraphSAGE并直接在两个完全未见过的测试图上评估。# 文件路径inductive_graphsage.py import torch import torch.nn.functional as F from sklearn.metrics import f1_score from torch_geometric.datasets import PPI from torch_geometric.loader import DataLoader from torch_geometric.nn import SAGEConv class GraphSAGE(torch.nn.Module): 多层 GraphSAGE 模型用于归纳式节点多标签分类。 def __init__(self, in_channels, hidden_channels, out_channels, num_layers3, dropout0.2): super().__init__() self.convs torch.nn.ModuleList() self.convs.append(SAGEConv(in_channels, hidden_channels)) for _ in range(num_layers - 2): self.convs.append(SAGEConv(hidden_channels, hidden_channels)) self.convs.append(SAGEConv(hidden_channels, out_channels)) self.dropout dropout def forward(self, x, edge_index): for i, conv in enumerate(self.convs): x conv(x, edge_index) if i len(self.convs) - 1: x F.relu(x) x F.dropout(x, pself.dropout, trainingself.training) return x def main(): device torch.device(cuda if torch.cuda.is_available() else cpu) print(Using device:, device) # 1. 加载 PPI 数据集20 个训练图 / 2 个验证图 / 2 个测试图 train_dataset PPI(root/tmp/PPI, splittrain) val_dataset PPI(root/tmp/PPI, splitval) test_dataset PPI(root/tmp/PPI, splittest) print(ftrain graphs: {len(train_dataset)}, fval graphs: {len(val_dataset)}, ftest graphs: {len(test_dataset)}) print(ffeature dim: {train_dataset.num_features}, flabel dim: {train_dataset.num_classes}) # 2. 数据加载器按图粒度 batch train_loader DataLoader(train_dataset, batch_size2, shuffleTrue) val_loader DataLoader(val_dataset, batch_size2, shuffleFalse) test_loader DataLoader(test_dataset, batch_size2, shuffleFalse) # 3. 模型、优化器、损失函数 model GraphSAGE( in_channelstrain_dataset.num_features, hidden_channels256, out_channelstrain_dataset.num_classes, num_layers3, dropout0.2, ).to(device) optimizer torch.optim.Adam(model.parameters(), lr0.005) loss_fn torch.nn.BCEWithLogitsLoss() # 4. 训练函数 def train(): model.train() total_loss 0.0 total_nodes 0 for batch in train_loader: batch batch.to(device) optimizer.zero_grad() out model(batch.x, batch.edge_index) loss loss_fn(out, batch.y) loss.backward() optimizer.step() total_loss loss.item() * batch.num_nodes total_nodes batch.num_nodes return total_loss / total_nodes # 5. 评估函数micro-F1 torch.no_grad() def evaluate(loader): model.eval() preds, targets [], [] for batch in loader: batch batch.to(device) out model(batch.x, batch.edge_index) preds.append((out 0).int().cpu()) targets.append(batch.y.cpu()) pred torch.cat(preds, dim0).numpy() target torch.cat(targets, dim0).numpy() return f1_score(target, pred, averagemicro) # 6. 训练循环 for epoch in range(60): loss train() if epoch % 5 0 or epoch 59: val_f1 evaluate(val_loader) print(fEpoch {epoch:03d} | Loss {loss:.4f} | fVal Micro-F1 {val_f1:.4f}) # 7. 在完全未见过的测试图上评估 test_f1 evaluate(test_loader) print(fTest Micro-F1: {test_f1:.4f}) if __name__ __main__: main()这段代码有几个关键设计需要单独说明。第一PPI 的每个样本是“一整张图”所以DataLoader的batch_size2表示一次加载两张完整的蛋白质图。PyG 的Batch会自动把多张图拼成一个大图并对节点编号做偏移model(batch.x, batch.edge_index)可以直接处理。第二PPI 是多标签分类每个节点有 121 个 0/1 标签所以输出层维度是 121损失函数必须用BCEWithLogitsLoss而不是单标签分类常用的CrossEntropyLoss。预测时对 logits 做out 0得到 0/1 预测结果。第三评估指标用 micro-F1。多标签场景下准确率这个指标没有意义因为一个节点同时有多个正确标签不能简单比较“预测类别是否等于真实类别”。micro-F1 把所有节点的所有标签汇总成全局的 TP、FP、FN再计算 F1能更公平地反映多标签分类效果。运行脚本python inductive_graphsage.py输出效果类似下面这样具体数值会随随机种子、初始化方式变化Using device: cpu train graphs: 20, val graphs: 2, test graphs: 2 feature dim: 50, label dim: 121 Epoch 000 | Loss 0.6923 | Val Micro-F1 0.3341 Epoch 005 | Loss 0.5684 | Val Micro-F1 0.6548 Epoch 010 | Loss 0.5042 | Val Micro-F1 0.7231 Epoch 015 | Loss 0.4653 | Val Micro-F1 0.7518 Epoch 020 | Loss 0.4417 | Val Micro-F1 0.7634 Epoch 025 | Loss 0.4251 | Val Micro-F1 0.7692 Epoch 030 | Loss 0.4114 | Val Micro-F1 0.7746 Epoch 035 | Loss 0.4038 | Val Micro-F1 0.7780 Epoch 040 | Loss 0.3962 | Val Micro-F1 0.7791 Epoch 045 | Loss 0.3907 | Val Micro-F1 0.7803 Epoch 050 | Loss 0.3854 | Val Micro-F1 0.7811 Epoch 055 | Loss 0.3813 | Val Micro-F1 0.7817 Epoch 059 | Loss 0.3799 | Val Micro-F1 0.7822 Test Micro-F1: 0.7805判断实验是否成功的标志很简单训练 loss 应该持续下降而不是震荡或上升。验证集 Micro-F1 应该在前 20 个 epoch 快速上升随后进入平缓区。测试集 Micro-F1 应该和验证集接近。因为测试图在训练时完全没见过这个数值直接代表模型的归纳泛化能力。如果测试集 F1 明显低于验证集比如低 0.1 以上说明模型过拟合了训练图需要增强正则化或减少模型容量。如果 loss 不下降先检查学习率是否合适再看节点特征是否做了标准化。7. 常见问题与排查思路问题现象可能原因排查方式解决方案训练 loss 不下降学习率设置不当或特征未归一化打印 loss 曲线检查输入 x 的数值范围调整学习率对节点特征做标准化测试 F1 明显低于验证 F1模型过拟合训练图或测试图分布不同对比 val/test 上的指标差距检查数据泄漏增加 dropout使用早停增加训练图数量显存不足OOM单张图过大或 batch_size 过大观察 batch 的 num_nodes 和 edge_index 数量减小 batch_size减少层数引入邻居采样PPI 报标签维度不匹配把多标签任务当成单标签分类处理检查 y 的形状和损失函数类型使用 BCEWithLogitsLoss输出维度改为 121新节点预测效果差新节点特征分布与训练节点不一致对比新旧节点的特征统计量做特征工程或周期性重训练模型训练时各 epoch F1 波动大随机种子不同或 batch 大小太小固定随机种子后多次运行固定 seed多次运行取平均结果新节点是孤立点没有邻居图连接关系缺失检查新节点是否真的连到了已有图结构设计孤立点的兜底策略单独处理特征缺失推理时预处理参数不一致训练时用了全图特征均值做归一化检查预处理代码是否保存了统计量上线前保存特征归一化的 mean/std 参数8. 最佳实践与工程建议8.1 数据划分优先按时间其次按社区构造归纳式切分时尽量不要用纯随机切分。随机切分会把紧密相连的节点打散到训练集和测试集虽然测试节点在训练时被遮挡但它的邻居几乎都在训练集里推理时能获得的信息量可能被高估。更贴近真实场景的做法是按时间划分让模型始终面对“未来的新节点”没有时间信息时再考虑按图连通分量或社区结构划分。8.2 模型设计深度控制在 2 到 3 层GNN 太深会出现过平滑over-smoothing问题随着层数增加所有节点的表示会趋于一致区分度下降。对大多数中小规模图2 到 3 层已经足够。如果你的任务需要更大的感受野优先考虑增加邻居采样范围或使用跳跃连接Jumping Knowledge而不是盲目加层。8.3 采样策略fanout 不是越大越好邻居采样中fanout 控制每个节点每层采样的邻居数量。fanout 太大计算图的节点数呈指数增长很快会耗尽内存fanout 太小模型可能采样不到足够信息。常见做法是第一层采 10 到 25 个邻居越往上层越多一些以补偿信息稀释。具体数值需要结合图的平均度数和任务复杂度验证不要照搬别人的配置。8.4 特征归一化要保留统计量很多 GNN 项目会先对节点特征做标准化例如减去均值再除以标准差。训练时如果用了全图的均值那么推理阶段也要用同一个均值而不是用新节点的实时均值重新计算。最稳妥的做法是训练结束后把特征工程的参数mean、std序列化保存随模型一起上线。8.5 评测结果要报告多次运行的波动范围GNN 训练受随机初始化影响较大。一次运行的结果可能带偶然性尤其在小数据集上。稳妥的做法是用不同的随机种子运行 5 次报告平均 micro-F1 和标准差。这比单次跑出高分更有说服力也更容易发现模型本身的不稳定性。8.6 上线前回答三个问题第一新节点是否真的会持续出现如果图是静态的转导式模型完全可以满足需求没必要增加归纳式实现的复杂度。第二新节点的特征和邻居模式是否与训练数据一致如果业务本身发生了变化再强的归纳模型也需要重新训练。第三推理时的局部子图构造和训练时是否一致很多“上线效果差”的问题根源是训练和推理的数据处理流程不一致。9. 总结与后续学习方向这篇文章真正讲清楚了这样几件事转导式学习和归纳式学习的边界在于训练阶段是否接触测试节点的结构和特征GCN 在标准实现下是转导式的典型代表GraphSAGE 则通过邻居采样天然支持归纳推理严格评测归纳能力必须使用 PPI、时间切分的 ogbn-arxiv 这类数据协议而不是直接拿 Cora 默认 mask 当归纳式实验。代码层面我们用 PyG 完成了一个完整的 GraphSAGE 归纳学习流程加载 PPI 数据集、按图粒度组织 DataLoader、定义多层 SAGEConv 模型、用 BCEWithLogitsLoss 做多标签分类、用 micro-F1 评估两个完全未见过测试图上的泛化效果。这套代码可以直接作为你后续做归纳式实验的模板。下一步可以沿着三个方向深入。第一个方向是模型升级把 SAGEConv 换成 GAT 或 GIN观察不同聚合器对归纳泛化的影响。第二个方向是大规模图研究 Cluster-GCN、GraphSAINT 这类基于采样的训练框架解决数亿节点场景下的内存和效率问题。第三个方向是动态图如果把“新节点出现”再细化成“图结构随时间演化”就需要引入时间维度那是时序图神经网络的研究范畴。最后留一个实用提醒在 GNN 上线之前先问自己一句——这个场景真的会出现新节点吗如果需要你的训练流程里就必须保证测试节点的特征和边从未被模型见过。很多项目在评测时精度很高上线后却翻车根源往往不是模型而是从一开始就把归纳式任务按转导式方式训练了。搞清了这一点你已经比大多数照着教程跑 Cora 的开发者走得更远。

最新新闻

日新闻

周新闻

月新闻