news 2026/9/13 17:31:58

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

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
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

导读

本文围绕 PyG(PyTorch Geometric)官方文档 Advanced Mini-Batching 展开,系统讲解图数据 mini-batch 的底层原理:如何把多个图以"对角堆叠邻接矩阵"的方式合并为一张巨型图,以及torch_geometric.loader.DataLoader内部如何通过Data.__inc__Data.__cat_dim__两个钩子控制合并行为。读完本文,你将掌握图级特征、图匹配任务(图对)、二分图等特殊场景下自定义 batching 的完整方案,并能读懂 PyG 数据流水线的核心源码。


一、为什么图数据不能像图像、文本那样做 Mini-Batch

在深度学习训练中,mini-batch 的核心作用是:将一批样本组织成统一表示,利用并行计算提高吞吐,从而支撑大规模数据集的训练。在图像或语言领域,这个流程通常通过**缩放或填充(padding)**实现:把每个样本调整到相同形状,再沿一个新的维度堆叠起来,这个维度的长度就是batch_size

但图是最一般化的数据结构,不同图的节点数、边数可以任意不同,直接填充会带来两个问题:

  1. 不可行:图没有固定的空间结构,无法通过简单的 resize 对齐;
  2. 浪费内存:即使强行 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 等价于"一张包含多个互不连通的子图的巨型图"。

这种方案有两个关键优势:

  1. GNN 算子零修改:基于消息传递(message passing)的 GNN 算子无需任何改动即可工作,因为对角堆叠天然保证了属于不同子图的节点之间不会交换消息;
  2. 零计算与内存开销:整个流程不需要任何 padding;邻接矩阵以稀疏形式存储,只保存非零元素(即边),因此对角堆叠不会引入额外的内存开销。

三、DataLoader:一个覆写 collate 的 PyTorch DataLoader

PyG 通过torch_geometric.loader.DataLoader自动完成"多图合并为巨型图"的工作。从源码看,这个类定义在 torch_geometric/loader/dataloader.py:

class DataLoader(torch.utils.data.DataLoader): def __init__(self, dataset, batch_size=1, shuffle=False, follow_batch=None, exclude_keys=None, **kwargs): kwargs.pop('collate_fn', None) self.follow_batch = follow_batch self.exclude_keys = exclude_keys super().__init__( dataset, batch_size, shuffle, collate_fn=Collater(dataset, follow_batch, exclude_keys), **kwargs, )

它本质上是 PyTorchtorch.utils.data.DataLoader的子类,唯一的关键改动是覆写了collate_fn(即"如何把一组样本合并起来"的定义),将其替换为Collater。因此,所有能传给 PyTorchDataLoader的参数(如num_workerspin_memoryprefetch_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_batch=self.follow_batch, exclude_keys=self.exclude_keys, ) ...

Batch.from_data_list(torch_geometric/data/batch.py)进一步委托给 torch_geometric/data/collate.py 中的collate函数完成逐属性合并,并记录两份辅助字典:

  • slice_dict:每个属性在合并结果中的切分位置,用于从 batch 中还原单个样本(对应get_example/to_data_list);
  • inc_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

要点解读:

  1. __inc__默认按num_nodes递增:只要属性名包含子串index(出于历史原因),PyG 就会将其按当前图的节点数平移,这对edge_indexnode_index等属性非常方便。但注意:如果某个属性名恰好包含index却不应递增(例如自定义的index类特征),就会产生意外行为——最佳实践是始终检查 batching 的输出结果
  2. __inc__需要num_nodes:如果数据中没有显式设置num_nodes且无法推断,现代版本会直接抛出RuntimeError,提醒用户显式设置num_nodes属性(例如后文新维度用例中的MyData(num_nodes=3, ...))。
  3. __cat_dim__决定拼接维度edge_index/face沿第 1 维(即最后一维)拼接,普通张量沿第 0 维拼接。
  4. 底层调用链:collate函数在合并每个属性时,会调用get_incs(torch_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_s=x_s, edge_index_s=edge_index_s, # Source graph. x_t=x_t, edge_index_t=edge_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_s=x_s, edge_index_s=edge_index_s, x_t=x_t, edge_index_t=edge_index_t) data_list = [data, data] loader = DataLoader(data_list, batch_size=2) 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 4),edge_index_sedge_index_t也被正确拼接:第二个样本的源边索引整体 +5、目标边索引整体 +4。

不过此时还缺少batch属性(用于把每个节点映射到所属图),因为 PyG 无法识别PairData中真正的"图"是谁。这就需要DataLoaderfollow_batch参数:指定要为哪些属性额外维护 batch 信息。

loader = DataLoader(data_list, batch_size=2, 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_sx_t分别生成分配向量x_s_batchx_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_sx_t的二分图:

from torch_geometric.data import Data class BipartiteData(Data): pass data = BipartiteData(x_s=x_s, x_t=x_t, edge_index=edge_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_s=x_s, x_t=x_t, edge_index=edge_index) data_list = [data, data] loader = DataLoader(data_list, batch_size=2) 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]])

第二个样本的源节点整体 +2(x_s.size(0)),目标节点整体 +3(x_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_nodes=3, edge_index=edge_index, foo=foo) data_list = [data, data] loader = DataLoader(data_list, batch_size=2) batch = next(iter(loader)) print(batch) >>> MyDataBatch(num_nodes=6, 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_nodes=3——由于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_fn=Collater) → Batch.from_data_list → collate → __cat_dim__ / __inc__,其中slice_dictinc_dict还支撑了get_exampleto_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),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/13 17:31:55

贝叶斯最小错误率分类器:从决策规则到手写数字识别实战

简介:这套基于贝叶斯最小错误率的手写数字识别项目,面向机器学习初学者与模式识别课程设计者,解决手写数字高效分类问题。项目在贝叶斯决策框架下,不仅计算后验概率,还引入错误成本,使不同误判损失能得到合…

作者头像 李华
网站建设 2026/9/13 17:31:49

两轮车地下定位技术:GPS失效后的四大替代方案

1. 为什么两轮车在隧道和地下车库会“失联”——不是设备坏了,是物理定律在发号施令你骑着电动自行车刚进地铁站口的斜坡,导航语音突然卡住:“前方…呃…请…保持…直行…”接着彻底静音;或者推着共享单车穿过商场B2层车库&#x…

作者头像 李华
网站建设 2026/9/13 17:30:15

IDEA打jar包全攻略:从普通Java到Spring Boot及外部依赖处理

干Java这行的,几乎没人能绕开“打jar包”这三个字。不管是把自己写的工具类发给同事,还是把一个Spring Boot服务部署到Windows服务器上,最后一步基本都得落到“怎么打出一个能跑的jar包”上。但我发现一个很有意思的现象:同样问“…

作者头像 李华
网站建设 2026/9/13 17:29:23

嵌入式开发三大方向:单片机、Linux驱动与汽车电子如何选择

1. 这不是危言耸听:为什么嵌入式入门前必须厘清这三个方向?“搞不懂这三个方向,千万别碰嵌入式!”——这句话在嵌入式圈子流传多年,不是导师吓唬新人,而是无数人踩坑后用项目延期、芯片烧毁、驱动崩溃换来的…

作者头像 李华