PyTorch中多损失训练及按需冻结网络部分的实现与代码验证
嘿,作为刚踩过类似坑的PyTorch开发者,我来帮你把这个问题捋明白~
首先得指出你当前代码里的几个关键问题,这些是导致需求无法实现的核心:
- 直接写
requires_grad = False/True完全没用,这只是定义了一个普通变量,根本没修改模型中「红色部分」参数的梯度属性 - 优化器
opt = SGD()没有传入任何模型参数,这样opt.step()根本不会更新任何参数,这是致命错误 - 训练紫色分类器时的
opt.step少了括号,应该是opt.step() - 冻结/解冻的时机不对,应该在计算对应loss前就调整好参数的梯度状态
正确实现方式
下面是符合你需求的完整可运行代码流程,我会一步步拆解:
1. 先明确模型结构
首先你的my_model要能清晰区分出「红色部分」和两个分类器模块,比如:
import torch import torch.nn as nn import torch.optim as optim class my_model(nn.Module): def __init__(self): super().__init__() # 红色部分:需要冻结/解冻的共享特征提取层 self.red_part = nn.Sequential( nn.Linear(10, 20), nn.ReLU(), nn.Linear(20, 15) ) # 黄色分类器 self.yellow_classifier = nn.Linear(15, 2) # 假设是2分类任务 # 紫色分类器 self.purple_classifier = nn.Linear(15, 3) # 假设是3分类任务 def forward(self, x): # 红色部分提取特征 shared_features = self.red_part(x) # 两个分类器分别输出 yellow_out = self.yellow_classifier(shared_features) purple_out = self.purple_classifier(shared_features) return yellow_out, purple_out
2. 初始化核心组件
注意优化器要传入所有模型参数,后续我们通过修改requires_grad来控制哪些参数会被更新:
model = my_model() criterion = nn.CrossEntropyLoss() # 传入所有模型参数,设置学习率(这里用0.01做示例) opt = optim.SGD(model.parameters(), lr=0.01)
3. 训练黄色分类器(冻结红色部分)
这一步要先冻结红色部分的所有参数,再进行前向传播、损失计算和参数更新:
# 冻结红色部分:遍历red_part的所有参数,关闭梯度计算 for param in model.red_part.parameters(): param.requires_grad = False # 示例输入和标签(实际用你的真实数据即可) x = torch.randn(32, 10) # batch_size=32,输入维度10 y_yellow = torch.randint(0, 2, (32,)) # 黄色分类的标签 y_purple = torch.randint(0, 3, (32,)) # 紫色分类的标签 # 前向传播 yellow_out, purple_out = model(x) # 计算黄色分类器的损失并更新参数 opt.zero_grad() # 清空上一轮残留的梯度 yellow_loss = criterion(yellow_out, y_yellow) yellow_loss.backward() opt.step() # 此时只有黄色分类器的参数会被更新,红色部分参数因requires_grad=False不会更新
4. 训练紫色分类器(解冻红色部分)
切换任务前,先解冻红色部分,再进行训练:
# 解冻红色部分:遍历red_part的所有参数,开启梯度计算 for param in model.red_part.parameters(): param.requires_grad = True # 计算紫色分类器的损失并更新参数 opt.zero_grad() # 必须清空梯度,避免和上一轮黄色分类的梯度累积 purple_loss = criterion(purple_out, y_purple) purple_loss.backward() opt.step() # 此时红色部分和紫色分类器的参数都会被更新
关键细节答疑
- zero_grad的使用:你代码里
zero_grad()的时机是对的,每次反向传播前都要清空梯度,避免不同任务的梯度相互干扰。切换训练任务时,一定要重新调用opt.zero_grad(),这部分你之前的思路是对的。 - 是否遗漏内容:你遗漏了对模型具体参数的
requires_grad赋值,优化器初始化时的参数传入,以及opt.step的括号。如果是多轮训练(比如训练100轮黄色再100轮紫色),要注意每一轮训练黄色前都要重新冻结红色部分,避免之前解冻后影响。 - 是否需要额外参数:不需要特殊的额外参数,只需要通过遍历对应模块的参数修改
requires_grad属性即可。如果你的「红色部分」是多个分散的模块,可以把它们放到一个列表里统一处理,比如red_modules = [model.layer1, model.layer3],然后循环每个模块修改参数。
可选优化:分组参数管理(进阶)
如果不想每次都遍历参数修改requires_grad,也可以在初始化优化器时把参数分成两组,分别控制:
# 分组定义参数 params = [ {"params": model.red_part.parameters(), "lr": 0.01}, {"params": model.yellow_classifier.parameters(), "lr": 0.01}, {"params": model.purple_classifier.parameters(), "lr": 0.01} ] opt = optim.SGD(params, lr=0.01) # 训练黄色时,直接关闭红色部分参数组的requires_grad for group in opt.param_groups: if "red_part" in str(group["params"][0]): # 或者通过位置判断,比如第0组是red_part group["requires_grad"] = False
不过对于新手来说,前面的遍历参数方法更直观,不容易出错。
内容的提问来源于stack exchange,提问作者Austin Moon
相关产品推荐
相关产品推荐

