PyTorch Geometric自定义Data类批处理时遇None传入与num_nodes问题
自定义PyTorch Geometric Data子类的批处理问题解决
问题根源
PyTorch Geometric的DataLoader批处理时,只会自动识别并处理内置属性(如x、edge_index),自定义张量属性如果不明确声明类型(节点级/图级/边级),框架不知道怎么拼接,就会在批处理时把它设为None。另外,在__init__里用自定义张量长度设num_nodes会报错,因为批处理时框架会先创建空实例,此时自定义张量为None,调用len()就会触发TypeError。
核心解决步骤
1. 给自定义属性声明拼接规则
必须在自定义Data类里实现__inc__和__cat_dim__方法,告诉框架如何拼接你的自定义属性:
__inc__:定义批处理时每个属性的增量值(比如节点级属性按num_nodes累加)__cat_dim__:定义属性的拼接维度
2. 避免在__init__中动态设置num_nodes
要么创建实例时显式传入num_nodes,要么让框架通过节点级属性或edge_index自动推断,不要在初始化方法里依赖自定义张量的长度。
完整可运行示例
import torch from torch_geometric.data import Data from torch_geometric.loader import DataLoader class CustomData(Data): def __inc__(self, key, value, *args, **kwargs): # 节点级属性a:每个图的增量是当前图的节点数 if key == 'a': return self.num_nodes # 边级属性(如果有的话):增量是当前图的边数 # elif key == 'custom_edge_attr': # return self.edge_index.size(1) # 其他属性沿用父类默认规则 return super().__inc__(key, value, *args, **kwargs) def __cat_dim__(self, key, value, *args, **kwargs): # 节点级属性按第0维拼接 if key == 'a': return 0 # 图级属性(如b)不用特殊处理,框架默认按第0维拼接 return super().__cat_dim__(key, value, *args, **kwargs) # 创建单个实例 data1 = CustomData(a=torch.tensor([1,2,3]), b=torch.tensor([10])) # b是图级属性 data2 = CustomData(a=torch.tensor([4,5]), b=torch.tensor([20])) # 批处理测试 loader = DataLoader([data1, data2], batch_size=2) batch = next(iter(loader)) print(batch.a) # 输出: tensor([1, 2, 3, 4, 5]) print(batch.b) # 输出: tensor([10, 20]) print(batch.batch) # 输出: tensor([0, 0, 0, 1, 1]),正常生成batch属性
关于batch属性的补充说明
batch属性是框架为节点级属性自动生成的,用来标记每个节点所属的图。只要你的CustomData实例包含节点级属性(并正确声明拼接规则)或edge_index,框架就能自动推断num_nodes,进而生成batch属性,不需要手动设置。
内容的提问来源于stack exchange,提问作者b-riley
相关产品推荐
相关产品推荐

