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

如何为torch_geometric.data Data对象添加属性并重构TUDataset?

问题

我尝试扩展TUDataset数据集的元素,该数据集通过代码dataset = TUDataset("PROTEIN", name=PROTEIN, use_node_attr=True)获取,希望为每个样本添加一个向量型特征。我尝试通过以下代码实现:

for i, current_g in enumerate(dataset):
    nxgraph = nx.to_numpy_array(torch_geometric.utils.to_networkx(current_g))
    feature = do_something(nxgraph)
    dataset[i].new_feature = feature

但该代码无法正常工作,直接为dataset元素添加属性会报错:

In [80]: dataset[2].test = 1

In [81]: dataset[2].test
---------------------------------------------------------------------------
AttributeError                            Traceback (most recent call last)
~/workspace/grouptheoretical/new-experiments/HGP-SL-myfork/main.py in <cell line: 1>()
----> 1 dataset[2].test

AttributeError: 'Data' object has no attribute 'test'

In [82]: dataset[2].__setattr__('test', 1)

In [83]: dataset[2].test
---------------------------------------------------------------------------
AttributeError                            Traceback (most recent call last)
~/workspace/grouptheoretical/new-experiments/HGP-SL-myfork/main.py in <cell line: 1>()
----> 1 dataset[2].test

AttributeError: 'Data' object has no attribute 'test'

数据集的元素是torch_geometric.data下的Data对象。我可以通过以下方式创建包含所需特征的新Data对象:

tmp=dataset[i].to_dict()
tmp['new_feature'] = feature
new_dataset[i]=torch_geometric.data.Data.from_dict(tmp)

但我不清楚如何将Data对象列表转换为TUDataset或其父类数据集,请问该如何解决这个问题?

解决方法

核心原因

TUDataset返回的Data对象是只读实例,每次通过dataset[i]访问时都会生成新的Data对象,直接修改的属性无法被持久化保存。

具体实现方案

你可以将生成的Data对象列表转换为InMemoryDataset(TUDataset的父类),有两种可行方式:

方式1:自定义InMemoryDataset子类(规范做法)

import torch
import torch_geometric
from torch_geometric.data import InMemoryDataset, Data

# 1. 生成所有带新特征的Data对象列表
new_data_list = []
for current_g in dataset:
    nxgraph = nx.to_numpy_array(torch_geometric.utils.to_networkx(current_g))
    feature = do_something(nxgraph)  # 你的特征计算逻辑
    tmp = current_g.to_dict()
    tmp['new_feature'] = feature
    new_data = Data.from_dict(tmp)
    new_data_list.append(new_data)

# 2. 定义自定义数据集类
class CustomProteinDataset(InMemoryDataset):
    def __init__(self, root, transform=None, pre_transform=None):
        super().__init__(root, transform, pre_transform)
        self.data, self.slices = torch.load(self.processed_paths[0])

    @property
    def processed_file_names(self):
        return ['custom_protein_data.pt']

    def process(self):
        data, slices = self.collate(new_data_list)
        torch.save((data, slices), self.processed_paths[0])

# 3. 实例化自定义数据集
custom_dataset = CustomProteinDataset(root='./custom_protein_dataset')

方式2:手动构建数据集(简化做法)

如果不想自定义子类,可直接通过collate方法处理数据列表后赋值:

from torch_geometric.data import InMemoryDataset, Data

# 1. 先生成带新特征的Data对象列表(同方式1的第一步)
new_data_list = []
for current_g in dataset:
    nxgraph = nx.to_numpy_array(torch_geometric.utils.to_networkx(current_g))
    feature = do_something(nxgraph)
    tmp = current_g.to_dict()
    tmp['new_feature'] = feature
    new_data = Data.from_dict(tmp)
    new_data_list.append(new_data)

# 2. 直接构建InMemoryDataset
data, slices = InMemoryDataset.collate(new_data_list)
custom_dataset = InMemoryDataset(root='./custom_protein_dataset')
custom_dataset.data = data
custom_dataset.slices = slices

使用验证

完成后即可像普通TUDataset一样操作:

# 验证新特征存在
print(custom_dataset[0].new_feature)

# 构建DataLoader
from torch_geometric.data import DataLoader
loader = DataLoader(custom_dataset, batch_size=32)

内容的提问来源于stack exchange,提问作者asdf

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.07 11:35:19