PyTorch Geometric中如何自定义HeteroData的mini-batch
问题背景
我在参考PyTorch Geometric官方高级mini-batch批处理教程做开发,需要实现自定义批处理逻辑:任务要在图结构中结合LSTM,用6个图预测单张图内4个节点对应的6个预测值,需要构造包含6个互不连通子图的图数据。同构图场景下用普通Data类很容易实现,但当前处理的是异构图场景:待预测的4个节点、提供辅助信息的11个节点分属不同类型,节点特征维度不同,节点间的边类型也存在差异。
目前官方没有提供异构图场景下该类自定义批处理的实现指南,同构图场景下已经实现的自定义SeqData类代码如下:
from torch_geometric.data import Data class SeqData(Data): def __init__(self, edge_indices=None, xs=None, y=None): super().__init__() self.edge_indices = edge_indices self.xs = xs self.y = y def __inc__(self, key, value, *args, **kwargs): if key == "edge_indices": return self.xs[0].size(0) else: return super().__inc__(key, value, *args, **kwargs) n_graphs = 2 n_nodes = 4 edge_index = torch.tensor([ [1, 2, 3], [0, 0, 0] ]) edge_indices = edge_index.repeat(2, 1, 1) xs = torch.stack([torch.randint(1, 5, (n_nodes, 8)) for _ in range(n_graphs)]) y = torch.randint(6, 10, (n_nodes, n_graphs)) data = SeqData(edge_indices, xs, y) data_list = [data, data] loader = DataLoader(data_list, batch_size=2, follow_batch=["xs", "y"]) batch = next(iter(loader))
运行输出:
SeqData(edge_indices=[2, 2, 3], xs=[2, 4, 8], y=[4, 2]) SeqDataBatch(edge_indices=[4, 2, 3], xs=[4, 4, 8], xs_batch=[4], y=[8, 2], y_batch=[8])
注:演示用了随机输入,示例里只用了2个图而非实际需求的6个,节点数量也和实际场景有差异。
遇到的问题
参照同构图逻辑在异构图场景下实现时得到空batch对象,无报错,自定义的__inc__方法也没有被调用,尝试的代码如下:
from torch_geometric.data import Data, HeteroData class SeqData(HeteroData): def __init__(self): super().__init__() def __inc__(self, key, value, store, *args, **kwargs): # 不确定此处应该如何实现 print(key, value) return 0 n_graphs = 2 n_nodes = 4 measurement_xs = torch.stack([torch.randint(1, 5, (n_nodes, 8)) for _ in range(n_graphs)]) measurement_y = torch.randint(6, 10, (n_nodes, n_graphs)) measurement_edge_index = torch.tensor([ [1, 2, 3], [0, 0, 0] ]) measurement_edge_indices = edge_index.repeat(2, 1, 1) data = SeqData() data["measurement"].x = measurement_xs data["measurement"].y = measurement_y data["measurement", "flows", "measurement"].edge_index = measurement_edge_indices data_list = [data, data] loader = DataLoader(data_list, batch_size=2) batch = next(iter(loader)) print(batch)
运行输出:
SeqDataBatch()
解决方案
出现空batch的核心原因是HeteroData的属性存储逻辑和普通Data不同,自定义批处理需要做两处修改:
- 重写
__cat_dim__方法指定序列维度的拼接方式,否则PyG会默认把多出来的序列维度当成节点/边维度做错误拼接,甚至直接丢弃属性 - 在
__inc__方法里针对不同节点类型、边类型分别返回正确的增量值,同时对非索引类属性返回0表示沿批次维度堆叠
正确的实现代码如下:
import torch from torch_geometric.data import HeteroData, NodeStorage, EdgeStorage from torch_geometric.loader import DataLoader class SeqHeteroData(HeteroData): def __inc__(self, key, value, store, *args, **kwargs): # 区分节点存储、边存储分别处理增量 if isinstance(store, NodeStorage): node_type = store._key if key == "edge_index": # 边索引增量为单图对应节点类型的节点数 return self[node_type].x.size(1) # 节点特征、标签等非索引属性不需要增量,沿batch维堆叠 return 0 elif isinstance(store, EdgeStorage): src_type, _, dst_type = store._key if key == "edge_index": # 边索引增量分别对应源节点、目标节点的单图节点数 return torch.tensor([ [self[src_type].x.size(1)], [self[dst_type].x.size(1)] ]) return 0 return super().__inc__(key, value, store, *args, **kwargs) def __cat_dim__(self, key, value, store, *args, **kwargs): # 所有属性沿第0维(batch维)拼接,不沿节点/边维度拼接 if key in ["x", "y", "edge_index"]: return 0 return super().__cat_dim__(key, value, store, *args, **kwargs) # 测试代码 n_graphs = 2 n_measure_nodes = 4 n_aux_nodes = 11 # 对应实际场景的辅助节点数 data = SeqHeteroData() # 测量节点属性:形状[序列长度, 节点数, 特征维度] data["measurement"].x = torch.stack([torch.randint(1,5,(n_measure_nodes,8)) for _ in range(n_graphs)]) data["measurement"].y = torch.randint(6,10,(n_measure_nodes, n_graphs)) # 辅助节点属性 data["aux"].x = torch.stack([torch.randn(n_aux_nodes, 16) for _ in range(n_graphs)]) # 同类型节点边 data["measurement", "flows", "measurement"].edge_index = torch.tensor([[1,2,3],[0,0,0]]).repeat(n_graphs,1,1) # 跨类型边示例 data["aux", "connects", "measurement"].edge_index = torch.randint(0, min(n_aux_nodes, n_measure_nodes), (2,5)).repeat(n_graphs,1,1) loader = DataLoader([data, data], batch_size=2) batch = next(iter(loader)) print(batch)
运行后可以得到正确的批处理结果,不会再出现空batch:
- 节点特征、标签会沿batch维度正确堆叠
- 边索引会自动加上对应节点数的偏移,保证不同样本的子图互不连通
- 多类型节点、多类型边都能正确处理,适配异构图场景
如果需要生成和同构图follow_batch对应的batch索引,可以在__inc__逻辑里补充对应属性的标记,或者在加载后手动根据张量形状生成即可。
内容的提问来源于stack exchange,提问作者PEREZje
相关产品推荐
相关产品推荐

