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

如何让PyTorch与Torch-Geometric复合模型的GinEncoder权重更新?

问题:GinEncoder权重无法随反向传播更新的解决方法

我构建了一个包含GinEncoder(基于torch-geometric实现)和Linear层的复合模型MainModel,训练时发现Linear层权重能正常更新,但GinEncoder的权重始终没有变化。

模型代码

import torch
from torch_geometric.nn import GINConv, global_add_pool
from torch.nn import Linear, BatchNorm1d, ReLU, Sequential

class GinEncoder(torch.nn.Module):
    def __init__(self):
        super(GinEncoder, self).__init__()
        self.gin_convs = torch.nn.ModuleList()
        self.gin_convs.append(GINConv(Sequential(Linear(1, 4),
                                                 BatchNorm1d(4), ReLU(),
                                                 Linear(4, 4), ReLU())))
        self.gin_convs.append(GINConv(Sequential(Linear(4, 4),
                                                 BatchNorm1d(4), ReLU(),
                                                 Linear(4, 4), ReLU())))

    def forward(self, x, edge_index, batch_node_id):
        # Node embeddings
        nodes_emb_layers = []
        for i in range(2):
            x = self.gin_convs[i](x, edge_index)
            nodes_emb_layers.append(x)

        # Graph-level readout
        nodes_emb_pools = [global_add_pool(nodes_emb, batch_node_id) for nodes_emb in nodes_emb_layers]

        # Concatenate and form the graph embeddings
        graph_embeds = torch.cat(nodes_emb_pools, dim=1)
        return graph_embeds

    def get_embeddings(self, x, edge_index, batch_node_id):
        with torch.no_grad():
            graph_embeds = self.forward(x, edge_index, batch_node_id).reshape(-1)

        return graph_embeds

class MainModel(torch.nn.Module):
    def __init__(self, graph_encoder:torch.nn.Module):
        super(MainModel, self).__init__()
        self.graph_encoder = graph_encoder
        self.lin1 = Linear(8, 4)
        self.lin2 = Linear(4, 8)

    def forward(self, x, edge_index, batch_node_id):
        graph_embeds = self.graph_encoder(x, edge_index, batch_node_id)
        out_lin1 = self.lin1(graph_embeds)
        pred = self.lin2(out_lin1)[-1]

        return pred

gin_encoder = GinEncoder().to("cuda")
model =  MainModel(gin_encoder).to("cuda")

训练代码及观测结果

我用以下训练代码观测权重变化:

from itertools import islice

gin_encoder = GinEncoder().to("cuda")
model =  MainModel(gin_encoder).to("cuda")
criterion = torch.nn.MSELoss()
optimizer = torch.optim.Adam(model.parameters())
epochs = 3

for epoch_i in range(epochs):
    model.train()
    train_loss = 0

    for batch_i, data in enumerate(train_loader):
        data.to("cuda")
        x, x_edge_index, x_batch_node_id = data.x, data.edge_index, data.batch
        y, y_edge_index, y_batch_node_id = data.y[-1].x, data.y[-1].edge_index, torch.zeros(data.y[-1].x.shape[0], dtype=torch.int64).to("cuda")
        optimizer.zero_grad()
        graph_embeds_pred = model(x, x_edge_index, x_batch_node_id)
        y_graph_embeds = model.graph_encoder.get_embeddings(y, y_edge_index, y_batch_node_id)
        loss =  criterion(graph_embeds_pred, y_graph_embeds)
        train_loss += loss
        loss.backward()
        optimizer.step()
        if batch_i == 0:
            print(f"NO. {epoch_i} EPOCH")
            print(f"MainModel weights in epoch_{epoch_i}_batch0:{next(islice(model.parameters(), 15, 16))}", end="\n\n")
            print(f"GinEncoder weights in epoch_{epoch_i}_batch0:{next(model.graph_encoder.parameters())}")
            print("*"*80)

输出结果显示Linear层权重持续变化,但GinEncoder权重完全不变:

NO. 0 EPOCH
MainModel weights in epoch_0_batch0:Parameter containing:
tensor([-0.1447, -0.3689, -0.2840, -0.3619, -0.2040,  0.2430,  0.4651,  0.3736],
       device='cuda:0', requires_grad=True)

GinEncoder weights in epoch_0_batch0:Parameter containing:
tensor([[-0.8312],
        [-0.5712],
        [-0.6963],
        [-0.1601]], device='cuda:0', requires_grad=True)
********************************************************************************
NO. 1 EPOCH
MainModel weights in epoch_1_batch0:Parameter containing:
tensor([-0.1842, -0.3333, -0.3170, -0.3247, -0.2424,  0.2627,  0.4272,  0.4119],
       device='cuda:0', requires_grad=True)

GinEncoder weights in epoch_1_batch0:Parameter containing:
tensor([[-0.8312],
        [-0.5712],
        [-0.6963],
        [-0.1601]], device='cuda:0', requires_grad=True)
********************************************************************************
NO. 2 EPOCH
MainModel weights in epoch_2_batch0:Parameter containing:
tensor([-0.2302, -0.3077, -0.3251, -0.2905, -0.2847,  0.2558,  0.3881,  0.4527],
       device='cuda:0', requires_grad=True)

GinEncoder weights in epoch_2_batch0:Parameter containing:
tensor([[-0.8312],
        [-0.5712],
        [-0.6963],
        [-0.1601]], device='cuda:0', requires_grad=True)
********************************************************************************

解决方法

1. 修正数据设备迁移

torch-geometric的Data对象to()方法默认不是原地操作,需要重新赋值确保所有张量都迁移到GPU:

data = data.to("cuda")  # 替换原有的data.to("cuda")

2. 修改模型Forward逻辑,避免丢弃梯度信号

当前模型只取输出的最后一个元素,会丢失批量中其他样本的梯度信息,甚至可能导致梯度回传路径异常。修改MainModel的forward方法,返回整个批量的输出:

def forward(self, x, edge_index, batch_node_id):
    graph_embeds = self.graph_encoder(x, edge_index, batch_node_id)
    out_lin1 = self.lin1(graph_embeds)
    pred = self.lin2(out_lin1)  # 移除[-1],返回完整批量输出
    return pred

同时需要调整标签y_graph_embeds的形状,使其与pred维度匹配(比如批量维度一致),确保loss计算合理。

3. 确认optimizer包含GinEncoder参数

虽然model.parameters()理论上包含子模块GinEncoder的参数,但可以手动验证:

# 训练前打印参数名称,确认GinEncoder参数被包含
for name, param in model.named_parameters():
    print(name, param.shape)

如果需要,也可以手动指定参数组,确保GinEncoder参数被优化:

optimizer = torch.optim.Adam([
    {'params': model.graph_encoder.parameters()},
    {'params': model.lin1.parameters()},
    {'params': model.lin2.parameters()}
], lr=1e-3)

4. 检查梯度回传是否正常

在loss.backward()后,打印GinEncoder参数的梯度,确认梯度存在:

loss.backward()
# 打印GinEncoder第一个卷积层的权重梯度
print("GinEncoder grad:", model.graph_encoder.gin_convs[0].lin_layers[0].weight.grad)

如果梯度为None,说明梯度回传路径被阻断;如果为0,需要检查loss计算逻辑是否合理。


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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 17:01:13