多彩编程 多彩编程MZPH · CODE BLOG
ARTICLE DETAIL

文章详情

深耕前端与后端开发技术的一线实战笔记与踩坑复盘。

PyTorch Geometric 高级 Mini-Batching 完全指南:对角堆叠原理、DataLoader 机制与自定义 collate 行为

PyTorch Geometric 高级 Mini-Batching 完全指南:对角堆叠原理、DataLoader 机制与自定义 collate 行为 PyTorch Geometric 高级 Mini-Batching 完全指南对角堆叠原理、DataLoader 机制与自定义 collate 行为【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric导读本文围绕 PyGPyTorch Geometric官方文档 Advanced Mini-Batching 展开系统讲解图数据 mini-batch 的底层原理如何把多个图以对角堆叠邻接矩阵的方式合并为一张巨型图以及torch_geometric.loader.DataLoader内部如何通过Data.__inc__与Data.__cat_dim__两个钩子控制合并行为。读完本文你将掌握图级特征、图匹配任务图对、二分图等特殊场景下自定义 batching 的完整方案并能读懂 PyG 数据流水线的核心源码。一、为什么图数据不能像图像、文本那样做 Mini-Batch在深度学习训练中mini-batch 的核心作用是将一批样本组织成统一表示利用并行计算提高吞吐从而支撑大规模数据集的训练。在图像或语言领域这个流程通常通过**缩放或填充padding**实现把每个样本调整到相同形状再沿一个新的维度堆叠起来这个维度的长度就是batch_size。但图是最一般化的数据结构不同图的节点数、边数可以任意不同直接填充会带来两个问题不可行图没有固定的空间结构无法通过简单的 resize 对齐浪费内存即使强行 padding 到相同形状大量填充位置尤其是邻接矩阵中的零元素会造成严重的存储浪费。PyG 选择了另一条完全不同的路不填充、不缩放而是把多个图拼成一张更大的图。二、核心原理邻接矩阵对角堆叠 特征沿节点维拼接假设有 n 个图其邻接矩阵为 $\mathbf{A}_1, \dots, \mathbf{A}_n$节点特征为 $\mathbf{X}_1, \dots, \mathbf{X}_n$标签为 $\mathbf{Y}_1, \dots, \mathbf{Y}_n$则 mini-batch 后的表示为$$ \mathbf{A} \begin{bmatrix} \mathbf{A}_1 \ \ddots \ \mathbf{A}_n \end{bmatrix}, \qquad \mathbf{X} \begin{bmatrix} \mathbf{X}_1 \ \vdots \ \mathbf{X}_n \end{bmatrix}, \qquad \mathbf{Y} \begin{bmatrix} \mathbf{Y}_1 \ \vdots \ \mathbf{Y}_n \end{bmatrix} $$也就是说邻接矩阵沿对角线堆叠节点特征与目标标签直接在节点维度上拼接。拼接后整个 batch 等价于一张包含多个互不连通的子图的巨型图。这种方案有两个关键优势GNN 算子零修改基于消息传递message passing的 GNN 算子无需任何改动即可工作因为对角堆叠天然保证了属于不同子图的节点之间不会交换消息零计算与内存开销整个流程不需要任何 padding邻接矩阵以稀疏形式存储只保存非零元素即边因此对角堆叠不会引入额外的内存开销。三、DataLoader一个覆写 collate 的 PyTorch DataLoaderPyG 通过torch_geometric.loader.DataLoader自动完成多图合并为巨型图的工作。从源码看这个类定义在 torch_geometric/loader/dataloader.pyclass DataLoader(torch.utils.data.DataLoader): def __init__(self, dataset, batch_size1, shuffleFalse, follow_batchNone, exclude_keysNone, **kwargs): kwargs.pop(collate_fn, None) self.follow_batch follow_batch self.exclude_keys exclude_keys super().__init__( dataset, batch_size, shuffle, collate_fnCollater(dataset, follow_batch, exclude_keys), **kwargs, )它本质上是 PyTorchtorch.utils.data.DataLoader的子类唯一的关键改动是覆写了collate_fn即如何把一组样本合并起来的定义将其替换为Collater。因此所有能传给 PyTorchDataLoader的参数如num_workers、pin_memory、prefetch_factor等都能直接传给 PyG 的DataLoader。Collater的核心逻辑torch_geometric/loader/dataloader.py在遇到Data/HeteroData对象时调用Batch.from_data_list完成合并class Collater: def __call__(self, batch): elem batch[0] if isinstance(elem, BaseData): return Batch.from_data_list( batch, follow_batchself.follow_batch, exclude_keysself.exclude_keys, ) ...而Batch.from_data_listtorch_geometric/data/batch.py进一步委托给 torch_geometric/data/collate.py 中的collate函数完成逐属性合并并记录两份辅助字典slice_dict每个属性在合并结果中的切分位置用于从 batch 中还原单个样本对应get_example/to_data_listinc_dict每个属性在合并时被累加的增量用于分离时反向减回原始值。在最一般的形式下合并规则是edge_index形状[2, num_edges]先按前面所有图累计的节点数整体平移增量再沿第 2 维拼接face网格面片索引与edge_index同样处理其他张量直接在第一个维度上拼接数值不做任何增量调整。四、默认的__inc__与__cat_dim__钩子当默认行为不满足需求时PyG 允许用户通过覆写torch_geometric.data.Data的两个方法来自定义 batching 过程__inc__(key, value, *args, **kwargs)定义相邻两个样本的同一属性之间数值需要递增多少increment__cat_dim__(key, value, *args, **kwargs)定义同一属性的张量应沿哪个维度拼接concatenation dimension。文档中给出的默认实现为def __inc__(self, key, value, *args, **kwargs): if index in key: return self.num_nodes else: return 0 def __cat_dim__(self, key, value, *args, **kwargs): if index in key: return 1 else: return 0从当前仓库 torch_geometric/data/data.py 的实际实现看最新版本在此基础上还增加了对稀疏邻接矩阵、face属性以及batch类属性的专门处理def __cat_dim__(self, key, value, *args, **kwargs): if is_sparse(value) and (adj in key or edge_index in key): return (0, 1) # 稀疏张量沿两个维度对角拼接 elif index in key or key face: return -1 # 最后一维等价于文档中的 1 else: return 0 def __inc__(self, key, value, *args, **kwargs): if batch in key and isinstance(value, Tensor): return int(value.max()) 1 # batch 向量按图数量递增 elif index in key or key face: num_nodes self.num_nodes if num_nodes is None: raise RuntimeError(...) # 无法推断 num_nodes 时报错 return num_nodes else: return 0要点解读__inc__默认按num_nodes递增只要属性名包含子串index出于历史原因PyG 就会将其按当前图的节点数平移这对edge_index、node_index等属性非常方便。但注意如果某个属性名恰好包含index却不应递增例如自定义的index类特征就会产生意外行为——最佳实践是始终检查 batching 的输出结果。__inc__需要num_nodes如果数据中没有显式设置num_nodes且无法推断现代版本会直接抛出RuntimeError提醒用户显式设置num_nodes属性例如后文新维度用例中的MyData(num_nodes3, ...)。__cat_dim__决定拼接维度edge_index/face沿第 1 维即最后一维拼接普通张量沿第 0 维拼接。底层调用链collate函数在合并每个属性时会调用get_incstorch_geometric/data/collate.py对每个样本逐一调用data.__inc__(key, value, store)得到增量列表再做前缀和cumsum得到每个样本的实际偏移量拼接维度则由data_list[0].__cat_dim__(key, elem, stores[0])决定。这两个方法属于内部接口PyG 官方建议仅在默认 mini-batch 流程对某个属性失效时才覆写。下面给出三个必须覆写这两个方法的典型场景。五、实战场景一图对Pairs of Graphs与 follow_batch在某些任务如图匹配中需要在单个Data对象里存多个图。例如把源图 $\mathcal{G}_s$ 与目标图 $\mathcal{G}_t$ 放在同一个PairData中from torch_geometric.data import Data class PairData(Data): pass data PairData(x_sx_s, edge_index_sedge_index_s, # Source graph. x_tx_t, edge_index_tedge_index_t) # Target graph.此时默认规则失效edge_index_s必须按源图节点数x_s.size(0)递增edge_index_t必须按目标图节点数x_t.size(0)递增二者互不相同。需要覆写__inc__class PairData(Data): def __inc__(self, key, value, *args, **kwargs): if key edge_index_s: return self.x_s.size(0) if key edge_index_t: return self.x_t.size(0) return super().__inc__(key, value, *args, **kwargs)用两个样本验证from torch_geometric.loader import DataLoader x_s torch.randn(5, 16) # 5 nodes. edge_index_s torch.tensor([ [0, 0, 0, 0], [1, 2, 3, 4], ]) x_t torch.randn(4, 16) # 4 nodes. edge_index_t torch.tensor([ [0, 0, 0], [1, 2, 3], ]) data PairData(x_sx_s, edge_index_sedge_index_s, x_tx_t, edge_index_tedge_index_t) data_list [data, data] loader DataLoader(data_list, batch_size2) batch next(iter(loader)) print(batch) PairDataBatch(x_s[10, 16], edge_index_s[2, 8], x_t[8, 16], edge_index_t[2, 6]) print(batch.edge_index_s) tensor([[0, 0, 0, 0, 5, 5, 5, 5], [1, 2, 3, 4, 6, 7, 8, 9]]) print(batch.edge_index_t) tensor([[0, 0, 0, 4, 4, 4], [1, 2, 3, 5, 6, 7]])可以看到即使源图与目标图的节点数不同5 vs 4edge_index_s与edge_index_t也被正确拼接第二个样本的源边索引整体 5、目标边索引整体 4。不过此时还缺少batch属性用于把每个节点映射到所属图因为 PyG 无法识别PairData中真正的图是谁。这就需要DataLoader的follow_batch参数指定要为哪些属性额外维护 batch 信息。loader DataLoader(data_list, batch_size2, follow_batch[x_s, x_t]) batch next(iter(loader)) print(batch) PairDataBatch(x_s[10, 16], edge_index_s[2, 8], x_s_batch[10], x_t[8, 16], edge_index_t[2, 6], x_t_batch[8]) print(batch.x_s_batch) tensor([0, 0, 0, 0, 0, 1, 1, 1, 1, 1]) print(batch.x_t_batch) tensor([0, 0, 0, 0, 1, 1, 1, 1])follow_batch[x_s, x_t]会为x_s、x_t分别生成分配向量x_s_batch、x_t_batch。从源码看这一步发生在 torch_geometric/data/collate.py当属性名出现在follow_batch中时会基于该属性的slices生成{attr}_batch与{attr}_ptr两个辅助张量。有了这些分配向量就可以对同一个Batch里的多张图执行归约操作例如全局池化 global pooling。六、实战场景二二分图Bipartite Graphs的非对称递增二分图的邻接矩阵定义了两类不同节点之间的关系两类节点数量一般不等因此邻接矩阵是非方阵$\mathbf{A} \in {0, 1}^{N \times M}$其中 $N \neq M$ 是可能的。在二分图的 mini-batch 中edge_index的源节点与目标节点需要独立递增。考虑一个带节点特征x_s、x_t的二分图from torch_geometric.data import Data class BipartiteData(Data): pass data BipartiteData(x_sx_s, x_tx_t, edge_indexedge_index)覆写__inc__让edge_index[0]源节点按x_s.size(0)递增、edge_index[1]目标节点按x_t.size(0)递增class BipartiteData(Data): def __inc__(self, key, value, *args, **kwargs): if key edge_index: return torch.tensor([[self.x_s.size(0)], [self.x_t.size(0)]]) return super().__inc__(key, value, *args, **kwargs)这里的关键是__inc__的返回值可以是一个形状[2, 1]的张量PyG 的get_incs会对这类张量做torch.stack后逐行累加见 torch_geometric/data/collate.py从而实现源、目标两个维度各自独立的递增偏移。验证from torch_geometric.loader import DataLoader x_s torch.randn(2, 16) # 2 nodes. x_t torch.randn(3, 16) # 3 nodes. edge_index torch.tensor([ [0, 0, 1, 1], [0, 1, 1, 2], ]) data BipartiteData(x_sx_s, x_tx_t, edge_indexedge_index) data_list [data, data] loader DataLoader(data_list, batch_size2) batch next(iter(loader)) print(batch) BipartiteDataBatch(x_s[4, 16], x_t[6, 16], edge_index[2, 8]) print(batch.edge_index) tensor([[0, 0, 1, 1, 2, 2, 3, 3], [0, 1, 1, 2, 3, 4, 4, 5]])第二个样本的源节点整体 2x_s.size(0)目标节点整体 3x_t.size(0)完全符合预期。七、实战场景三沿新维度 Batching新增 batch 维度有时我们希望某些图级属性graph-level property / target按经典 mini-batch 的方式新增一个 batch 维度例如把一组形状为[num_features]的属性合并成[num_examples, num_features]而不是默认的[num_examples * num_features]。PyG 的实现方式是在__cat_dim__中返回None表示不拼接而是新增一个维度from torch_geometric.data import Data from torch_geometric.loader import DataLoader class MyData(Data): def __cat_dim__(self, key, value, *args, **kwargs): if key foo: return None return super().__cat_dim__(key, value, *args, **kwargs) edge_index torch.tensor([ [0, 1, 1, 2], [1, 0, 2, 1], ]) foo torch.randn(16) data MyData(num_nodes3, edge_indexedge_index, foofoo) data_list [data, data] loader DataLoader(data_list, batch_size2) batch next(iter(loader)) print(batch) MyDataBatch(num_nodes6, edge_index[2, 8], foo[2, 16])如预期batch.foo变为两维batch 维 特征维。从源码看__cat_dim__返回None时collate会先对每个值执行unsqueeze(0)再沿第 0 维拼接torch_geometric/data/collate.py从而得到[num_examples, num_features]的形状。同时注意示例中显式传入了num_nodes3——由于edge_index的递增需要知道节点数这是保证__inc__正常工作的前提。八、总结与实践建议场景修改方法返回值图对 / 图匹配多图共存于一个 Data覆写__inc__区分各子图的edge_index各自对应的节点数标量二分图源/目标独立递增覆写__inc__对edge_index分别返回形状[2, 1]的张量图级属性需要新增 batch 维覆写__cat_dim__对目标属性返回None需要为某属性生成分配向量使用DataLoader(follow_batch[...])生成{attr}_batch/{attr}_ptr不需要某属性参与 batching使用DataLoader(exclude_keys[...])跳过该顶层属性实践建议总是检查 batching 输出由于默认__inc__对属性名包含index的属性一律按num_nodes递增很容易误伤自定义属性养成打印 batch 形状与数值的习惯非常必要显式设置num_nodes当数据缺少edge_index或节点数无法推断时应在Data中显式传入num_nodes避免现代版本直接抛出RuntimeError理解底层调用链一次合并的本质是DataLoader(collate_fnCollater) → Batch.from_data_list → collate → __cat_dim__ / __inc__其中slice_dict与inc_dict还支撑了get_example、to_data_list等逆向还原操作理解这条链路有助于排查任何 batching 相关的问题利用follow_batch完成图级归约在自定义多图Data中follow_batch生成的分配向量是执行全局池化等归约操作的必备输入。掌握了__inc__与__cat_dim__这两个钩子你就拥有了定制 PyG 数据流水线的最后一公里能力无论是图匹配、二分图推荐还是任意异构的自定义数据结构都能无缝接入标准训练流程。【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考
返回列表