修改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
关键改动说明
- 节点ID映射:创建
id_map字典,将原节点ID映射为从0开始的连续新ID - 节点特征重排:按照新ID对应的原节点顺序提取特征,保证
node_features的索引与新ID一致 - 边索引更新:将边中的原节点ID替换为新ID,确保
edge_index中的所有索引都在node_features的合法范围内
这样处理后,删除节点后的图数据就能正常传入GCN进行训练了。
内容的提问来源于stack exchange,提问作者Protik Nag
相关产品推荐
相关产品推荐

