PyTorch Geometric 数据加载器(torch_geometric.loader)完整 API 指南:从批量训练到大规模图采样
【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric
导读
torch_geometric.loader是 PyTorch Geometric(PyG)中负责把图数据组织成 mini-batch 的核心模块,覆盖了从整图小规模训练到百万节点大规模图采样的全部场景:既包含将多个Data/HeteroData对象合并为 mini-batch 的DataLoader,也包含面向节点、边、时序事件和异构图的大规模采样加载器。本文以官方 API 参考文档 docs/source/modules/loader.rst 为骨架,结合模块源码(torch_geometric/loader)与测试用例(test/loader),系统讲解每个加载器的核心参数、运行机制与实战用法,帮助你按图索骥地选择并正确使用合适的加载器。
说明:
loader.rst是 Sphinxautosummary自动生成的模块 API 索引,其内容即模块中公开类的完整文档;本文逐类展开这些 API,并补充源码实现细节。
一、模块总览:25 个公开类
模块入口 torch_geometric/loader/init.py 通过__all__ = classes = [...]导出了全部公开类,它们按功能可分为以下几组:
| 分组 | 类 | 典型用途 |
|---|---|---|
| 基础批量加载器 | DataLoader、DataListLoader、DenseDataLoader | 把小图集合合并成 mini-batch,适合 Planetoid、TUDataset 等中小数据集 |
| 通用采样器加载器 | NodeLoader、LinkLoader | 基于BaseSampler的通用节点级/链路级采样入口 |
| 邻居采样 | NeighborLoader、LinkNeighborLoader、NeighborSampler | GraphSAGE 式采样,支持同构图与异构图 |
| 异构图采样 | HGTLoader | 保持各节点类型预算均衡的异构采样 |
| 图分区/图级采样 | ClusterData、ClusterLoader、GraphSAINTSampler系列、ShaDowKHopSampler、RandomNodeLoader | Cluster-GCN、GraphSAINT、ShaDow k-hop 等方法 |
| 时序 | TemporalDataLoader | 面向TemporalData事件流的批量加载与负采样 |
| 采样器/装饰器 | ImbalancedSampler、DynamicBatchSampler、PrefetchLoader、CachedLoader、ZipLoader、AffinityMixin | 采样策略与加载性能增强 |
模块同时保留了RandomNodeSampler作为RandomNodeLoader的弃用别名(使用它会触发 deprecation 警告,提示改用loader.RandomNodeLoader)。注意IBMBBatchLoader、IBMBNodeLoader在导出列表中被注释掉(对应测试文件 test/loader/test_ibmb_loader.py 仍存在),说明它们当前并未作为公开 API 开放。
二、基础批量加载器:DataLoader与 Collater 机制
2.1DataLoader
DataLoader继承自torch.utils.data.DataLoader(torch_geometric/loader/dataloader.py),其职责是把Dataset中的图对象合并为 mini-batch。构造参数如下:
dataset:数据来源,可为Dataset、Sequence[BaseData]或DatasetAdapter;batch_size:每个 batch 的样本数,默认1;shuffle:每个 epoch 是否重新打乱数据,默认False;follow_batch:对列表中每个 key 额外生成 batch 赋值向量(如follow_batch=['x']会生成x_batch),用于图分类中对齐多尺度信息;exclude_keys:从 mini-batch 中排除的 key;**kwargs:其余参数透传给torch.utils.data.DataLoader(如num_workers、drop_last)。
与标准 PyTorchDataLoader的本质差异在于collate_fn:源码中DataLoader将Collater(dataset, follow_batch, exclude_keys)注入构造(dataloader.py),并主动kwargs.pop('collate_fn', None)以兼容 PyTorch Lightning 的调用方式。
2.2Collater的类型分发
Collater.__call__(dataloader.py)按 batch 首元素的类型进行分发,这是理解 PyG mini-batch 的关键:
BaseData(Data/HeteroData):调用Batch.from_data_list(batch, follow_batch=..., exclude_keys=...),将多个图沿节点维拼接,并通过batch向量记录每个节点属于哪个图;torch.Tensor:走default_collate;TensorFrame:走torch_frame.cat(batch, dim=0);float/int:分别打包为torch.float32/ 默认 dtype 的张量;str:直接返回列表;Mapping:递归对每个 key 聚合;- 具名元组与一般
Sequence:递归逐元素聚合; - 其他类型:抛出
TypeError。
这种设计使得DataLoader不仅能处理图对象,还能处理特征张量、字符串标签等混合数据,为图级任务(如分子性质预测)提供了统一入口。
2.3 其余基础加载器
DataListLoader:将图列表按"每 batch 一个列表"的方式加载(配合Batch.from_data_list使用),适合图特征尺寸差异较大的场景;DenseDataLoader:面向data.adj稠密邻接矩阵表示的数据,输出DenseBatch,用于图分类中的稠密图批处理。
三、节点级采样:NeighborLoader与HGTLoader
当整图无法放入显存时,需要采样子图进行 mini-batch 训练。这两类加载器都继承自NodeLoader(torch_geometric/loader/node_loader.py),后者是承载通用BaseSampler的抽象加载器。
3.1NeighborLoader:GraphSAGE 式邻居采样
NeighborLoader实现了 GraphSAGE 论文("Inductive Representation Learning on Large Graphs",arXiv:1706.02216)中的邻居采样策略(torch_geometric/loader/neighbor_loader.py)。核心参数:
num_neighbors:每一跳为每个节点采样的邻居数。同构图传List[int],例如[30, 30]表示两跳各采样 30 个邻居;异构图可传Dict[EdgeType, List[int]],对每条边类型分别指定每跳数量;某跳设为-1表示采样该节点的全部邻居;input_nodes:作为采样种子的节点索引,可为LongTensor/BoolTensor;异构图须传(node_type, indices)元组;默认None表示所有节点;input_time:覆盖种子节点时间戳的可选张量,需要同时设置time_attr;replace:是否放回采样,默认False;subgraph_type:返回子图类型,取值"directional"(默认,仅保留计算种子节点表示所需的有向边)、"bidirectional"(转为双向边)、"induced"(所有采样节点的导出子图);disjoint:若为True,每个种子节点构建独立子图,mini-batch 携带batch向量;时序采样下自动置为True;temporal_strategy:时序采样策略,"uniform"(默认)或"last"(取满足时序约束的最后num_neighbors个邻居);time_attr:节点/边时间戳属性名,设置后保证邻居时间戳不晚于中心节点;weight_attr:边权属性名,设置后按权重偏置采样(权重不必归一化,但须非负、有限且局部邻域内和非零);is_sorted:若edge_index已按列排序(设置time_attr时还要求行内按时间排序)可跳过内部重排序以提升性能;filter_per_worker:过滤发生位置,True在 worker 子进程、False在主进程、None(默认)自动推断(数据部分在 GPU 时为True);directed:旧版参数,已被subgraph_type取代(默认True)。
官方示例(Cora):
from torch_geometric.datasets import Planetoid from torch_geometric.loader import NeighborLoader data = Planetoid(path, name='Cora')[0] loader = NeighborLoader( data, num_neighbors=[30] * 2, # 两跳各采样 30 个邻居 batch_size=128, # 每 batch 128 个训练种子节点 input_nodes=data.train_mask, ) sampled_data = next(iter(loader)) print(sampled_data.batch_size) # >>> 128异构图场景(OGB-MAG),可对每条边类型独立控制采样量:
from torch_geometric.datasets import OGB_MAG from torch_geometric.loader import NeighborLoader hetero_data = OGB_MAG(path)[0] loader = NeighborLoader( hetero_data, num_neighbors={key: [30] * 2 for key in hetero_data.edge_types}, batch_size=128, input_nodes=('paper', hetero_data['paper'].train_mask), ) sampled_hetero_data = next(iter(loader)) print(sampled_hetero_data['paper'].batch_size) # >>> 128返回 mini-batch 的附加属性(源码在 node_loader.py 中写入):
batch_size:种子节点数(batch 中最前面的节点);n_id:每个采样节点对应的全局节点索引(同构图为data.n_id,异构图为data[type].n_id);e_id:每个采样边的全局边索引;input_id:input_nodes的全局索引;num_sampled_nodes/num_sampled_edges:每一跳采样的节点数/边数(异构图下为按类型组织的set_value_dict);- 时序/分布式场景下还会写入
seed_time与_orig_edge_index。
训练要点:默认subgraph_type="directional"仅包含原始采样边,适用于"跳数 = GNN 层数"的情形;若层数多于跳数,应设置"induced"或"bidirectional"以保留采样节点间的更多连接(代价是稍慢)。NodeLoader.filter_fn(node_loader.py)负责把采样结果与特征合并成Data/HeteroData,并兼容FeatureStore/GraphStore远程后端及DistNeighborSampler分布式场景。
3.2HGTLoader:面向异构图的预算均衡采样
HGTLoader实现了 HGT(Heterogeneous Graph Transformer,arXiv:2003.01332)论文中的异构采样策略(torch_geometric/loader/hgt_loader.py),目标有二:让每种节点/边类型保持相近数量,并保持子图稠密以降低采样方差与信息损失。它内部为每种节点类型维护"节点预算",采样概率由节点与已采样节点的连接数及其度决定。
核心参数:
num_samples:每轮每节点类型采样的节点数。传List[int]表示对所有类型使用相同数量,或传Dict[str, List[int]]按类型分别指定;input_nodes:必须传(node_type, indices)元组,None表示该类型全部节点;transform/**kwargs:同NodeLoader。
官方示例:
from torch_geometric.loader import HGTLoader from torch_geometric.datasets import OGB_MAG hetero_data = OGB_MAG(path)[0] loader = HGTLoader( hetero_data, num_samples={key: [512] * 4 for key in hetero_data.node_types}, batch_size=128, input_nodes=('paper', hetero_data['paper'].train_mask), )HGTLoader 同样基于NodeLoader构建,训练范式可参考 examples/hetero/to_hetero_mag.py。
3.3NodeLoader与RandomNodeLoader
NodeLoader本身是通用基类(node_loader.py),接受任意实现了sample_from_nodes的BaseSampler,参数包括node_sampler、input_nodes、input_time、transform、transform_sampler_output、filter_per_worker、custom_cls(远程后端下自定义返回的HeteroData类)等;内部把输入包装为NodeSamplerInput并以range(input_nodes.size(0))作为迭代对象。RandomNodeLoader:在每次迭代中随机采样一批节点构成子图,是节点级随机抽样的轻量选择。
四、链路级采样:LinkNeighborLoader与LinkLoader
4.1LinkNeighborLoader
LinkNeighborLoader是NeighborLoader的链路扩展(torch_geometric/loader/link_neighbor_loader.py):先从edge_label_index中选出一批边,再以这些边两端的节点为种子做邻居采样。它继承了NeighborLoader的num_neighbors、replace、subgraph_type、disjoint、temporal_strategy、time_attr、is_sorted等全部参数,并新增链路相关参数:
edge_label_index:作为采样种子的边索引([2, num_edges]张量),异构图传(edge_type, indices);默认None表示所有边;edge_label:与edge_label_index等长的标签张量,默认None时内部置为torch.zeros(...);edge_label_time:边的时间戳,设置后启用时序约束采样(邻居时间戳早于输出边),需要time_attr;neg_sampling:负采样配置(NegativeSampling对象),详见 4.3;neg_sampling_ratio:已弃用,请改用neg_sampling。
官方示例:
from torch_geometric.datasets import Planetoid from torch_geometric.loader import LinkNeighborLoader data = Planetoid(path, name='Cora')[0] loader = LinkNeighborLoader( data, num_neighbors=[30] * 2, batch_size=128, edge_label_index=data.edge_index, ) sampled_data = next(iter(loader)) # >>> Data(x=[1368, 1433], edge_index=[2, 3103], y=[1368], # train_mask=[1368], val_mask=[1368], test_mask=[1368], # edge_label_index=[2, 128])带标签版本:
loader = LinkNeighborLoader( data, num_neighbors=[30] * 2, batch_size=128, edge_label_index=data.edge_index, edge_label=torch.ones(data.edge_index.size(1)), ) # >>> Data(..., edge_label_index=[2, 128], edge_label=[128])返回 mini-batch 附带的属性与NeighborLoader类似:n_id、e_id、input_id(edge_label_index的全局索引)、num_sampled_nodes、num_sampled_edges。
两个重要注意事项(见 link_neighbor_loader.py):
- 负采样是近似实现,负样本中可能混入假阴性(
false negatives); - 采样过程独立于待预测边——默认情况下
edge_label_index中的监督边不会在采样时被掩蔽。若data.edge_index与edge_label_index存在重叠,可能采到正在预测的边本身。建议通过RandomLinkSplit变换及其disjoint_train_ratio参数(torch_geometric/transforms/random_link_split.py)让两组边不相交。
4.2LinkLoader
LinkLoader(torch_geometric/loader/link_loader.py)是LinkNeighborLoader的通用基类:接受实现了sample_from_edges的BaseSampler,参数包括link_sampler、edge_label_index、edge_label、edge_label_time、neg_sampling、neg_sampling_ratio(弃用)、transform、transform_sampler_output、filter_per_worker、custom_cls等。若需自定义链路采样逻辑,可基于它扩展。
4.3 负采样配置(neg_sampling)
neg_sampling接受NegativeSampling对象,支持两种模式:
"binary"模式:负样本通过返回 mini-batch 对应边类型的edge_label_index与edge_label访问。若原edge_label不存在则自动创建,表示二分类任务(0= 负边,1= 正边);若已存在,则须为0到num_classes-1的分类标签,负采样后0表示负边、1..num_classes表示正边标签。注意:二分类返回torch.float标签(便于直接使用F.binary_cross_entropy),多分类返回torch.long(便于F.cross_entropy);"triplet"模式:通过返回 mini-batch 节点类型的src_index、dst_pos_index、dst_neg_index访问,此时edge_label必须为None。
五、图级采样与分区:ClusterData/ClusterLoader、GraphSAINT、ShaDow k-hop
5.1ClusterData与ClusterLoader(Cluster-GCN)
ClusterData(torch_geometric/loader/cluster.py)基于 METIS 算法把图划分为多个子图分区,对应 Cluster-GCN 论文(arXiv:1905.07953)。参数:
data:图数据对象;num_parts:分区数量;recursive:是否使用多层次递归二分替代多层次 k-way 划分,默认False;save_dir/filename:设置后分区结果会缓存到磁盘(默认文件名metis.pt,目录形如part_{num_parts}{_recursive}/),便于重复使用;log:是否打印分区进度,默认True;keep_inter_cluster_edges:是否保留簇间边连接,默认False;sparse_format:分区计算所用的稀疏格式,"csr"(默认)或"csc"。
注意:底层 METIS 算法要求输入为无向图(cluster.py)。ClusterLoader则负责按分区逐一产出 mini-batch,训练时可配合examples/cluster_gcn_reddit.py、examples/cluster_gcn_ppi.py 等示例使用。从源码看,分区对象Partition(cluster.py)保存了indptr、index、partptr、node_perm、edge_perm与稀疏格式,用于把子图索引映射回原图。
5.2 GraphSAINT 系列采样器
GraphSAINTSampler(torch_geometric/loader/graph_saint.py)是 GraphSAINT 论文(arXiv:1907.04931)采样器的基类,返回的每个 mini-batch 带有归一化系数属性node_norm与edge_norm,用于无偏估计。公共参数:
data:图数据对象;batch_size:每 batch 的近似样本数;num_steps:每个 epoch 的迭代步数,默认1;sample_coverage:用于计算归一化统计量的每节点采样次数,默认0(不计算归一化);save_dir:设置后把归一化统计量缓存到磁盘(文件名形如{sampler_name}_{sample_coverage}.pt);log:是否打印预处理进度,默认True。
三个具体实现(均继承基类并实现_sample_nodes):
GraphSAINTNodeSampler:随机采样节点;GraphSAINTEdgeSampler:随机采样边及其端点;GraphSAINTRandomWalkSampler:随机游走采样。
使用示例见 examples/graph_saint.py,测试覆盖见 test/loader/test_graph_saint.py。注意基类要求data.edge_index存在且位于 CPU,且数据中不能已有node_norm/edge_norm属性(graph_saint.py)。
5.3ShaDowKHopSampler:浅层局部子图
ShaDowKHopSampler(torch_geometric/loader/shadow.py)实现 ShaDow(Decoupling the Depth and Scope of Graph Neural Networks,arXiv:2201.07858)中的 k 跳采样:为每个种子节点构建浅层、局部的 k 跳子图,再由深层 GNN 在这些局部图上平滑信息。参数:
depth:局部子图的跳数;num_neighbors:每一跳每个节点采样的邻居数;node_idx:参与 mini-batch 的节点,默认None(全部节点);replace:是否放回采样,默认False;**kwargs:透传torch.utils.data.DataLoader参数。
注意:该采样器依赖torch-sparse(源码在初始化时检查WITH_TORCH_SPARSE,未安装会抛出ImportError,见 shadow.py)。使用示例见 examples/shadow.py。
六、时序数据加载:TemporalDataLoader
TemporalDataLoader(torch_geometric/loader/temporal_dataloader.py)面向TemporalData(时序事件流)加载数据,把连续的事件合并为 mini-batch:
data:TemporalData对象;batch_size:每 batch 的事件数,默认1;neg_sampling_ratio:负目标节点数相对正目标节点数的比例,默认0.0(即默认不做负采样)。
负采样时,neg_dst通过在[data.dst.min(), data.dst.max()]区间内均匀随机采样生成(temporal_dataloader.py),数量为round(neg_sampling_ratio * batch.dst.size(0))。该类内部以步长为batch_size的range作为迭代序列,shuffle会被强制移除(时序数据须保持时间顺序)。适用于 TGN 等时序图网络训练,示例见 examples/tgn.py。
七、采样器与性能增强工具
7.1NeighborSampler
NeighborSampler是基于torch_sparse的经典邻居采样器(torch_geometric/loader/neighbor_sampler.py),它本身不是一个 DataLoader,而是返回采样结果(n_id、edge_index、e_id)的工具类。NeighborLoader在内部正是通过构造NeighborSampler来完成采样(neighbor_loader.py),并把share_memory=kwargs.get('num_workers', 0) > 0传入以便多进程共享。旧代码中的NeighborSampler用法在新版本中建议迁移到NeighborLoader。相关测试见 test/loader/test_neighbor_sampler.py。
7.2ImbalancedSampler
ImbalancedSampler(torch_geometric/loader/imbalanced_sampler.py)针对节点类别不均衡的数据集,根据节点标签分布计算采样权重,让每个 batch 中各类别保持相对均衡。适用于节点分类中类别分布极度倾斜的场景。
7.3DynamicBatchSampler
DynamicBatchSampler(torch_geometric/loader/dynamic_batch_sampler.py)根据样本的num_nodes动态决定 batch 组成,使每个 batch 的总节点数不超过预设上限(而非固定样本数),适合图规模差异大的数据集以充分利用显存。测试见 test/loader/test_dynamic_batch_sampler.py。
7.4PrefetchLoader与CachedLoader
PrefetchLoader(torch_geometric/loader/prefetch.py)在 GPU 训练时预取下一批数据,与当前 batch 的计算重叠,隐藏数据搬运延迟;参数num_workers控制预取 worker 数,默认2;CachedLoader(torch_geometric/loader/cache.py)缓存 loader 已产出过的 batch 结果,避免重复计算,适合迭代式算法(如 GNN Explainer、标签传播)中多次遍历相同数据。
7.5ZipLoader
ZipLoader(torch_geometric/loader/zip_loader.py)并行迭代多个 loader(如正样本 loader 与负样本 loader),按位置把各 loader 输出打包为元组,供对比学习等需要成对数据的任务使用。测试见 test/loader/test_zip_loader.py。
7.6AffinityMixin
AffinityMixin(torch_geometric/loader/mixin.py)为加载器提供 CPU 亲和性(CPU affinity)设置能力:在 NUMA 架构下将 worker 进程绑定到特定 CPU 核心,减少跨 NUMA 节点的内存访问,提升采样与加载吞吐。NodeLoader/LinkLoader均继承了该 Mixin,可通过其 API 在初始化后启用亲和性优化。
八、如何选择加载器:决策参考
| 你的任务 | 推荐加载器 | 关键参数 |
|---|---|---|
| 中小图集图分类/回归 | DataLoader | batch_size、follow_batch、exclude_keys |
| 大规模同构图节点分类 | NeighborLoader | num_neighbors、input_nodes、subgraph_type |
| 大规模异构图节点分类 | NeighborLoader(按边类型指定)或HGTLoader | num_neighbors/num_samples、input_nodes=('type', idx) |
| 大规模链路预测 | LinkNeighborLoader | edge_label_index、edge_label、neg_sampling |
| 极深 GNN / 图分区训练 | ClusterData+ClusterLoader | num_parts、recursive、save_dir |
| 图级采样(带归一化) | GraphSAINTNodeSampler/EdgeSampler/RandomWalkSampler | batch_size、num_steps、sample_coverage |
| 局部浅层子图 + 深 GNN | ShaDowKHopSampler | depth、num_neighbors |
| 时序事件流(TGN 等) | TemporalDataLoader | batch_size、neg_sampling_ratio |
| 类别不均衡节点分类 | ImbalancedSampler(配合任意节点 loader) | — |
| 图规模差异大的图分类 | DynamicBatchSampler(配合DataLoader) | 动态节点数上限 |
通用提示:大规模采样加载器的公共透传参数(**kwargs)直接进入torch.utils.data.DataLoader,包括batch_size、shuffle、drop_last、num_workers、pin_memory等;num_workers > 0时采样器会自动启用共享内存模式,同时留意filter_per_worker=True在内存数据集上会把全部特征移入共享内存,可能造成文件句柄过多(node_loader.py)。
九、验证与进一步探索
- 模块导出清单:见 torch_geometric/loader/init.py 的
classes列表,与本文第一部分的 25 个类一一对应; - 单元测试:
test/loader/下为每个加载器都配备了测试(如 test_neighbor_loader.py、test_link_neighbor_loader.py、test_hgt_loader.py、test_temporal_dataloader.py),可作为参数语义与边界行为的行为规范参考; - 端到端示例:大规模采样训练可参考 examples/reddit.py、examples/ogbn_train.py、examples/hetero/to_hetero_mag.py;图分区可参考 examples/cluster_gcn_reddit.py 与 examples/cluster_gcn_ppi.py;GraphSAINT 见 examples/graph_saint.py;ShaDow k-hop 见 examples/shadow.py;时序见 examples/tgn.py。
从源码结构看,torch_geometric.loader正在逐步收敛为"通用采样器(sampler模块)+ 通用加载器(NodeLoader/LinkLoader)"的架构,NeighborLoader、LinkNeighborLoader、HGTLoader都是这一架构下的具体实例,因此理解NodeLoader/LinkLoader的参数与filter_fn合并逻辑,是深入掌握整个 loader 模块的关键。
【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考