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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.08 09:18:15