如何在PyTorch中拼接两个模型并使第一个模型不可训练
在PyTorch中冻结预训练模型并拼接训练的解决方案
这事儿处理起来很直观,核心就是冻结预训练模型的参数,让它们在反向传播时不被更新。我帮你调整代码的同时,给你拆解里面的关键细节:
核心操作:冻结modelA的参数
加载预训练的modelA后,你需要遍历它的所有参数,把requires_grad属性设为False——这样反向传播时这些参数不会计算梯度,优化器自然也不会对它们进行更新。
另外提个细节:如果你的modelA里包含dropout、BatchNorm这类依赖训练/评估模式的层,训练时最好把modelA切换到评估模式(eval()),避免这些层的行为干扰预训练模型的输出(这一步和冻结参数是独立操作,主要是保证预训练模型的表现和它训练时一致)。
修改后的完整可运行代码
import torch import torch.nn as nn class MyModelA(nn.Module): def __init__(self): super(MyModelA, self).__init__() self.fc1 = nn.Linear(10, 2) def forward(self, x): x = self.fc1(x) return x class MyModelB(nn.Module): def __init__(self): super(MyModelB, self).__init__() # 这里修正了维度不匹配问题:modelA输出维度是2,所以modelB输入要对应 self.fc1 = nn.Linear(2, 20) self.fc2 = nn.Linear(20, 2) def forward(self, x): x = self.fc1(x) x = self.fc2(x) return x class MyEnsemble(nn.Module): def __init__(self, modelA, modelB): super(MyEnsemble, self).__init__() self.modelA = modelA self.modelB = modelB def forward(self, x): x1 = self.modelA(x) x2 = self.modelB(x1) return x2 # 初始化模型并加载预训练权重 modelA = MyModelA() modelB = MyModelB() # 替换为你的预训练权重路径 modelA.load_state_dict(torch.load(PATH)) # 1. 冻结modelA的所有参数,禁止梯度更新 for param in modelA.parameters(): param.requires_grad = False # 2. 如果modelA包含dropout/BatchNorm,切换到评估模式 modelA.eval() # 拼接得到完整模型 model = MyEnsemble(modelA, modelB) # 优化器仅传入需要训练的参数(modelB的参数),节省计算资源 optimizer = torch.optim.Adam(model.modelB.parameters(), lr=0.001) # 测试前向传播 x = torch.randn(1, 10) output = model(x)
额外补充说明
- 解冻参数:如果之后需要解冻
modelA的部分或全部参数,只需把对应参数的requires_grad改回True,同时将模型切换回训练模式(modelA.train()),并更新优化器的参数列表。 - 优化器参数筛选:如果之后有部分解冻的场景,也可以用
filter(lambda p: p.requires_grad, model.parameters())自动筛选所有需要训练的参数传入优化器。 - 维度匹配:你原来的代码里
modelA输出维度是2,modelB输入维度是20,会触发维度不匹配错误,我已经帮你修正了这一点,你可以根据实际需求调整维度。
内容的提问来源于stack exchange,提问作者Nagabhushan S N
相关产品推荐
相关产品推荐

