PyTorch Geometric 数据集切分实战指南:RandomNodeSplit 与 RandomLinkSplit 深入解析
2026/9/12 12:51:06 网站建设 项目流程

PyTorch Geometric 数据集切分实战指南:RandomNodeSplit 与 RandomLinkSplit 深入解析

【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric

图机器学习中,数据集切分(Dataset Splitting)是决定模型评估是否可靠的关键步骤:我们需要把数据集划分为训练、验证和测试三个子集,从而在防止过拟合的同时,准确衡量模型的泛化能力。本文围绕 PyTorch Geometric(PyG)中三种基础任务——节点预测、链接预测与图预测——系统讲解数据集切分的主流方案,重点剖析RandomNodeSplitRandomLinkSplit两个内置变换的完整参数、底层实现原理与典型用法,并介绍如何基于业务场景创建自定义切分。读完本文,你将能够在 PyG 中熟练完成节点级、边级和图级的数据集切分,并理解其背后的掩码(mask)与负采样机制。

为什么图数据需要专门的切分策略

与传统表格数据不同,图数据中的样本并非相互独立:节点通过边相互关联,边又依赖于节点的存在。如果简单地随机切分,训练集、验证集与测试集之间可能发生信息泄漏(data leakage)——例如测试集中出现与训练集共享连边的节点,模型就能"偷看"到测试信息。因此,PyG 针对不同任务提供了语义明确的切分工具:

任务类型切分粒度核心工具
节点预测(Node Prediction)节点RandomNodeSplit
链接预测(Link Prediction)RandomLinkSplit
图预测(Graph Prediction)整张图数据集内建索引 / sklearn / numpy

其中RandomNodeSplitRandomLinkSplit均为 PyG 的 transform(变换),可同时作用于同构图Data与异构图HeteroData对象(对应 torch_geometric/data/data.py 与 torch_geometric/data/hetero_data.py)。

节点预测:使用 RandomNodeSplit 划分节点

在节点级任务(如半监督节点分类)中,我们通常给每个节点打上"是否属于训练 / 验证 / 测试集"的布尔掩码。RandomNodeSplit变换正是为此设计:它会在DataHeteroData对象上附加train_maskval_masktest_mask三个属性,对应 torch_geometric/transforms/random_node_split.py 的实现。

核心参数说明

  • split:切分类型,可选"train_rest"(默认)、"test_rest""random"三种,源码中以assert split in ['train_rest', 'test_rest', 'random']校验;
  • num_splits:要生成的切分数目(默认1)。当大于 1 时,掩码形状为[num_nodes, num_splits],否则为[num_nodes]
  • num_train_per_class:在"test_rest""random"模式下,每个类别采样的训练节点数(默认20);
  • num_val:验证集节点数;传入浮点数时表示占节点总数的比例(默认500);
  • num_test:测试集节点数;传入浮点数时表示占节点总数的比例(默认1000),仅在"train_rest""random"模式下生效;
  • key:存放真实标签的属性名(默认"y")。设置为None时对所有节点存储切分;否则只对含有该属性的节点存储执行切分。

三种 split 模式的语义

从 random_node_split.py 的实现可以清楚看到三种模式的区别:

  • "train_rest":先对全部节点做随机置换(torch.randperm),前num_val个作验证、随后num_test个作测试,剩余全部节点作训练。这与 FastGCN 论文中的设置一致;
  • "test_rest":先按类别为每个类各取num_train_per_class个节点作训练(保持类别平衡),再从剩余节点中取num_val个作验证,其余全部作测试。这与 "Pitfalls of Graph Neural Network Evaluation" 论文的设置一致;
  • "random":同样按类别采样训练节点,再取验证节点与测试节点各num_val/num_test个,其余节点不参与任何掩码(常用于 GCN 论文 "Semi-supervised Classification with Graph Convolutional Networks" 的协议)。

num_valnum_test为浮点数时,源码会执行round(num_nodes * ratio)将其换算为节点数,例如传0.1表示取全部节点的 10%。

动手实践:8 节点小图示例

官方教程给出了一个完整的可运行示例(8 个节点、取 2 个验证、3 个测试):

import torch from torch_geometric.data import Data from torch_geometric.transforms import RandomNodeSplit x = torch.randn(8, 32) # Node features of shape [num_nodes, num_features] y = torch.randint(0, 4, (8, )) # Node labels of shape [num_nodes] edge_index = torch.tensor([ [2, 3, 3, 4, 5, 6, 7], [0, 0, 1, 1, 2, 3, 4]], ) # 0 1 # / \/ \ # 2 3 4 # | | | # 5 6 7 data = Data(x=x, y=y, edge_index=edge_index) node_transform = RandomNodeSplit(num_val=2, num_test=3) node_splits = node_transform(data)

变换完成后,train_maskval_masktest_mask会作为属性挂载到图数据上:

node_splits.train_mask >>> tensor([ True, False, False, False, True, True, False, False]) node_splits.val_mask >>> tensor([False, False, False, False, False, False, True, True]) node_splits.test_mask >>> tensor([False, True, True, True, False, False, False, False])

本示例中 8 个节点中取节点0, 4, 5作为训练集,节点6, 7作为验证集,节点1, 2, 3作为测试集。底层逻辑(见 random_node_split.py)是:为每个掩码初始化全Falsetorch.bool张量,通过随机置换索引定位并置True

训练循环中如何使用掩码

一个真实可运行的代表性用例是 examples/cora.py,它在加载 Cora 数据集时通过T.Compose组合了RandomNodeSplitTargetIndegree变换:

transform = T.Compose([ T.RandomNodeSplit(num_val=500, num_test=500), T.TargetIndegree(), ]) dataset = Planetoid(path, dataset, transform=transform) data = dataset[0]

训练时只需用掩码索引即可完成监督信号的选取与精度计算:

F.nll_loss(model()[data.train_mask], data.y[data.train_mask]).backward() # ... for _, mask in data('train_mask', 'test_mask'): pred = log_probs[mask].max(1)[1] acc = pred.eq(data.y[mask]).sum().item() / mask.sum().item()

多折切分与异构图支持

num_splits > 1时,每个store上会通过列表推导对每折分别执行_split,再以torch.stack(..., dim=-1)堆叠成[num_nodes, num_splits]的掩码(见 random_node_split.py)。对应测试 test/transforms/test_random_node_split.py 中对train_resttest_restrandom三种模式以及浮点比例、异构图(如data['paper'].train_mask)均做了断言验证,例如校验三组掩码互不重叠且并集覆盖全集。

链接预测:使用 RandomLinkSplit 划分边

链接预测任务需要判断两个节点之间是否存在(或属于何种类型的)连边,因此切分对象是RandomLinkSplit会一次性返回三个数据对象train_data, val_data, test_data(而非附加掩码),其核心保证是:训练切分不含验证与测试边,验证切分不含测试边,从而避免标签泄漏。实现位于 torch_geometric/transforms/random_link_split.py。

核心参数说明

  • num_val:验证边数量;浮点数表示占边总数的比例(默认0.1);
  • num_test:测试边数量;浮点数表示比例(默认0.2);
  • is_undirected:是否将图视为无向图(默认False)。为True时,正负样本不会在不同切分间泄漏反向连边;对二部边类型或edge_type != rev_edge_type时该选项被忽略;
  • key:边标签属性名(默认"edge_label")。若该属性不存在,会自动创建并视为二分类任务(1表示边、0表示非边);若存在,则须为0num_classes - 1的类别标签,负采样后0表示负边,1num_classes表示正边类别;
  • split_labels:为True时将正负标签分别存入pos_edge_label/neg_edge_label及对应_index(默认False);
  • add_negative_train_samples:是否为训练集添加负样本(默认True)。若模型内部已自行做负采样,应设为False,否则同一批负样本会在每次迭代中重复使用;
  • neg_sampling_ratio:负边数量相对正边数量的比例(默认1.0);
  • disjoint_train_ratio:大于0时,将训练边拆分为"用于消息传递"与"用于监督"两部分,避免消息传递与监督信号重叠(默认0.0);
  • edge_types/rev_edge_types:作用于HeteroData时指定要切分的边类型及其反向边类型,以确保反向边被一致切分、防止数据泄漏。

底层切分流程

从 random_link_split.py 的实现可以看到完整流程:

  1. 随机置换边索引:无向图下先通过edge_index[0] <= edge_index[1]取无向边的一半,再做torch.randperm;有向图直接置换全部边;
  2. 划分三个切分:按比例(int(num_val * perm.numel()))切出train_edgesval_edgestest_edges,并保证num_train > 0,否则抛出"Insufficient number of edges for training"
  3. 构造消息传递边:训练与验证阶段使用训练边train_edges,测试阶段使用训练边与验证边的并集train_val_edges(这一点在教程中有明确强调:测试时可以基于训练+验证边的并集传播信息);
  4. 负采样:调用 torch_geometric/utils/_negative_sampling.py 中的negative_sampling(edge_index, size, num_neg_samples=..., method='sparse')生成负边;若负边数量不足,会调整采样比例并发出warnings.warn(见test_random_link_split_insufficient_negative_edges对应测试);
  5. 生成标签_create_label将正负边拼接为edge_labeledge_label_index(或pos_*/neg_*属性)。

动手实践:8 节点小图示例

沿用教程示例,为每条边附带标签edge_y

import torch from torch_geometric.data import Data from torch_geometric.transforms import RandomLinkSplit x = torch.randn(8, 32) # Node features of shape [num_nodes, num_features] y = torch.randint(0, 4, (8, )) # Node labels of shape [num_nodes] edge_index = torch.tensor([ [2, 3, 3, 4, 5, 6, 7], [0, 0, 1, 1, 2, 3, 4]], ) edge_y = torch.tensor([0, 0, 0, 0, 1, 1, 1]) # 0 1 # / \/ \ # 2 3 4 # | | | # 5 6 7 data = Data(x=x, y=y, edge_index=edge_index, edge_y=edge_y) edge_transform = RandomLinkSplit(num_val=0.2, num_test=0.2, key='edge_y', is_undirected=False, add_negative_train_samples=False) train_data, val_data, test_data = edge_transform(data)

切分结果如下:

train_data >>> Data(x=[8, 32], edge_index=[2, 5], y=[8], edge_y=[5], edge_y_index=[2, 5]) val_data >>> Data(x=[8, 32], edge_index=[2, 5], y=[8], edge_y=[2], edge_y_index=[2, 2]) test_data >>> Data(x=[8, 32], edge_index=[2, 6], y=[8], edge_y=[2], edge_y_index=[2, 2])

注意这里key='edge_y'会生成edge_y_index属性用于评估,而默认key='edge_label'时生成的则是edge_label_index。语义总结:

  • train_data.edge_indexval_data.edge_index用于消息传递的边:训练与验证阶段只允许基于训练边传播信息;
  • 测试阶段基于训练边 + 验证边的并集传播信息;
  • val_data.edge_label_indextest_data.edge_label_index各保存一批正负样本,专用于评估和测试模型。

结合真实示例:Cora 链接预测

examples/link_pred.py 给出了端到端的链接预测管线,其中切分配置为:

transform = T.Compose([ T.NormalizeFeatures(), T.ToDevice(device), T.RandomLinkSplit(num_val=0.05, num_test=0.1, is_undirected=True, add_negative_train_samples=False), ]) dataset = Planetoid(path, name='Cora', transform=transform) train_data, val_data, test_data = dataset[0]

应用该变换后,数据集元素从单个Data对象变成(train_data, val_data, test_data)三元组。由于设置了add_negative_train_samples=False,训练阶段每个 epoch 都重新进行负采样,避免固定负样本导致模型退化:

neg_edge_index = negative_sampling( edge_index=train_data.edge_index, num_nodes=train_data.num_nodes, num_neg_samples=train_data.edge_label_index.size(1), method='sparse') edge_label_index = torch.cat([train_data.edge_label_index, neg_edge_index], dim=-1) edge_label = torch.cat([ train_data.edge_label, train_data.edge_label.new_zeros(neg_edge_index.size(1)) ], dim=0)

评估阶段则直接使用data.edge_label_indexdata.edge_label计算 AUC 等指标。

异构图与无向图场景

HeteroData上使用RandomLinkSplit时,必须指定edge_types,否则会抛出ValueError(见 random_link_split.py);同时建议通过rev_edge_types声明反向边类型,确保反向边被同步切分而不泄漏。is_undirected=True时,切分边会通过torch.cat([edge_index, edge_index.flip([0])], dim=-1)恢复双向表示(random_link_split.py)。这些行为在 test/transforms/test_random_link_split.py 中均有对应的test_random_link_split_on_hetero_datatest_random_link_split_on_undirected_hetero_data等用例覆盖。

图预测:按图划分数据集

在图级任务中,每张图都是一个独立样本,通常需要按比例将整个图数据集划分为训练、验证、测试子集。PyG 的部分数据集(如 PPI)已经内置了对应的划分索引,直接按split参数加载即可:

from torch_geometric.datasets import PPI path = './data/PPI' train_dataset = PPI(path, split='train') val_dataset = PPI(path, split='val') test_dataset = PPI(path, split='test')

从 torch_geometric/datasets/ppi.py 的源码可以看到,PPI 数据集通过assert split in ['train', 'val', 'test']校验参数,并从原始文件中分别加载train_graph.json/valid_graph.json/test_graph.json及对应的特征、标签与图编号文件,天然支持train / val / test三种划分。教程配套的完整训练示例可参考 examples/ppi.py。

对于没有内建划分的数据集,可以直接借助scikit-learntrain_test_splitnumpy的随机索引完成划分,例如对dataset的索引数组做np.random.permutation后按比例切片,再将划分结果传入Subset构造子数据集。

创建自定义切分

当随机切分无法满足特定业务场景时(这在真实工业界非常常见),可以自行构造自定义切分。教程给出的典型场景是电商领域的超大规模异构图:节点代表用户、商品、商家等多种类型,业务上往往需要按"新老用户"划分数据集,以评估模型对新用户的泛化能力——这类按节点属性、时间戳或社区结构划分的需求,随机切分无法直接表达。

PyG 为此提供了充分的灵活性:你可以在加载数据后手动为Data/HeteroData的节点存储(node_stores)或边存储(edge_stores)直接赋值train_maskval_masktest_mask等布尔张量,其形状与数值完全由业务规则决定;也可以编写自定义BaseTransform子类并注册为函数式变换,在forward中实现任意切分逻辑后与内置变换一起组合使用。切分完成后,训练、验证、测试阶段与使用内置变换完全一致——PyG 的训练循环只关心掩码或edge_label_index是否存在。

小结

本文围绕节点预测、链接预测与图预测三类任务,完整梳理了 PyG 数据集切分的三种主流方案:

  1. 节点预测:使用RandomNodeSplitData/HeteroData上附加train_mask/val_mask/test_mask,支持train_resttest_restrandom三种语义与整数 / 浮点两种计数方式;
  2. 链接预测:使用RandomLinkSplit返回训练、验证、测试三个数据对象,自动完成边划分、负采样与标签生成,并严格保证切分间不泄漏边信息;
  3. 图预测:优先利用数据集内建划分(如 PPI),或借助 sklearn / numpy 按比例切分图集合。

掌握这些工具的关键在于理解其底层机制:节点切分的本质是布尔掩码,边切分的本质是带消息传递约束的索引划分与负采样,而自定义切分则赋予你完全掌控数据划分的能力。通过 test/transforms/test_random_node_split.py 与 test/transforms/test_random_link_split.py 中的测试用例,可以进一步验证上述行为在不同参数组合下的正确性。

【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询