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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.15 15:22:17