如何为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
相关产品推荐
相关产品推荐

