神经网络通常以批量方式训练。PyG 通过创建稀疏块对角邻接矩阵(由 edge_index 定义)并在节点维度上连接特征矩阵和目标矩阵来实现对 mini-batch 的并行化。这种组合允许一个 batch 中的不同示例具有不同数量的节点和边:

PyG 自身包含 torch_geometric.loader.DataLoader ,它已经负责了这个连接过程。
from torch_geometric.datasets import TUDataset
from torch_geometric.loader import DataLoaderdataset = TUDataset(root='/tmp/ENZYMES', name='ENZYMES', use_node_attr=True)
loader = DataLoader(dataset, batch_size=32, shuffle=True)for batch in loader:print(batch)
"""
DataBatch(edge_index=[2, 3590], x=[988, 21], y=[32], batch=[988], ptr=[33])
DataBatch(edge_index=[2, 4260], x=[1155, 21], y=[32], batch=[1155], ptr=[33])
DataBatch(edge_index=[2, 3718], x=[968, 21], y=[32], batch=[968], ptr=[33])
DataBatch(edge_index=[2, 3506], x=[1017, 21], y=[32], batch=[1017], ptr=[33])
DataBatch(edge_index=[2, 4182], x=[1100, 21], y=[32], batch=[1100], ptr=[33])
DataBatch(edge_index=[2, 3890], x=[1021, 21], y=[32], batch=[1021], ptr=[33])
"""
torch_geometric.data.Batch 继承自 torch_geometric.data.Data ,并包含一个额外的属性称为 batch 。
batch 是一个列向量,将每个节点映射到批次中的相应图:

ptr 表示每个图在批处理数据中起始节点的索引,如果 ptr = [0, 12, 30, 45],就代表:第 1 个图的节点索引为 [0, 12),即 0 到 11,第 2 个图为 [12, 30),第 3 个图为 [30, 45)。