如何在PyTorch Geometric教程中替换Cora为自定义数据集?
替换PyTorch Geometric教程中Cora数据集为自定义数据集的实现方法
一、明确自定义数据集的核心结构
PyG处理图数据(和Cora这类单图任务一致)依赖以下核心张量:
x: 节点特征矩阵,形状为[num_nodes, num_features],类型为torch.FloatTensoredge_index: 边索引矩阵,形状为[2, num_edges],类型为torch.LongTensor(存储每条边的源节点、目标节点索引)y: 节点标签,分类任务下形状为[num_nodes],类型为torch.LongTensor
若需要训练/验证/测试划分,还需train_mask、val_mask、test_mask三个布尔张量,形状均为[num_nodes]
二、实现自定义数据集的两种常用方式
方式1:基于InMemoryDataset(小数据集首选,加载至内存)
这是和Cora数据集逻辑一致的实现方式,适合数据量较小的场景:
import torch import pandas as pd from torch_geometric.data import InMemoryDataset, Data class CustomGraphDataset(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 raw_file_names(self): # 替换为你的原始数据文件名,比如节点特征、边列表、标签文件 return ['node_features.csv', 'edge_list.csv', 'node_labels.csv'] @property def processed_file_names(self): # 处理后保存的文件名 return ['custom_data.pt'] def process(self): # 1. 加载原始数据(根据你的数据格式调整,这里以CSV为例) node_df = pd.read_csv(self.raw_paths[0]) edge_df = pd.read_csv(self.raw_paths[1]) label_df = pd.read_csv(self.raw_paths[2]) # 2. 转换为PyG要求的张量格式 x = torch.tensor(node_df.drop('node_id', axis=1).values, dtype=torch.float) edge_index = torch.tensor(edge_df[['source', 'target']].values.T, dtype=torch.long) y = torch.tensor(label_df['label'].values, dtype=torch.long) # 3. 生成训练/验证/测试划分(若已有固定划分,直接加载对应mask即可) num_nodes = x.size(0) train_mask = torch.zeros(num_nodes, dtype=torch.bool) val_mask = torch.zeros(num_nodes, dtype=torch.bool) test_mask = torch.zeros(num_nodes, dtype=torch.bool) # 示例:随机选20个训练节点、500个验证节点、1000个测试节点(按需调整) train_idx = torch.randperm(num_nodes)[:20] val_idx = torch.randperm(num_nodes)[20:520] test_idx = torch.randperm(num_nodes)[520:1520] train_mask[train_idx] = True val_mask[val_idx] = True test_mask[test_idx] = True # 4. 创建PyG Data对象 data = Data(x=x, edge_index=edge_index, y=y, train_mask=train_mask, val_mask=val_mask, test_mask=test_mask) # 应用预转换(如果有) if self.pre_transform is not None: data = self.pre_transform(data) # 保存处理后的数据 data, slices = self.collate([data]) torch.save((data, slices), self.processed_paths[0])
方式2:基于Dataset(大数据场景,按需加载)
若数据集过大无法一次性加载到内存,可继承torch_geometric.data.Dataset,实现len()和get()方法按需加载数据。不过Cora是单图任务,大部分场景下方式1足够覆盖需求。
三、替换教程中的数据集加载代码
原教程加载Cora的代码:
from torch_geometric.datasets import Planetoid dataset = Planetoid(root='data/Planetoid', name='Cora')
替换为自定义数据集加载:
# 确保CustomGraphDataset类已导入 dataset = CustomGraphDataset(root='data/CustomGraph')
后续的模型训练、评估逻辑可完全复用教程代码,只要自定义数据集的Data对象结构与Cora一致(包含x、edge_index、y、各类mask)。
四、验证自定义数据集正确性
加载后可打印关键信息确认结构匹配:
print(f'图数量: {len(dataset)}') data = dataset[0] print(f'节点数量: {data.num_nodes}') print(f'边数量: {data.num_edges}') print(f'特征维度: {data.num_features}') print(f'类别数量: {dataset.num_classes}') print(f'训练节点数: {data.train_mask.sum()}') print(f'验证节点数: {data.val_mask.sum()}') print(f'测试节点数: {data.test_mask.sum()}')
五、关键注意事项
- 节点索引一致性:边列表中的节点ID必须和节点特征的索引一一对应,避免索引不匹配
- 标签编码:分类任务的标签需转换为从0开始的整数编码,不要直接用字符串
- 数据格式适配:若原始数据不是CSV,只需在
process()方法中调整加载逻辑(如读取JSON、txt等) - 划分逻辑:若数据集已有官方训练/验证/测试划分,直接加载对应mask即可,无需随机生成
内容的提问来源于stack exchange,提问作者SILA
相关产品推荐
相关产品推荐

