如何让PyTorch Geometric的RGCNConv适配HeteroData进行链接预测?
异构图链接预测:RGCNConv适配HeteroData的问题
我用PyTorch Geometric基于外部数据集构建了HeteroData对象,要做链接预测任务,目标是用RGCNConv模型,但始终无法构建出能接收HeteroData作为输入的模型结构,已经参考过官方异构图教程。
处理后的HeteroData结构
HeteroData( admission={ x=[1024] }, medicine={ x=[1024] }, diagnosis={ x=[1024] }, procedure={ x=[1024] }, (admission, prescribed_to, medicine)={ edge_index=[2, 1024], edge_weight=1, }, (admission, diagnosed_with, diagnosis)={ edge_index=[2, 1024], edge_weight=1, }, (admission, procedure_done, procedure)={ edge_index=[2, 1024], edge_weight=1, }, (medicine, interacts_with, medicine)={ edge_index=[2, 2048], edge_weight=-1, }, (medicine, rev_prescribed_to, admission)={ edge_index=[2, 1024], edge_weight=1, }, (diagnosis, rev_diagnosed_with, admission)={ edge_index=[2, 1024], edge_weight=1, }, (procedure, rev_procedure_done, admission)={ edge_index=[2, 1024], edge_weight=1, } )
现有模型代码
class RelationPredictionModel(nn.Module): def __init__(self, input_dim, hidden_dim, output_dim, num_relations, num_layers, dropout_rate=0.5): super(RelationPredictionModel, self).__init__() # Define RGCNConv layers self.convs = nn.ModuleList() for _ in range(num_layers): conv = HeteroConv({ # Using the Heterogeneous Convolution Wrapper ('admission', 'prescribed_to', 'medicine'): RGCNConv(in_channels=input_dim, out_channels=hidden_dim, num_relations=num_relations), ... # similar code for each relation ('procedure', 'rev_procedure_done', 'admission'): RGCNConv(in_channels=input_dim, out_channels=hidden_dim, num_relations=num_relations) }, aggr='sum') self.convs.append(conv) self.relu = nn.ReLU() self.dropout = nn.Dropout(p=dropout_rate) def forward(self, x_dict, edge_index_dict, edge_types): for conv in self.convs: x_dict = conv(x_dict, edge_index_dict, edge_types) x_dict = {key: x.relu() for key, x in x_dict.items()} x_dict = self.dropout(x_dict) return x_dict
训练代码
model = RelationPredictionModel2(input_dim, hidden_dim, output_dim, num_relations, num_layers) model = model.to(device) for epoch in range(num_epochs): model.train() total_loss = 0 for batch_data in train_loader: batch_data = batch_data.to(device) optimizer.zero_grad() x_dict = batch_data.x_dict edge_index_dict = batch_data.edge_index_dict edge_types = batch_data.edge_types # Forward pass output = model(x_dict, edge_index_dict, edge_types) ...
运行报错
TypeError Traceback (most recent call last) <ipython-input-46-a33be3f75ca2> in <cell line: 19>() 35 36 # Forward pass ---> 37 output = model(x_dict, edge_index_dict, edge_types) 38 39 # Compute loss /usr/local/lib/python3.10/dist-packages/torch/nn/modules/module.py in _wrapped_call_impl(self, *args, **kwargs) 1516 return self._compiled_call_impl(*args, **kwargs) # type: ignore[misc] 1517 else: -> 1518 return self._call_impl(*args, **kwargs) 1519 1520 def _call_impl(self, *args, **kwargs): /usr/local/lib/python3.10/dist-packages/torch/nn/modules/module.py in _call_impl(self, *args, **kwargs) 1525 or _global_backward_pre_hooks or _global_backward_hooks 1526 or _global_forward_hooks or _global_forward_pre_hooks): -> 1527 return forward_call(*args, **kwargs) 1528 1529 try: <ipython-input-25-30ce90f7d49d> in forward(self, x_dict, edge_index_dict, edge_types) 36 # Graph convolutional layers 37 for conv in self.convs: -> 38 x_dict = conv(x_dict, edge_index_dict, edge_types) 39 x_dict = {key: x.relu() for key, x in x_dict.items()} 40 /usr/local/lib/python3.10/dist-packages/torch/nn/modules/module.py in _wrapped_call_impl(self, *args, **kwargs) 1516 return self._compiled_call_impl(*args, **kwargs) # type: ignore[misc] 1517 else: -> 1518 return self._call_impl(*args, **kwargs) 1519 1520 def _call_impl(self, *args, **kwargs): /usr/local/lib/python3.10/dist-packages/torch/nn/modules/module.py in _call_impl(self, *args, **kwargs) 1525 or _global_backward_pre_hooks or _global_backward_hooks 1526 or _global_forward_hooks or _global_forward_pre_hooks): -> 1527 return forward_call(*args, **kwargs) 1528 1529 try: /usr/local/lib/python3.10/dist-packages/torch_geometric/nn/conv/hetero_conv.py in forward(self, *args_dict, **kwargs_dict) 125 if edge_type in value_dict: 126 has_edge_level_arg = True -> 127 args.append(value_dict[edge_type]) 128 elif src == dst and src in value_dict: 129 args.append(value_dict[src]) TypeError: list indices must be integers or slices, not tuple
尝试to_hetero转换后的报错
Traceback (most recent call last): File "/usr/local/lib/python3.10/dist-packages/torch/fx/graph_module.py", line 274, in __call__ return super(self.cls, obj).__call__(*args, **kwargs) # type: ignore[misc] File "/usr/local/lib/python3.10/dist-packages/torch/nn/modules/module.py", line 1518, in _wrapped_call_impl return self._call_impl(*args, **kwargs) File "/usr/local/lib/python3.10/dist-packages/torch/nn/modules/module.py", line 1527, in _call_impl return forward_call(*args, **kwargs) File "<eval_with_key>.5", line 16, in forward edge_type__user__to__artist = edge_type_dict.get(('user', 'to', 'artist'), None) AttributeError: 'list' object has no attribute 'get' Call using an FX-traced Module, line 16 of the traced Module's generated forward function: edge_type_dict = torch_geometric_nn_to_hetero_transformer_get_dict(edge_type); edge_type = None edge_type__user__to__artist = edge_type_dict.get(('user', 'to', 'artist'), None)
求能将HeteroData传入RGCNConv模型的解决方案。
PS:抱歉帖子较长,这是我第一次发帖,尽量做到内容详尽。
内容的提问来源于stack exchange,提问作者Rajat Mishra
相关产品推荐
相关产品推荐

