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

无监督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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 21:38:08