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

为HeteroData创建NeighborLoader时出现EdgeStorage无num_nodes属性错误求助

修复方案及替代思路

问题根因

这个错误是因为NeighborLoader在异构图采样时,需要从边存储(EdgeStorage)关联的节点类型中获取num_nodes属性,但你的数据结构中该属性缺失或未正确关联——无批处理时全图训练不需要动态采样节点,所以不会触发这个校验。


具体修复步骤

1. 显式设置所有节点类型的num_nodes

PyG有时会通过edge_index自动推断节点数,但划分数据集后可能破坏这个关联。手动给每个节点类型添加num_nodes:

# 假设你的异构图有'user'和'item'两种节点类型
data['user'].num_nodes = data['user'].x.shape[0]  # 用特征矩阵行数直接设置
data['item'].num_nodes = data['item'].x.shape[0]

如果节点类型没有特征矩阵(比如只有边连接),可以用edge_index中的最大索引+1来设置:

max_user_idx = data[('user', 'interacts', 'item')].edge_index[0].max().item()
data['user'].num_nodes = max_user_idx + 1

2. 修正RandomNodeSplit的用法

异构图下RandomNodeSplit默认只处理单一节点类型,若需要划分多个节点类型,需明确指定目标类型,或分别生成mask:

from torch_geometric.transforms import RandomNodeSplit

# 针对目标节点类型划分(比如'user')
transform = RandomNodeSplit(
    num_val=0.1,
    num_test=0.2,
    node_type='user',  # 明确要划分的节点类型
    key='train_mask'
)
data = transform(data)

# 若需多节点类型划分,手动生成mask
for node_type in data.node_types:
    total_nodes = data[node_type].num_nodes
    perm = torch.randperm(total_nodes)
    train_idx = perm[:int(0.7*total_nodes)]
    val_idx = perm[int(0.7*total_nodes):int(0.8*total_nodes)]
    test_idx = perm[int(0.8*total_nodes):]
    
    data[node_type].train_mask = torch.zeros(total_nodes, dtype=torch.bool)
    data[node_type].val_mask = torch.zeros(total_nodes, dtype=torch.bool)
    data[node_type].test_mask = torch.zeros(total_nodes, dtype=torch.bool)
    
    data[node_type].train_mask[train_idx] = True
    data[node_type].val_mask[val_idx] = True
    data[node_type].test_mask[test_idx] = True

3. 正确配置NeighborLoader参数

异构图的NeighborLoader必须指定input_nodes为节点类型+索引的元组,同时针对每个边类型设置采样数量:

from torch_geometric.loader import NeighborLoader

# 获取目标节点类型的训练索引
train_indices = data['user'].train_mask.nonzero().flatten()

# 初始化NeighborLoader
loader = NeighborLoader(
    data,
    input_nodes=('user', train_indices),  # 明确节点类型和采样起始节点
    # 为每个边类型设置采样层数(比如两层分别采10、5个邻居)
    num_neighbors={
        ('user', 'follows', 'user'): [10, 5],
        ('user', 'interacts', 'item'): [10, 5],
        ('item', 'interacted_by', 'user'): [10, 5]
    },
    batch_size=64,
    shuffle=True,
)

替代方案

如果上述调整仍无法解决问题,可以改用HeteroDataLoader做简单的节点批处理(不采样邻居,仅按批次取节点),适合小图场景:

from torch_geometric.loader import HeteroDataLoader
from torch_geometric.data import Batch

def collate_fn(batch):
    return Batch.from_data_list(batch)

# 按节点类型生成批次数据列表
train_data_list = [data.subgraph({
    'user': data['user'].train_mask,
    'item': torch.ones(data['item'].num_nodes, dtype=torch.bool)  # 保留所有item节点
})]
# 若需要更小批次,手动拆分train_indices为多个子集合生成subgraph

loader = HeteroDataLoader(train_data_list, batch_size=64, shuffle=True, collate_fn=collate_fn)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 08:44:51