无监督GNN训练参数未更新、损失呈噪声状问题求助
我想要实现一个无监督GNN以完成节点标注任务,自定义了用于描述节点与其邻居取值关系的损失函数。但训练后发现损失曲线呈噪声状,网络参数完全没有更新,模型似乎未学到任何内容。相关代码如下:
import torch import torch.nn as nn import torch.nn.functional as F from torch_geometric.data import Data from torch_geometric.nn import GCNConv import networkx as nx import numpy as np from time import time import random from itertools import chain, islice def setup_seed(seed): torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) np.random.seed(seed) random.seed(seed) setup_seed(3) TORCH_DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu') class GCN_Net(torch.nn.Module): def __init__(self, num_features, num_classes, nerouns, dropout=0.1): super(GCN_Net, self).__init__() self.dropout = dropout self.conv1 = GCNConv(num_features, nerouns) self.conv2 = GCNConv(nerouns, 2*nerouns) self.linear = torch.nn.Linear(2*nerouns, num_classes) self.softmax = torch.nn.Softmax() def forward(self, data): h = self.conv1(data.x, data.edge_index) h = torch.relu(h) h = F.dropout(h, p=self.dropout) h = self.conv2(h, data.edge_index) h = self.linear(h) h = self.softmax(h) return h def quardaic_loss(a,b): return (a-b)**2 def loss_function(outputs, edges, func=quardaic_loss): loss = [] for n1, n2 in edges: loss.append((outputs[n1].item() - outputs[n2].item())**2) return torch.sum(torch.tensor(loss), dtype=float) def data_transformation(graph): nodes = list(graph.nodes()) edge_idx = [[], []] for node in nodes: for neighbor in [n for n in graph.neighbors(node)]: edge_idx[0].append(node) edge_idx[0].append(neighbor) edge_idx[1].append(neighbor) edge_idx[1].append(node) edge_index = torch.tensor(edge_idx, dtype=torch.long) return edge_index def create_net(nodes_num, gnn_hypers, opt_params, torch_device, torch_dtype): num_features = gnn_hypers['num_features'] number_classes = gnn_hypers['number_classes'] dropout = gnn_hypers['dropout'] neurons = gnn_hypers['neurons'] embed = nn.Embedding(nodes_num, num_features) embed = embed.type(torch_dtype).to(torch_device) net = GCN_Net(num_features, number_classes, neurons, dropout) net = net.type(torch_dtype).to(torch_device) params = chain(net.parameters(), embed.parameters()) optimizer = torch.optim.Adam(params, **opt_params) return net, embed, optimizer def train(graph, net, embed, num_epoch, optimizer, loss_function, device, max_state=15, tol=5, patience=4): edges = list(graph.edges()) torch.manual_seed(666) x = embed.weight edge_index = data_transformation(graph) data = Data(x=x, edge_index=edge_index.contiguous()) data = data.to(device) prevloss = len(list(graph.nodes())) * max_state**2 best_loss = len(list(graph.nodes())) * max_state**2 start_time = time() no_improve_count = 0 # print("Initial loss is {}".format(prevloss)) # print("Tranining starts...") losses = [] for epoch in range(num_epoch): out = net(data) out = torch.argmax(out, 1) + 1 loss = loss_function(out, edges) loss.requires_grad_(True) loss_ = loss.detach().item() losses.append(loss_) best_loss = min(best_loss, loss_) if abs(loss_ - prevloss) < tol or loss_ - prevloss > 0: no_improve_count += 1 else: no_improve_count = 0 if no_improve_count >= patience: print("Early stopping at epoch {} with patience {}".format(epoch, patience)) break if epoch % 100 == 0: print("The loss after epoch {} is {}".format(epoch, loss_)) prevloss = loss_ optimizer.zero_grad() loss.backward() optimizer.step() edge_index = data_transformation(graph) data = Data(x=x, edge_index=edge_index.contiguous()) data = data.to(device) print("The training time is {} for total {} epochs.".format(time() - start_time, epoch+1)) print("The final loss is {}".format(loss_)) print("The best loss is {}".format(best_loss)) return best_loss, losses nodes_num = 100 g = nx.erdos_renyi_graph(nodes_num, p=0.1) # print("The obs_0 is {}".format(observations[0])) gnn_hypers = {'num_features': 10, 'number_classes': 10, "neurons": 500, 'dropout': 0.1} opt_params = {'lr': 0.01} torch_device = TORCH_DEVICE torch_dtype = torch.float num_epoch = 100 dim_embedding = 10 net, embed, optimizer = create_net(nodes_num, gnn_hypers, opt_params, torch_device, torch_dtype) train(g, net, embed, num_epoch, optimizer, loss_function, torch_device, 10)
以下是导致模型无法训练的核心问题及对应修复:
1. 离散化操作切断梯度回传
训练中执行了out = torch.argmax(out, 1) + 1,torch.argmax是不可微分操作,会把模型输出的连续概率分布转换成离散类别标签,直接切断了梯度从损失函数向网络参数的传播路径,导致参数完全无法更新。
修复方案:
去掉离散化步骤,直接基于模型输出的概率分布计算损失,让邻居节点的概率分布尽可能相似(无监督节点标注的核心思路)。
2. 损失函数构造破坏梯度链
原损失函数中使用.item()将张量元素转为Python数值,再重新构造张量,这会彻底丢失梯度信息;同时循环遍历边的方式效率极低。
修复方案:
用张量索引直接批量计算损失,避免循环和.item()操作:
def loss_function(outputs, edge_index): # edge_index是[2, E]格式,转为[E,2] edges = edge_index.T n1_out = outputs[edges[:, 0]] n2_out = outputs[edges[:, 1]] # 用MSE损失让邻居节点的概率分布尽可能接近 return F.mse_loss(n1_out, n2_out)
3. Softmax层维度错误
原代码中self.softmax = torch.nn.Softmax()默认对所有维度做softmax,应该指定dim=1,对每个节点的类别维度进行归一化:
self.softmax = nn.Softmax(dim=1)
4. 冗余的Data对象重构
训练循环中每次都重新构造edge_index和Data对象,而图结构并未变化,这完全是冗余操作,会拖慢训练速度。
修复方案:
将edge_index和Data的构造移到训练循环外。
5. Early Stopping逻辑不合理
原逻辑中abs(loss_ - prevloss) < tol or loss_ - prevloss > 0会把损失小幅下降的情况也计入"无改进",导致模型过早停止训练。调整为只有当损失没有明显下降时才计数。
import torch import torch.nn as nn import torch.nn.functional as F from torch_geometric.data import Data from torch_geometric.nn import GCNConv import networkx as nx import numpy as np from time import time import random from itertools import chain def setup_seed(seed): torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) np.random.seed(seed) random.seed(seed) setup_seed(3) TORCH_DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu') class GCN_Net(torch.nn.Module): def __init__(self, num_features, num_classes, neurons, dropout=0.1): super(GCN_Net, self).__init__() self.dropout = dropout self.conv1 = GCNConv(num_features, neurons) self.conv2 = GCNConv(neurons, 2*neurons) self.linear = torch.nn.Linear(2*neurons, num_classes) self.softmax = nn.Softmax(dim=1) # 指定维度 def forward(self, data): h = self.conv1(data.x, data.edge_index) h = torch.relu(h) h = F.dropout(h, p=self.dropout, training=self.training) # 仅训练时dropout h = self.conv2(h, data.edge_index) h = self.linear(h) h = self.softmax(h) return h def loss_function(outputs, edge_index): # edge_index是[2, E]格式,转为[E,2] edges = edge_index.T n1_out = outputs[edges[:, 0]] n2_out = outputs[edges[:, 1]] # 用MSE损失让邻居节点的概率分布尽可能接近 return F.mse_loss(n1_out, n2_out) def data_transformation(graph): edge_idx = [[], []] for u, v in graph.edges(): edge_idx[0].append(u) edge_idx[1].append(v) edge_idx[0].append(v) edge_idx[1].append(u) edge_index = torch.tensor(edge_idx, dtype=torch.long) return edge_index def create_net(nodes_num, gnn_hypers, opt_params, torch_device, torch_dtype): num_features = gnn_hypers['num_features'] number_classes = gnn_hypers['number_classes'] dropout = gnn_hypers['dropout'] neurons = gnn_hypers['neurons'] embed = nn.Embedding(nodes_num, num_features) embed = embed.type(torch_dtype).to(torch_device) net = GCN_Net(num_features, number_classes, neurons, dropout) net = net.type(torch_dtype).to(torch_device) params = chain(net.parameters(), embed.parameters()) optimizer = torch.optim.Adam(params, **opt_params) return net, embed, optimizer def train(graph, net, embed, num_epoch, optimizer, loss_function, device, max_state=15, tol=1e-4, patience=4): torch.manual_seed(666) x = embed.weight edge_index = data_transformation(graph) data = Data(x=x, edge_index=edge_index.contiguous()) data = data.to(device) prevloss = float('inf') best_loss = float('inf') start_time = time() no_improve_count = 0 losses = [] for epoch in range(num_epoch): net.train() out = net(data) loss = loss_function(out, data.edge_index) loss_ = loss.detach().item() losses.append(loss_) best_loss = min(best_loss, loss_) # 调整Early Stopping逻辑 if loss_ >= prevloss - tol: no_improve_count += 1 else: no_improve_count = 0 if no_improve_count >= patience: print("Early stopping at epoch {} with patience {}".format(epoch, patience)) break if epoch % 10 == 0: print("The loss after epoch {} is {}".format(epoch, loss_)) prevloss = loss_ optimizer.zero_grad() loss.backward() optimizer.step() print("The training time is {:.2f}s for total {} epochs.".format(time() - start_time, epoch+1)) print("The final loss is {:.6f}".format(loss_)) print("The best loss is {:.6f}".format(best_loss)) return best_loss, losses nodes_num = 100 g = nx.erdos_renyi_graph(nodes_num, p=0.1) gnn_hypers = {'num_features': 10, 'number_classes': 10, "neurons": 500, 'dropout': 0.1} opt_params = {'lr': 0.001} # 调整学习率,避免震荡 torch_device = TORCH_DEVICE torch_dtype = torch.float num_epoch = 200 net, embed, optimizer = create_net(nodes_num, gnn_hypers, opt_params, torch_device, torch_dtype) train(g, net, embed, num_epoch, optimizer, loss_function, torch_device)
内容的提问来源于stack exchange,提问作者ppxwdy

