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

如何将NetworkX异质多重图转换为PyTorch Geometric的HeteroData?

将NetworkX异质Multigraph转换为PyTorch Geometric HeteroData的方法

PyTorch Geometric目前没有内置函数直接将nx.Multigraph转换为HeteroData,需要手动处理节点类型、边类型以及对应的属性,核心思路是按类型分组后转成PyTorch张量格式。以下是具体实现步骤和示例代码:

步骤说明

  1. 节点处理:按节点类型分组,将每个类型的节点映射为连续整数索引,并提取节点属性转换为张量。
  2. 边处理:按(源节点类型, 边类型, 目标节点类型)的组合分组,将Multigraph中的边转换为边索引张量,同时提取边属性(如果存在)。
  3. 构建HeteroData:将处理后的节点和边数据逐一添加到HeteroData对象中。

示例代码

1. 构造示例异质Multigraph

import networkx as nx
import torch
from torch_geometric.data import HeteroData

# 创建包含用户、物品节点,以及购买、浏览边的异质Multigraph
G = nx.MultiGraph()
# 用户节点:类型'user',属性'age'
G.add_node('u1', node_type='user', age=25)
G.add_node('u2', node_type='user', age=30)
# 物品节点:类型'item',属性'price'
G.add_node('i1', node_type='item', price=100)
G.add_node('i2', node_type='item', price=200)
# 边:支持同一对节点的多条同类型边
G.add_edge('u1', 'i1', edge_type='buy', amount=1)
G.add_edge('u1', 'i1', edge_type='buy', amount=2)
G.add_edge('u1', 'i2', edge_type='view', duration=5)
G.add_edge('u2', 'i1', edge_type='view', duration=3)

2. 转换为HeteroData

hetero_data = HeteroData()

# 处理节点:按类型生成索引映射和特征张量
node_types = set(nx.get_node_attributes(G, 'node_type').values())
for node_type in node_types:
    # 获取当前类型的所有节点
    nodes = [n for n, attr in G.nodes(data=True) if attr['node_type'] == node_type]
    # 节点到连续索引的映射
    node_map = {n: idx for idx, n in enumerate(nodes)}
    # 提取节点属性并转为张量(根据实际属性调整)
    if node_type == 'user':
        features = torch.tensor([attr['age'] for n, attr in G.nodes(data=True) if attr['node_type'] == node_type], dtype=torch.float).unsqueeze(1)
    elif node_type == 'item':
        features = torch.tensor([attr['price'] for n, attr in G.nodes(data=True) if attr['node_type'] == node_type], dtype=torch.float).unsqueeze(1)
    # 添加到HeteroData
    hetero_data[node_type].x = features
    hetero_data[node_type].node_map = node_map  # 保存映射用于边处理

# 处理边:按(源类型, 边类型, 目标类型)分组
edge_groups = {}
for u, v, attr in G.edges(data=True):
    u_type = G.nodes[u]['node_type']
    v_type = G.nodes[v]['node_type']
    edge_type = attr['edge_type']
    edge_key = (u_type, edge_type, v_type)
    
    if edge_key not in edge_groups:
        edge_groups[edge_key] = {'sources': [], 'targets': [], 'attrs': []}
    
    # 转换为连续索引
    u_idx = hetero_data[u_type].node_map[u]
    v_idx = hetero_data[v_type].node_map[v]
    edge_groups[edge_key]['sources'].append(u_idx)
    edge_groups[edge_key]['targets'].append(v_idx)
    
    # 提取边属性(根据实际属性调整)
    if edge_type == 'buy':
        edge_groups[edge_key]['attrs'].append(attr['amount'])
    elif edge_type == 'view':
        edge_groups[edge_key]['attrs'].append(attr['duration'])

# 将边数据添加到HeteroData
for edge_key, data in edge_groups.items():
    edge_index = torch.tensor([data['sources'], data['targets']], dtype=torch.long)
    hetero_data[edge_key].edge_index = edge_index
    # 添加边属性(如果存在)
    if data['attrs']:
        edge_attr = torch.tensor(data['attrs'], dtype=torch.float).unsqueeze(1)
        hetero_data[edge_key].edge_attr = edge_attr

# 查看转换结果
print(hetero_data)

注意事项

  • 如果节点/边包含多个属性,可以将多个属性拼接成一个特征向量,比如用torch.cat合并不同属性的张量。
  • 若节点无属性,可跳过x的赋值,或根据需求初始化默认特征(如全零张量)。
  • Multigraph中的多条边会被完整保留在edge_index中,PyTorch Geometric支持这种多边结构。

内容的提问来源于stack exchange,提问作者tiurina

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 11:17:49