You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.31 21:09:21