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

如何让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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.01 13:17:06