PyTorch与Torch-Geometric拆分模型后权重不更新问题
拆分模型权重不更新的原因及解决方法
你遇到的问题核心是拆分法中GinEncoder的参数没有被纳入优化器的更新流程,以下是具体原因分析和解决办法:
最可能的原因:优化器初始化错误
这是这类问题最常见的诱因。在拆分法中,如果你初始化优化器时只传入了MainModel的顶层参数(比如lin1和lin2的参数),而没有包含GinEncoder的参数,自然只有MainModel的权重会更新。而合并法中,所有层都属于同一个类,model.parameters()会包含所有参数,优化器能正常更新所有权重。
比如错误的优化器写法:
# 只优化MainModel的线性层,漏掉GinEncoder的参数 optimizer = torch.optim.Adam([model.lin1.parameters(), model.lin2.parameters()], lr=0.001)
正确的写法应该传入整个模型的参数:
# 包含所有子模块的参数 optimizer = torch.optim.Adam(model.parameters(), lr=0.001)
其他可能的原因
- 子模块未正确注册
PyTorch要求子模块必须作为父模块的属性(即赋值给self.xxx),才会被自动注册,其参数才会被纳入父模块的parameters()中。虽然你的代码里MainModel的__init__中self.graph_encoder = graph_encoder是正确的注册方式,但可以通过打印参数列表确认:
# 查看所有可优化的参数 for name, param in model.named_parameters(): print(name)
如果输出中没有graph_encoder.gin_convs开头的参数,说明子模块注册失败,检查是否在初始化后意外修改了self.graph_encoder。
- 梯度被主动禁用
如果在调用GinEncoder的forward时,使用了torch.no_grad()或者对输出做了detach()操作,会切断梯度回传路径,导致GinEncoder的参数无法更新。检查你的MainModel的forward函数,确保没有类似代码:
# 错误示例:禁用梯度 with torch.no_grad(): graph_embeds = self.graph_encoder(x, edge_index, batch_node_id) # 或者 graph_embeds = self.graph_encoder(x, edge_index, batch_node_id).detach()
- 模型模式设置错误
如果将GinEncoder设置为eval模式(encoder.eval()),虽然不会直接阻止权重更新,但会影响BatchNorm层的统计量更新,导致模型表现异常。训练时确保所有模块都处于train模式(model.train())。
解决步骤
- 检查优化器初始化代码,确保传入
model.parameters()而不是仅部分参数。 - 打印
model.named_parameters()确认GinEncoder的参数是否在列表中。 - 排查forward函数中是否有禁用梯度的操作。
- 训练前调用
model.train()确保所有模块处于训练模式。
内容的提问来源于stack exchange,提问作者theabc50111
相关产品推荐
相关产品推荐

