PyG 远程后端(Remote Backends)完全指南:借助 FeatureStore 与 GraphStore 将 GNN 扩展到单机内存之外
【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric
本指南深入解析 PyTorch Geometric(PyG,2.2 及更高版本)提供的可扩展图机器学习基础设施:通过
FeatureStore(特征存储)与GraphStore(图结构存储)两个抽象接口,将节点特征与图结构迁移到远程/外置存储,配合采样器(Sampler)与数据加载器(NodeLoader/LinkLoader),让用户能够在远超单机内存容量的超大图上训练 GNN。读完本文,你将掌握远程后端的背景动机、两大存储抽象的接口设计与 CRUD 用法、采样器与加载器的协同机制,以及如何在当前仓库中动手实现自己的远程后端。
引言:为什么需要远程后端
一个实例化的图神经网络,其数据本质由两部分组成:
- 节点/边特征信息:图中节点与边对应的稠密向量(attribute);
- 图结构信息:图中的节点以及连接这些节点的边。
观察 GNN 的训练方式可以立刻得到一个结论:当数据规模超出所选加速器(accelerator,如 GPU)的可用内存时,就必须放弃全图训练(full-batch training),改为在采样子图上训练(mini-batch training)。虽然这种方式给学习过程引入了随机性,但它把加速器的内存需求降低到了采样子图的规模,这正是经典 mini-batch GNN 训练范式的核心思想(见图 1)。
图 1:经典 mini-batch GNN 训练范式:全图与特征整体存放在 CPU DRAM,每轮迭代将采样子图及对应特征送入加速器。
然而,mini-batch 训练并非解决所有图学习扩展性问题的银弹。由于每一轮学习都要采样子图再送入加速器,传统做法要求图和特征常驻在用户机器的 CPU DRAM 中。在大规模场景下,这个要求会变得相当沉重:
- 采购一台拥有足够 CPU DRAM 来容纳整张图和特征的机器非常困难;
- 采用数据并行训练时,需要把图和特征复制到每一个计算节点;
- 图和特征很容易超过单台机器的内存上限。
因此,要扩展到超出单机内存的极大图和特征,就必须把这些数据结构移出内存(out-of-core),只让执行计算的那个节点处理采样子图。为了达成这一目标,PyG 引入了两个核心抽象来分别存储特征信息与图结构(见图 2):
- 特征存放在一个键值式的
FeatureStore中,它必须支持高效的随机访问; - 图信息存放在一个
GraphStore中,它必须支持采样器对其进行高效采样。
图 2:图数据在远程存储与训练实例之间的布局。左侧分布式存储中,Graph Store 负责节点与边结构,Feature Store 负责节点/边张量;右侧训练实例内,采样子图与特征拼接后即为前向/反向传播所需的全部数据。
在 PyG(2.2 及更高版本)中,图数据被拆分为特征与结构两部分、这些信息被存放在可能远离实际训练节点的地方,以及它们之间的交互——这一切对最终用户是完全透明的。只要FeatureStore与GraphStore被恰当地定义(并牢记上文提到的性能要求),剩下的工作全部由 PyG 处理。
⚠️ 注意事项
- 本文讨论的远程后端 API 仍处于演进之中,PyG 团队会持续改进其易用性与通用性,未来可能发生变化。
- 目前
FeatureStore与GraphStore仅支持异质图(heterogeneous graphs),且不支持边特征;同质图与边特征的支持即将到来。这一点在仓库源码中亦有印证:DataType.from_data(见 sampler/base.py)将(FeatureStore, GraphStore)元组识别为remote类型,而 NodeLoader/LinkLoader 的返回类型为HeteroData。
FeatureStore:特征的键值抽象
FeatureStore持有图中节点与边的特征。特征存储通常是图学习应用中最主要的存储瓶颈,因为图布局信息(即edge_index)本身相对廉价(每条边约 32 字节)。
PyG 为各种FeatureStore实现提供了一个公共接口,使其能够接入核心学习 API。其抽象基类定义在 feature_store.py,实现细节通过一套CRUD 风格的接口与 PyG 解耦。实现者主要需要覆写三个方法:
| 方法 | 作用 | 源码位置 |
|---|---|---|
put_tensor(tensor, *args, **kwargs) | 向存储中同步写入一个特征张量,返回是否成功 | feature_store.py |
get_tensor(*args, **kwargs) | 从存储中同步读取一个特征张量 | feature_store.py |
remove_tensor(*args, **kwargs) | 从存储中删除一个特征张量,返回是否成功 | feature_store.py |
这三个方法都以TensorAttr(见 feature_store.py)为参数标识。TensorAttr包含三个字段,其顺序即索引调用时属性必须给出的顺序:
group_name:张量所属的分组名(例如异质图中的节点类型);attr_name:张量在分组内的名字(例如x或edge_attr);index:张量行对应的节点索引(可为torch.Tensor、numpy.ndarray、slice或单个整数)。
底层方法(_put_tensor/_get_tensor/_remove_tensor)是抽象方法,由子类实现;公开方法则在调用前通过TensorAttr.cast完成属性解析,并要求属性必须被完整指定(fully specified),否则抛出ValueError。除 CRUD 外,接口还提供了multi_get_tensor(批量读取,默认实现逐条调用get_tensor,实现类可覆写以获得更高性能)、get_tensor_size、get_all_tensor_attrs以及update_tensor(默认先删后插)等辅助方法。
这一设计同时赋予用户pythonic 的接口来检查和修改FeatureStore中的元素。以下是文档给出的完整示例:
feature_store = CustomFeatureStore() paper_features = ... # [num_papers, num_paper_features] author_features = ... # [num_authors, num_author_features] # 写入特征: feature_store['paper', 'x', None] = paper_features feature_store['author', 'x', None] = author_features # 访问特征: assert torch.equal(feature_store['paper', 'x'], paper_features) assert torch.equal(feature_store['paper'].x, paper_features) assert torch.equal(feature_store['author', 'x', 0:20], author_features[0:20])上述索引语法由FeatureStore.__getitem__/__setitem__(见 feature_store.py)与AttrView(见 feature_store.py)共同实现:完全指定的键会直接产出张量;部分指定的键会返回一个AttrView视图,该视图可继续按属性名或索引取值,也可通过调用store[group, attr]()强制触发 GET 操作。例如feature_store['paper'].x正是先得到'paper'的视图、再以属性访问方式补全attr_name的链式写法。
从设计意图看,FeatureStore抽象做出如下关键假设(见 feature_store.py):特征可通过TensorAttr中指定的任意属性唯一标识;实现者负责妥善处理这些假设——例如一个简单的内存实现可以把所有元数据值与特征索引拼接,作为键值存储中的唯一键;更复杂的实现可以基于元数据对特征做有趣的分区。常见的FeatureStore实现形态是键值存储(key-value store),例如memcached、LevelDB、RocksDB都是可行的性能选项。源码中还标注了未来的重要 TODO:异步put与get功能。
GraphStore:面向高效采样的图结构抽象
GraphStore持有定义节点间关系的边索引。其目标是以支持从根节点高效采样的方式存储图信息,采样算法由开发者自行选择。
与FeatureStore类似,PyG 为各种GraphStore实现提供了接入核心学习 API 的公共接口;但与FeatureStore不同的是,GraphStore不需要对全部元素提供随机访问,而需要定义一种能提供高效子图采样的表示。其抽象基类定义在 graph_store.py,核心 CRUD 方法包括:
| 方法 | 作用 | 源码位置 |
|---|---|---|
put_edge_index(edge_index, *args, **kwargs) | 以EdgeAttr指定的格式写入边索引 | graph_store.py |
get_edge_index(*args, **kwargs) | 读取边索引;找不到时抛出KeyError | graph_store.py |
remove_edge_index(*args, **kwargs) | 删除边索引 | graph_store.py |
边索引通过EdgeAttr(见 graph_store.py)唯一标识,其字段包括:
edge_type:边类型(在 PyG 中为源节点、关系类型、目标节点的三元组);layout:边表示格式,取值为EdgeLayout枚举中的COO、CSC或CSR(见 graph_store.py);is_sorted:边索引是否按目标节点排序(对 COO 有意义,CSC 天然有序,CSR 定义上即无序);size:该边类型的源/目标节点数量。
GraphStore抽象的关键假设是:边索引仅以 COO、CSC 或 CSR 格式表示,且一旦存入即静态不变(不支持动态修改),这一点在源码 docstring 中有明确说明。接口还内置了布局转换能力:coo()/csr()/csc()方法(见 graph_store.py)可在三种格式间转换,内部通过_edge_to_layout/_edges_to_layout完成,转换时可选择是否将结果回写存储(store=True)。测试用例 test_graph_store.py 覆盖了基本的读写与 COO↔CSR↔CSC 转换逻辑。
接口用法示例如下:
graph_store = CustomGraphStore() edge_index = torch.tensor([[0, 1, 1, 2], [1, 0, 2, 1]]) # 写入边: graph_store['edge', 'coo'] = coo # 访问边: row, col = graph_store['edge', 'coo'] assert torch.equal(row, edge_index[0]) assert torch.equal(col, edge_index[1])常见的GraphStore实现是图数据库(graph database),例如Neo4j、TigerGraph、ArangoDB、Kùzu都是可行的性能选项。当前仓库提供了一个与Kùzu图数据库结合的示例(见 examples/distributed/kuzu):Kùzu 是一个为查询速度与可扩展性而构建的进程内属性图数据库,其 Python API 直接输出可接入 PyG 接口的FeatureStore与GraphStore,从而允许直接在存储在 Kùzu 中的图上训练 GNN;示例包含papers_100M场景(约 1.11 亿节点、16 亿边的ogbn-papers100M数据在单机上配合远程后端使用)。
采样器:与 GraphStore 紧密耦合
图采样器(graph sampler)与给定的GraphStore紧密耦合,它操作GraphStore从输入节点出发产出采样子图。不同采样算法实现在BaseSampler接口(见 sampler/base.py)背后:
- 默认情况下,PyG 的默认内存采样器会把所有边索引从
GraphStore拉取到训练节点内存中,转换为压缩稀疏列(CSC)格式,然后复用预构建的内存采样例程; - 自定义采样器实现则可以选择覆写
BaseSampler.sample_from_nodes(见 sampler/base.py)和/或BaseSampler.sample_from_edges(见 sampler/base.py),调用GraphStore的专有方法以获得效率提升(例如直接在远程GraphStore上执行采样)。
sample_from_nodes接收一个NodeSamplerInput(包含输入示例索引、种子节点索引、可选时间戳与输入节点类型),返回SamplerOutput或HeteroSamplerOutput;sample_from_edges则接收EdgeSamplerInput并支持可选的neg_sampling配置(用于负采样)。此外,BaseSampler还提供edge_permutation属性(见 sampler/base.py)报告采样过程中对边顺序的置换,供数据加载器还原原始边 ID。源码还给出重要提示:采样器中存放的任何数据都会在数据加载 worker 间被复制(每个 worker 持有采样器的独立实例),因此建议限制采样器内保存的信息量。
# `CustomGraphSampler` 知道如何在 `CustomGraphStore` 上采样: node_sampler = CustomGraphSampler( graph_store=graph_store, num_neighbors=[10, 20], ... )数据加载器:NodeLoader 与 LinkLoader
PyG 并未为GraphStore定义必须实现的采样领域专用语言(DSL);相反,采样器与GraphStore通过数据加载器紧密耦合在一起。
PyG 开箱即用地提供了两种数据加载器:
NodeLoader(见 node_loader.py):从输入节点采样子图,用于节点分类任务;LinkLoader:从一条边的任一侧采样子图,用于链接预测任务。
这两种加载器都以FeatureStore、GraphStore和一个图采样器为输入,内部调用采样器的sample_from_nodes或sample_from_edges方法执行子图采样。NodeLoader的__init__签名(见 node_loader.py)支持data(Data、HeteroData或(FeatureStore, GraphStore)元组)、node_sampler、input_nodes(异质图中需以(node_type, indices)元组形式传入)、input_time、transform、transform_sampler_output、filter_per_worker(自动推断过滤发生在 worker 子进程还是主进程)、custom_cls(远程后端场景下返回的自定义HeteroData类)等参数,其余**kwargs直接透传给torch.utils.data.DataLoader(如batch_size、shuffle、drop_last、num_workers)。
核心用法如下:
# 不再传入 PyG data 对象,而是传入 # `FeatureStore` 与 `GraphStore` 组成的元组作为输入数据: loader = NodeLoader( data=(feature_store, graph_store), node_sampler=node_sampler, batch_size=20, input_nodes='paper', ) for batch in loader: pass在内部,NodeLoader.collate_fn会调用node_sampler.sample_from_nodes完成采样;随后filter_fn(见 node_loader.py)负责把采样结果与特征拼接:对远程后端(Tuple[FeatureStore, GraphStore]),调用filter_custom_store/filter_custom_hetero_store从特征存储中取出采样节点对应的特征,构造出Data或HeteroData对象(并可附上n_id、e_id、batch、num_sampled_nodes、num_sampled_edges、input_id等元数据),最终送到加速器。加载器还支持分布式场景:当采样器为DistNeighborSampler时走专门的分布式过滤路径。
整体架构:组件如何协同工作
从高层次看,上述组件共同协作,为 PyG 内扩展 GNN 训练提供支撑(见图 3):
- 数据加载器(准确地说,是每个 worker)借助一个
BaseSampler向GraphStore发起采样请求; - 收到响应后,数据加载器随后向
FeatureStore查询采样子图中节点与边对应的特征; - 数据加载器从图结构与特征信息中构造最终的 mini-batch,发送给加速器执行前向/反向传播;
- 循环往复,直至收敛。
图 3:统一FeatureStore、GraphStore、图采样器与数据加载器的公共接口与数据流:(1) 根节点 → (2) 采样节点 → (3) 采样节点 → (4) 采样节点特征。
上述所有类都通过公共接口通信,因此它们是可扩展、可泛化的,并且易于与用户日常使用的 PyG 集成——DataType.from_data(见 sampler/base.py)将(FeatureStore, GraphStore)元组判定为remote数据类型的逻辑,正是这一"即插即用"设计的直接体现。
动手实践:从零实现远程后端
要开始使用这一扩展能力,推荐按以下步骤进行:
- 阅读接口:通读
FeatureStore(feature_store.py)与GraphStore(graph_store.py)的抽象基类定义,理解TensorAttr/EdgeAttr的字段语义与 CRUD 方法契约; - 实现
FeatureStore:覆写_put_tensor/_get_tensor/_remove_tensor,保证对随机访问的高效支持(键值存储是首选形态); - 实现
GraphStore:覆写_put_edge_index/_get_edge_index/_remove_edge_index,保证对子图采样的高效支持(图数据库是首选形态); - 实现采样器:根据采样算法覆写
BaseSampler.sample_from_nodes与/或sample_from_edges,尽量把采样逻辑下沉到远程GraphStore侧以省去整图搬运; - 接入加载器:将三者作为参数传给
NodeLoader或LinkLoader,其余 PyG 功能将像纯内存应用一样无缝工作。
一旦FeatureStore、GraphStore和BaseSampler实现正确,只需把它们作为参数传递给NodeLoader或LinkLoader,PyG 的其余部分便会无缝运行,与任何纯内存应用别无二致。
可参考的仓库样例包括:
- examples/distributed/kuzu/README.md:Kùzu 远程后端示例说明,覆盖 PubMed 与
papers_100M(约 1.11 亿节点 / 16 亿边)两种规模; - examples/distributed/kuzu/papers_100M:
ogbn-papers100M大规模图在单机上的远程后端训练代码; - torch_geometric/data/feature_store.py 与 torch_geometric/data/graph_store.py:两个抽象基类的完整定义与接口契约;
- torch_geometric/sampler/base.py:
BaseSampler、NodeSamplerInput、EdgeSamplerInput与输出类型定义; - torch_geometric/loader/node_loader.py 与 torch_geometric/loader/link_loader.py:两种开箱即用的远程后端数据加载器;
- test/data/test_graph_store.py:
GraphStore读写与格式转换的测试用例,可作为实现正确性的验证参考。
结语
PyG 的远程后端通过FeatureStore与GraphStore两个简洁、易用且可扩展的抽象,将"特征"与"结构"的存储彻底解耦,并以数据加载器为纽带与采样器紧密协作,从而把可扩展 GNN 训练的复杂度从用户侧完全抽离。值得注意的是,该特性仍处于密集开发阶段:目前仅支持异质图、暂不支持边特征,API 细节未来仍可能调整。如果你在使用过程中有任何问题、意见或顾虑,可以前往 PyG 的 GitHub Discussions 或 Slack 与 PyG 核心团队交流,共同推动这一方向的演进。
【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考