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

PyGeometric中KNNGraph Transform在数据集上的使用问题排查

问题描述

我有一个存储点云及节点数据的DataFrame,包含X、Y、Z坐标和节点特征,需要将其转换为图结构用于GNN模型。我尝试用PyGeometric的KNNGraph Transform提取近邻边连接信息,参考官方文档实现了自定义InMemoryDataset并设置了transform,但调用dataset_cloud[0]时只返回带pos属性的Data对象,没有生成近邻边连接。虽然直接调用KNNGraph方法能得到预期的edge_index属性,但想知道原方法失效的原因。

尝试代码

import pandas as pd
import numpy as np
import torch
from torch_geometric.data import Data, Dataset,InMemoryDataset
from torch_geometric.transforms import SamplePoints, KNNGraph
import torch_geometric.transforms as T
from torch_geometric.datasets import GeometricShapes

class CustomDataset(InMemoryDataset):
    def __init__(self, listOfDataObjects):
        super().__init__()
        self.data, self.slices = self.collate(listOfDataObjects)
    
    def __len__(self):
        return len(self.slices)
    
    def __getitem__(self, idx):
        sample = self.get(idx)
        return sample

## 生成测试数据
data_fake = pd.DataFrame(data=np.random.rand(20,5), columns =['X','Y', 'Z','Node_feature_1', 'Node_feature_2'])

# 创建Data对象和Dataset
my_fake_data = Data()
my_fake_data.pos =  torch.from_numpy(data_fake[['X','Y', 'Z']].values)
dataset_cloud = CustomDataset([my_fake_data])

# 设置Transform
dataset_cloud.transform = T.Compose([SamplePoints(num=20), KNNGraph(k=5)])

可行的直接调用代码

knn = KNNGraph(k=3)
data_knn = knn(dataset_cloud[0])
失效原因分析

自定义InMemoryDataset没有正确集成PyGeometric的transform机制,问题出在两点:

  • 初始化未传递transform参数:InMemoryDataset父类构造函数需要接收transform参数并维护相关逻辑,但你的__init__方法没有将transform传递给父类,导致类无法识别这个属性。
  • __getitem__未应用transform:PyGeometric的Dataset默认会在__getitem__中自动执行transform,但你重写了该方法,直接返回原始样本,没有调用transform处理数据。
修复方案

修改自定义InMemoryDataset的实现,正确集成transform机制:

import pandas as pd
import numpy as np
import torch
from torch_geometric.data import Data, InMemoryDataset
from torch_geometric.transforms import SamplePoints, KNNGraph
import torch_geometric.transforms as T

class CustomDataset(InMemoryDataset):
    def __init__(self, listOfDataObjects, transform=None):
        # 将transform传递给父类构造函数
        super().__init__(transform=transform)
        self.data, self.slices = self.collate(listOfDataObjects)
    
    def __len__(self):
        return len(self.slices)
    
    def __getitem__(self, idx):
        sample = self.get(idx)
        # 应用transform(如果存在)
        if self.transform is not None:
            sample = self.transform(sample)
        return sample

## 生成测试数据
data_fake = pd.DataFrame(data=np.random.rand(20,5), columns =['X','Y', 'Z','Node_feature_1', 'Node_feature_2'])

# 创建Data对象和Dataset,初始化时传入transform
my_fake_data = Data()
my_fake_data.pos =  torch.from_numpy(data_fake[['X','Y', 'Z']].values)
dataset_cloud = CustomDataset(
    [my_fake_data],
    transform=T.Compose([SamplePoints(num=20), KNNGraph(k=5)])
)

# 现在调用dataset_cloud[0]会自动生成edge_index
print(dataset_cloud[0].edge_index)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.13 22:46:16