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

如何在PyTorch Geometric教程中替换Cora为自定义数据集?

替换PyTorch Geometric教程中Cora数据集为自定义数据集的实现方法

一、明确自定义数据集的核心结构

PyG处理图数据(和Cora这类单图任务一致)依赖以下核心张量:

  • x: 节点特征矩阵,形状为[num_nodes, num_features],类型为torch.FloatTensor
  • edge_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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 04:32:36