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

资讯详情

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

PyTorch Geometric 数据加载器(torch_geometric.loader)完整 API 指南:从批量训练到大规模图采样

PyTorch Geometric 数据加载器(torch_geometric.loader)完整 API 指南:从批量训练到大规模图采样 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 GeometricPyG中负责把图数据组织成 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、NeighborSamplerGraphSAGE 式采样支持同构图与异构图异构图采样HGTLoader保持各节点类型预算均衡的异构采样图分区/图级采样ClusterData、ClusterLoader、GraphSAINTSampler系列、ShaDowKHopSampler、RandomNodeLoaderCluster-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.1DataLoaderDataLoader继承自torch.utils.data.DataLoadertorch_geometric/loader/dataloader.py其职责是把Dataset中的图对象合并为 mini-batch。构造参数如下dataset数据来源可为Dataset、Sequence[BaseData]或DatasetAdapterbatch_size每个 batch 的样本数默认1shuffle每个 epoch 是否重新打乱数据默认Falsefollow_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 的关键BaseDataData/HeteroData调用Batch.from_data_list(batch, follow_batch..., exclude_keys...)将多个图沿节点维拼接并通过batch向量记录每个节点属于哪个图torch.Tensor走default_collateTensorFrame走torch_frame.cat(batch, dim0)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 训练。这两类加载器都继承自NodeLoadertorch_geometric/loader/node_loader.py后者是承载通用BaseSampler的抽象加载器。3.1NeighborLoaderGraphSAGE 式邻居采样NeighborLoader实现了 GraphSAGE 论文Inductive Representation Learning on Large GraphsarXiv: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_attrreplace是否放回采样默认Falsesubgraph_type返回子图类型取值directional默认仅保留计算种子节点表示所需的有向边、bidirectional转为双向边、induced所有采样节点的导出子图disjoint若为True每个种子节点构建独立子图mini-batch 携带batch向量时序采样下自动置为Truetemporal_strategy时序采样策略uniform默认或last取满足时序约束的最后num_neighbors个邻居time_attr节点/边时间戳属性名设置后保证邻居时间戳不晚于中心节点weight_attr边权属性名设置后按权重偏置采样权重不必归一化但须非负、有限且局部邻域内和非零is_sorted若edge_index已按列排序设置time_attr时还要求行内按时间排序可跳过内部重排序以提升性能filter_per_worker过滤发生位置True在 worker 子进程、False在主进程、None默认自动推断数据部分在 GPU 时为Truedirected旧版参数已被subgraph_type取代默认True。官方示例Corafrom torch_geometric.datasets import Planetoid from torch_geometric.loader import NeighborLoader data Planetoid(path, nameCora)[0] loader NeighborLoader( data, num_neighbors[30] * 2, # 两跳各采样 30 个邻居 batch_size128, # 每 batch 128 个训练种子节点 input_nodesdata.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_size128, 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_ide_id每个采样边的全局边索引input_idinput_nodes的全局索引num_sampled_nodes/num_sampled_edges每一跳采样的节点数/边数异构图下为按类型组织的set_value_dict时序/分布式场景下还会写入seed_time与_orig_edge_index。训练要点默认subgraph_typedirectional仅包含原始采样边适用于跳数 GNN 层数的情形若层数多于跳数应设置induced或bidirectional以保留采样节点间的更多连接代价是稍慢。NodeLoader.filter_fnnode_loader.py负责把采样结果与特征合并成Data/HeteroData并兼容FeatureStore/GraphStore远程后端及DistNeighborSampler分布式场景。3.2HGTLoader面向异构图的预算均衡采样HGTLoader实现了 HGTHeterogeneous Graph TransformerarXiv: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_size128, input_nodes(paper, hetero_data[paper].train_mask), )HGTLoader 同样基于NodeLoader构建训练范式可参考 examples/hetero/to_hetero_mag.py。3.3NodeLoader与RandomNodeLoaderNodeLoader本身是通用基类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与LinkLoader4.1LinkNeighborLoaderLinkNeighborLoader是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_attrneg_sampling负采样配置NegativeSampling对象详见 4.3neg_sampling_ratio已弃用请改用neg_sampling。官方示例from torch_geometric.datasets import Planetoid from torch_geometric.loader import LinkNeighborLoader data Planetoid(path, nameCora)[0] loader LinkNeighborLoader( data, num_neighbors[30] * 2, batch_size128, edge_label_indexdata.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_size128, edge_label_indexdata.edge_index, edge_labeltorch.ones(data.edge_index.size(1)), ) # Data(..., edge_label_index[2, 128], edge_label[128])返回 mini-batch 附带的属性与NeighborLoader类似n_id、e_id、input_idedge_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.2LinkLoaderLinkLoadertorch_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_samplingneg_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_entropytriplet模式通过返回 mini-batch 节点类型的src_index、dst_pos_index、dst_neg_index访问此时edge_label必须为None。五、图级采样与分区ClusterData/ClusterLoader、GraphSAINT、ShaDow k-hop5.1ClusterData与ClusterLoaderCluster-GCNClusterDatatorch_geometric/loader/cluster.py基于 METIS 算法把图划分为多个子图分区对应 Cluster-GCN 论文arXiv:1905.07953。参数data图数据对象num_parts分区数量recursive是否使用多层次递归二分替代多层次 k-way 划分默认Falsesave_dir/filename设置后分区结果会缓存到磁盘默认文件名metis.pt目录形如part_{num_parts}{_recursive}/便于重复使用log是否打印分区进度默认Truekeep_inter_cluster_edges是否保留簇间边连接默认Falsesparse_format分区计算所用的稀疏格式csr默认或csc。注意底层 METIS 算法要求输入为无向图cluster.py。ClusterLoader则负责按分区逐一产出 mini-batch训练时可配合examples/cluster_gcn_reddit.py、examples/cluster_gcn_ppi.py 等示例使用。从源码看分区对象Partitioncluster.py保存了indptr、index、partptr、node_perm、edge_perm与稀疏格式用于把子图索引映射回原图。5.2 GraphSAINT 系列采样器GraphSAINTSamplertorch_geometric/loader/graph_saint.py是 GraphSAINT 论文arXiv:1907.04931采样器的基类返回的每个 mini-batch 带有归一化系数属性node_norm与edge_norm用于无偏估计。公共参数data图数据对象batch_size每 batch 的近似样本数num_steps每个 epoch 的迭代步数默认1sample_coverage用于计算归一化统计量的每节点采样次数默认0不计算归一化save_dir设置后把归一化统计量缓存到磁盘文件名形如{sampler_name}_{sample_coverage}.ptlog是否打印预处理进度默认True。三个具体实现均继承基类并实现_sample_nodesGraphSAINTNodeSampler随机采样节点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浅层局部子图ShaDowKHopSamplertorch_geometric/loader/shadow.py实现 ShaDowDecoupling the Depth and Scope of Graph Neural NetworksarXiv: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。六、时序数据加载TemporalDataLoaderTemporalDataLoadertorch_geometric/loader/temporal_dataloader.py面向TemporalData时序事件流加载数据把连续的事件合并为 mini-batchdataTemporalData对象batch_size每 batch 的事件数默认1neg_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.1NeighborSamplerNeighborSampler是基于torch_sparse的经典邻居采样器torch_geometric/loader/neighbor_sampler.py它本身不是一个 DataLoader而是返回采样结果n_id、edge_index、e_id的工具类。NeighborLoader在内部正是通过构造NeighborSampler来完成采样neighbor_loader.py并把share_memorykwargs.get(num_workers, 0) 0传入以便多进程共享。旧代码中的NeighborSampler用法在新版本中建议迁移到NeighborLoader。相关测试见 test/loader/test_neighbor_sampler.py。7.2ImbalancedSamplerImbalancedSamplertorch_geometric/loader/imbalanced_sampler.py针对节点类别不均衡的数据集根据节点标签分布计算采样权重让每个 batch 中各类别保持相对均衡。适用于节点分类中类别分布极度倾斜的场景。7.3DynamicBatchSamplerDynamicBatchSamplertorch_geometric/loader/dynamic_batch_sampler.py根据样本的num_nodes动态决定 batch 组成使每个 batch 的总节点数不超过预设上限而非固定样本数适合图规模差异大的数据集以充分利用显存。测试见 test/loader/test_dynamic_batch_sampler.py。7.4PrefetchLoader与CachedLoaderPrefetchLoadertorch_geometric/loader/prefetch.py在 GPU 训练时预取下一批数据与当前 batch 的计算重叠隐藏数据搬运延迟参数num_workers控制预取 worker 数默认2CachedLoadertorch_geometric/loader/cache.py缓存 loader 已产出过的 batch 结果避免重复计算适合迭代式算法如 GNN Explainer、标签传播中多次遍历相同数据。7.5ZipLoaderZipLoadertorch_geometric/loader/zip_loader.py并行迭代多个 loader如正样本 loader 与负样本 loader按位置把各 loader 输出打包为元组供对比学习等需要成对数据的任务使用。测试见 test/loader/test_zip_loader.py。7.6AffinityMixinAffinityMixintorch_geometric/loader/mixin.py为加载器提供 CPU 亲和性CPU affinity设置能力在 NUMA 架构下将 worker 进程绑定到特定 CPU 核心减少跨 NUMA 节点的内存访问提升采样与加载吞吐。NodeLoader/LinkLoader均继承了该 Mixin可通过其 API 在初始化后启用亲和性优化。八、如何选择加载器决策参考你的任务推荐加载器关键参数中小图集图分类/回归DataLoaderbatch_size、follow_batch、exclude_keys大规模同构图节点分类NeighborLoadernum_neighbors、input_nodes、subgraph_type大规模异构图节点分类NeighborLoader按边类型指定或HGTLoadernum_neighbors/num_samples、input_nodes(type, idx)大规模链路预测LinkNeighborLoaderedge_label_index、edge_label、neg_sampling极深 GNN / 图分区训练ClusterDataClusterLoadernum_parts、recursive、save_dir图级采样带归一化GraphSAINTNodeSampler/EdgeSampler/RandomWalkSamplerbatch_size、num_steps、sample_coverage局部浅层子图 深 GNNShaDowKHopSamplerdepth、num_neighbors时序事件流TGN 等TemporalDataLoaderbatch_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_workerTrue在内存数据集上会把全部特征移入共享内存可能造成文件句柄过多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.pyGraphSAINT 见 examples/graph_saint.pyShaDow 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),仅供参考
返回列表