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

