PyG HeteroData数据集构建异常:节点数与图数量不符排查
PyG异构图构建问题:图数量与节点数不符合预期
问题描述
我是PyG新手,尝试从含5条记录的JSON文件构建数据集,原数据对应1张图、5个节点、8条边,但构建后查看属性发现图数量为3、节点数为20,不符合预期。推测是设置了org、event、player、rated四种节点类型导致的。目前仅需构建正确数据集,暂不考虑节点分类、链接预测等下游任务,其中rated字段应为标签(y)。
当前构建数据集代码
def build_dataset(self, edge_index, org_X, player_X, event_X, rated_X, labels_y): data = HeteroData() data['org'].x = org_X data['player'].x = player_X data['event'].x = event_X data['rated'].x = rated_X data['event', 'is_related_to', 'event'].edge_index = edge_index data['player', 'is_rated', 'rated'].y = labels_y return data
提取player节点特征代码
def extract_player_node_features(self, df): sorted_player_df = df.sort_values(by='player_id').set_index('player_id') sorted_player_df = sorted_player_df.reset_index(drop=False) player_id_mapping = sorted_player_df['player_id'] #print(f'\nPlayer ID mapping:\n{player_id_mapping}') # select player node features player_node_features_df = df[['player_name', 'age', 'school']] player_name_features_df = pd.DataFrame(player_node_features_df.player_name.values.tolist(), player_node_features_df.index).add_prefix('player_name_') player_name_features_ohe = pd.get_dummies(player_name_features_df) player_age_features_df = pd.DataFrame(player_node_features_df.age.values.tolist(), player_node_features_df.index).add_prefix('age_') player_age_features_ohe = pd.get_dummies(player_age_features_df) player_school_features_df = pd.DataFrame(player_node_features_df.school.values.tolist(), player_node_features_df.index).add_prefix('school_') player_school_features_ohe = pd.get_dummies(player_school_features_df) player_node_features = pd.concat([player_node_features_df, player_name_features_ohe], axis=1) player_node_features = pd.concat([player_node_features, player_age_features_ohe], axis=1) player_node_features = pd.concat([player_node_features, player_school_features_ohe], axis=1) player_node_features.drop(columns=['player_name', 'age', 'school'], axis=1, inplace=True) player_node_X = player_node_features.to_numpy(dtype='int32') player_node_X = torch.from_numpy(player_node_X) return player_node_X
输入数据
原始输入DataFrame
event_id event_type org_id org_name org_location player_id player_name age school related_event_id rated 0 1-ab3 qualifiers 305 milan tennis club Milan 1-b7a3-52d2 Alex 20 BCE [4-ab3, 3-ab3] no 1 2-ab3 under 18 finals 76 Nadal tennis academy madrid 2-b7a3-52d2 Bob 20 BCMS [5-ab3, 1-ab3] yes 2 3-ab3 womens tennis qualifiers 185 Griz tennis club budapest 3-b7a3-52d2 Mary 21 BCE [4-ab3] no 3 4-ab3 US professional tennis club 285 Nick Bolletieri Tennis Academy tampa 4-b7a3-52d2 Joe 21 BCMS [1-ab3, 3-ab3] yes 4 5-ab3 womens tennis circuit 305 milan tennis club Milan 5-b7a3-52d2 Bolt 22 LTHS [4-ab3] no
数值化排序后的DataFrame
related_event_id org_id org_name org_location player_id player_name event_id event_type age school rated 1 [4, 0] 0 1 2 1 1 1 2 0 1 1 2 [3] 1 0 1 2 4 2 4 1 0 0 3 [0, 2] 2 2 3 3 3 3 0 1 1 1 0 [3, 2] 3 3 0 0 0 0 1 0 0 0 4 [3] 3 3 0 4 2 4 3 2 2 0
报错信息
若不将rated设为节点类型,调用validate()函数会报错:
ValueError: The node types {'rated'} are referenced in edge types but do not exist as node types
现需排查图数量和节点数不符合预期的原因,可提供完整代码及输入文件用于复现。
内容的提问来源于stack exchange,提问作者user1717931
相关产品推荐
相关产品推荐

