使用GNN(Node2Vec)构建分类算法时遇权重和非有限值错误
图神经网络分类算法调试:解决
ValueError: Total of weights must be finite错误 问题描述
在Python中基于图神经网络构建分类算法时,触发错误:ValueError: Total of weights must be finite。
已完成以下排查步骤但问题未解决:
- 移除了数据框中的NaN和无限值
- 检查并验证了图构建逻辑、节点及边属性
- 在训练Node2Vec模型后,使用其嵌入训练Random Forest分类器时仍报错
使用的代码如下:
# Função para criar um grafo a partir de uma linha do dataframe def criar_grafo_passe(row, index): G = nx.Graph() passer = f"Player_PP" # Identificador único para o jogador que passa a bola receiver = f"Player_PR" # Identificador único para o jogador que recebe a bola em t0 receiver_t1 = f"Player_PRt1" # Identificador único para o jogador que recebe a bola em t1 opponent_pp = f"Opponent_PP" # Identificador único para o oponente mais próximo do passador opponent_pr = f"Opponent_PR" # Identificador único para o oponente mais próximo do receptor em t0 opponent_pr_t1 = f"Opponent_PRt1" # Identificador único para o oponente mais próximo do receptor em t1 distance = row['Passing distance'] # Distância do passe # Adiciona nós para passador, receptor, receptor em t1 e oponentes mais próximos G.add_node(passer, role='passer') G.add_node(receiver, role='receiver') G.add_node(receiver_t1, role='receiver_t1') G.add_node(opponent_pp, role='opponent') G.add_node(opponent_pr, role='opponent') G.add_node(opponent_pr_t1, role='opponent') # Adiciona arestas entre passador e receptor G.add_edge(passer, receiver, weight=distance) # Adiciona arestas entre passador e oponente mais próximo G.add_edge(passer, opponent_pp, weight=row['Nearest opp. PPt0']) # Adiciona arestas entre receptor e oponente mais próximo em t0 e t1 G.add_edge(receiver, opponent_pr, weight=row['Nearest opp. PRt0']) G.add_edge(receiver_t1, opponent_pr_t1, weight=row['Nearest opp. PRt1']) # Adiciona aresta entre receptor em t0 e t1 G.add_edge(receiver, receiver_t1, weight=row['Displacement PR']) # Adiciona atributos aos nós do passador G.nodes[passer].update({ 'Foot or not': row['Foot or not'], 'Nearest opp. PPt0': row['Nearest opp. PPt0'], 'Velocity nearest opp. PPt0': row['Velocity nearest opp. PPt0'], 'Opponet angle': row['Opponet angle'], 'Density PPt0': [row['Density PPt0 (1m)'], row['Density PPt0 (2m)'], row['Density PPt0 (5m)'], row['Density PPt0 (10m)']], 'Velocity PPt0': row['Velocity PPt0'], 'Distance PPt0 to target': row['Distance PPt0 to target'] }) # Adiciona atributos aos nós do receptor em t0 G.nodes[receiver].update({ 'Nearest opp. PRt0': row['Nearest opp. PRt0'], 'Density PRt0': [row['Density PRt0 (1m)'], row['Density PRt0 (2m)'], row['Density PRt0 (5m)'], row['Density PRt0 (10m)']], 'Velocity PRt0': row['Velocity PRt0'], 'Distance PRt0 to target': row['Distance PRt0 to target'] }) # Adiciona atributos aos nós do receptor em t1 G.nodes[receiver_t1].update({ 'Nearest opp. PRt1': row['Nearest opp. PRt1'], 'Density PRt1': [row['Density PRt1 (1m)'], row['Density PRt1 (2m)'], row['Density PRt1 (5m)'], row['Density PRt1 (10m)']], 'Velocity PRt1': row['Velocity PRt1'], 'Distance PRt1 to target': row['Distance PRt1 to target'], 'Out ball angle': row['Out ball angle'] }) # Adiciona atributos aos nós dos oponentes mais próximos G.nodes[opponent_pp].update({ 'role': 'opponent', 'distance to passer': row['Nearest opp. PPt0'] }) G.nodes[opponent_pr].update({ 'role': 'opponent', 'distance to receiver t0': row['Nearest opp. PRt0'] }) G.nodes[opponent_pr_t1].update({ 'role': 'opponent', 'distance to receiver t1': row['Nearest opp. PRt1'] }) # Adiciona atributos à aresta entre passador e receptor G[passer][receiver].update({ 'Passing distance': row['Passing distance'], 'Ball velocity': row['Ball velocity'], 'Passing angle': row['Passing angle'], 'Ball progression': row['Ball progression'], 'Outplayed opp. ': row['Outplayed opp. '], 'Opp. btw PRt1 and target': row['Opp. btw PRt1 and target'], 'Accuracy': row['Accuracy'], 'One touth': row['One touth'] }) return G # Remover valores NaN ou infinitos do dataframe df.dropna(inplace=True) df.replace([np.inf, -np.inf], np.nan, inplace=True) df.dropna(inplace=True) # Criar uma lista de grafos, um para cada passe grafos_passes = [criar_grafo_passe(row, index) for index, row in df.iterrows()] # Treinar o algoritmo Node2Vec em cada grafo embeddings = [] for G in grafos_passes: # Verificar e remover qualquer peso não finito do grafo for u, v, data in G.edges(data=True): if not np.isfinite(data.get("weight", 1.0)): G.remove_edge(u, v) # Configurar e treinar o modelo Node2Vec node2vec = Node2Vec(G, dimensions=64, walk_length=30, num_walks=200, workers=4) model = node2vec.fit(window=10, min_count=1) # Obter os embeddings dos nós node_embeddings = {node: model.wv[node] for node in G.nodes} embeddings.append(node_embeddings) # Converter os embeddings em um dataframe df_embeddings = pd.DataFrame(embeddings) # Separar os dados em recursos (X) e rótulos (y) X = df_embeddings.values y = df['Difficulty'].values # Substitua 'Difficulty' pelo nome da coluna que contém os rótulos de dificuldade # Dividir os dados em conjuntos de treinamento e teste X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2, random_state=42) # Treinar o classificador Random Forest clf = RandomForestClassifier() clf.fit(X_train, y_train) # Fazer previsões y_pred = clf.predict(X_test) # Avaliar o desempenho do classificador accuracy = accuracy_score(y_test, y_pred) print("Acurácia do Random Forest:", accuracy)
调试与修复建议
- 提前校验边权重,避免无效边创建:当前在图构建后才检查边权重,建议在添加每条边前就验证权重是否为有限值,若无效则用默认值替代或跳过添加:
# 示例:添加边前校验权重 if np.isfinite(row['Passing distance']): G.add_edge(passer, receiver, weight=row['Passing distance']) else: G.add_edge(passer, receiver, weight=1.0) # 用默认有效值替代
- 禁用Node2Vec的权重加权测试:Node2Vec默认使用边的
weight属性计算随机游走概率,如果权重总和非有限会触发错误。可以先禁用权重验证问题根源:
node2vec = Node2Vec(G, dimensions=64, walk_length=30, num_walks=200, workers=4, weight_key=None)
若禁用后错误消失,说明需要彻底清理边权重数据。
- 校验嵌入向量的有效性:Node2Vec生成的嵌入可能包含非有限值,需在收集嵌入后做清理:
# 遍历嵌入并替换非有限值 for idx, emb_dict in enumerate(embeddings): for node, vec in emb_dict.items(): emb_dict[node] = np.nan_to_num(vec, nan=0.0, posinf=0.0, neginf=0.0)
- 清理Random Forest的特征矩阵:转换后的特征矩阵
X可能残留非有限值,训练前做最后一次清理:
X = np.nan_to_num(df_embeddings.values, nan=0.0, posinf=0.0, neginf=0.0)
- 检查图的连通性:如果单个图存在多个孤立节点(连通分量数>1),Node2Vec的随机游走无法正常进行,会导致权重相关错误。可以添加连通性检查:
for idx, G in enumerate(grafos_passes): comp_count = nx.number_connected_components(G) if comp_count != 1: print(f"图{idx}包含{comp_count}个不连通分量") # 可添加临时边连接孤立节点,或跳过该图
内容的提问来源于stack exchange,提问作者Val Secco
相关产品推荐
相关产品推荐

