异构GNN(HeteroGAT)运行时出现索引错误的解决方案咨询
异构GNN(HeteroGAT)运行时出现索引错误的解决方案咨询
看起来你遇到了PyTorch Geometric(PyG)异构图模型里的索引匹配问题,我来帮你理清楚原因和解决办法~
错误原因分析
PyG的HeteroData对每种节点类型的索引是局部独立的,不是全局连续计数的:
- 你的
user节点特征是[100,16],意味着user类型的节点局部索引范围是0~99 keyword节点特征是[321,16],对应局部索引范围是0~320tweet节点特征是[1000,16],对应局部索引范围是0~999
但你构建edge_index时用了全局连续索引(比如keyword用100420,tweet用4211420),这就导致模型在访问节点特征时,出现了超出对应类型局部索引范围的数值(比如1420超过了tweet的最大局部索引999),从而触发报错。
解决方案:将全局索引转换为局部索引
你需要把edge_index里的全局索引,转换成对应节点类型的局部索引,具体步骤如下:
- 先定义每种节点类型的全局起始偏移量:
# 全局索引的起始值,对应你定义的节点索引规则 offset_dict = { 'user': 0, 'keyword': 100, 'tweet': 100 + 321 # user的数量 + keyword的数量 = 421 }
- 逐个处理每个边类型的
edge_index,将全局索引减去对应类型的偏移量:
import torch # 处理(user, follow, user):两边都是user,偏移量为0,无需修改 # 处理(user, tweetedby, tweet):右侧是tweet节点,转换为局部索引 edge_index_ut = edge_index_dict[('user', 'tweetedby', 'tweet')].clone() edge_index_ut[1] -= offset_dict['tweet'] # 全局索引 - 偏移量 = 局部索引 edge_index_dict[('user', 'tweetedby', 'tweet')] = edge_index_ut # 处理(keyword, haskeyword, tweet):左右两侧分别转换 edge_index_kt = edge_index_dict[('keyword', 'haskeyword', 'tweet')].clone() edge_index_kt[0] -= offset_dict['keyword'] # keyword节点转局部索引 edge_index_kt[1] -= offset_dict['tweet'] # tweet节点转局部索引 edge_index_dict[('keyword', 'haskeyword', 'tweet')] = edge_index_kt
- 验证转换结果:
- 确保
user类型的索引范围是0~99 keyword类型的索引范围是0~320tweet类型的索引范围是0~999
- 确保
完成上述转换后,你的edge_index里的所有索引都会匹配对应节点类型的特征矩阵维度,再运行HeteroGATBinaryClassifier模型就不会出现索引错误了。
备注:内容来源于stack exchange,提问作者aliiiiiiiiiiiiiiiiiiiii
相关产品推荐
相关产品推荐

