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

修改Networkx图节点数量用于GNN训练时触发RuntimeError

问题:删除Networkx节点后传入GCN出现索引越界错误

使用10节点图训练GCN时正常,但删除一个节点后传入GCN训练,抛出错误:

RuntimeError: index 9 is out of bounds for dimension 0 with size 9

但重新生成9节点图却能正常运行,相关代码如下:

GCN模型代码

class GCN(nn.Module):
    def __init__(self, input_size, hidden_size, num_classes):
        super(GCN, self).__init__()
        self.layer1 = GCNConv(input_size, hidden_size)
        self.layer2 = GCNConv(hidden_size, hidden_size)
        self.layer3 = GCNConv(hidden_size, num_classes)
        self.softmax = nn.Softmax(dim=0)

    def forward(self, node_features, edge_index):
        output = self.layer1(node_features, edge_index)
        output = torch.relu(output)
        output = self.layer2(output, edge_index)
        output = torch.relu(output)
        output = self.layer3(output, edge_index)
        output = self.softmax(output)

        return output

图生成与节点删除代码

def generate_graph(num_nodes):
    # generate weighted and connected graph
    Graph = nx.gnm_random_graph(num_nodes, random.randint(num_nodes, num_nodes*2), seed=42)
    while not nx.is_connected(Graph):
        Graph = nx.gnm_random_graph(num_nodes, random.randint(num_nodes, num_nodes*2), seed=42)

    # add features to nodes
    # node 0 will be the source node
    # each node will have a feature of 3
    # first feature will represent the node's bias (a random value between 0 and 1)
    # second feature will represent if the node is a source node (0 or 1, 1 if the node is the source node)
    # third feature will represent the node's degree
    for node in Graph.nodes:
        Graph.nodes[node]['feature'] = [random.random(), 1 if node == 0 else 0, Graph.degree[node]]

    node_features = Graph.nodes.data('feature')
    node_features = torch.tensor([node_feature[1] for node_feature in node_features])
    edge_index = torch.tensor(list(Graph.edges)).t().contiguous()

    return Graph, node_features, edge_index


def remove_node_from_graph(Graph, node):
    # remove the node from the graph
    Graph.remove_node(node)

    # update the features of the nodes
    for node in Graph.nodes:
        Graph.nodes[node]['feature'][2] = Graph.degree[node]

    node_features = Graph.nodes.data('feature')
    node_features = torch.tensor([node_feature[1] for node_feature in node_features])
    edge_index = torch.tensor(list(Graph.edges)).t().contiguous()

    return Graph, node_features, edge_index

问题原因

Networkx删除节点后,剩余节点的ID不会自动重新连续编号。比如初始10节点的ID是0-9,删除节点9后,剩余节点ID是0-8;但如果删除的是中间节点(比如5),剩余节点ID会出现断层(0-4,6-9)。此时生成的edge_index中仍会保留原节点ID,而node_features的长度是9(对应9个节点),索引范围仅为0-8。GCNConv在计算时会用edge_index里的原ID去索引node_features,当ID大于等于9时就会触发索引越界错误。

而重新生成9节点图时,节点ID是连续的0-8,edge_index中的索引全部合法,因此不会报错。


解决方案

删除节点后,需要手动对节点ID重新映射为连续的0~N-1(N为剩余节点数),同时更新edge_index中的节点索引。修改remove_node_from_graph函数如下:

def remove_node_from_graph(Graph, node):
    # remove the node from the graph
    Graph.remove_node(node)

    # 更新节点特征的度数
    for node in Graph.nodes:
        Graph.nodes[node]['feature'][2] = Graph.degree[node]
    
    # 重新映射节点ID为连续的0~N-1
    nodes = sorted(Graph.nodes())
    id_map = {old_id: new_id for new_id, old_id in enumerate(nodes)}
    
    # 生成新的节点特征(按新ID顺序)
    node_features = torch.tensor([Graph.nodes[old_id]['feature'] for old_id in nodes])
    
    # 生成新的edge_index,替换为新ID
    edges = [(id_map[u], id_map[v]) for u, v in Graph.edges()]
    edge_index = torch.tensor(edges).t().contiguous()

    return Graph, node_features, edge_index

关键改动说明

  1. 节点ID映射:创建id_map字典,将原节点ID映射为从0开始的连续新ID
  2. 节点特征重排:按照新ID对应的原节点顺序提取特征,保证node_features的索引与新ID一致
  3. 边索引更新:将边中的原节点ID替换为新ID,确保edge_index中的所有索引都在node_features的合法范围内

这样处理后,删除节点后的图数据就能正常传入GCN进行训练了。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.04 04:01:29