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

PyTorch预训练模型分层冻结:如何自动确定需解冻层?

解决SingleStageFSDV2增量训练中自动确定解冻层的问题

由于print(model)输出的模块顺序和实际前向执行顺序不匹配,要实现自动冻结backbone及之前的层,核心是先获取模块真实执行顺序,再据此划分冻结/解冻范围。

步骤1:追踪模块实际执行顺序

利用PyTorch的register_forward_hook钩子,记录前向传播时每个模块的调用顺序:

import torch
from collections import OrderedDict

# 初始化并加载预训练的SingleStageFSDV2模型
model = ...  # 替换为你的模型实例

# 存储模块执行顺序:key=模块实例,value=执行顺序索引
execution_order = OrderedDict()
order_counter = 0

def forward_hook(module, input, output):
    global order_counter
    # 仅记录顶层模块(避免子模块重复,可根据需求调整)
    if module not in execution_order:
        execution_order[module] = order_counter
        order_counter += 1

# 给所有模块注册前向钩子
for _, module in model.named_modules():
    module.register_forward_hook(forward_hook)

# 用dummy输入跑一次前向传播,触发钩子记录顺序
dummy_input = torch.randn(1, 3, 640, 640)  # 匹配模型输入尺寸
model(dummy_input)

# 按执行顺序排序并打印
sorted_modules = sorted(execution_order.items(), key=lambda x: x[1])
print("模块实际执行顺序:")
for idx, (module, _) in enumerate(sorted_modules):
    print(f"{idx}: {module.__class__.__name__}")

步骤2:定位backbone模块的执行位置

找到model.backbone在执行顺序中的索引:

backbone_module = model.backbone
backbone_order_idx = execution_order.get(backbone_module, -1)
if backbone_order_idx == -1:
    raise ValueError("未在执行顺序中找到backbone模块,请检查钩子注册逻辑")

步骤3:自动冻结/解冻模块

根据执行顺序索引,冻结backbone及之前的所有层,解冻之后的层:

for name, module in model.named_modules():
    # 跳过模型本身
    if module is model:
        continue
    module_idx = execution_order.get(module, -1)
    if module_idx != -1:
        if module_idx <= backbone_order_idx:
            # 冻结:关闭参数梯度
            for param in module.parameters():
                param.requires_grad = False
            print(f"已冻结模块:{name}")
        else:
            # 解冻:开启参数梯度
            for param in module.parameters():
                param.requires_grad = True
            print(f"已解冻模块:{name}")

验证与优化器配置

  • 验证参数梯度状态:
print("\n参数梯度状态验证:")
for name, param in model.named_parameters():
    print(f"{name}: requires_grad={param.requires_grad}")
  • 优化器仅传入可训练参数:
optimizer = torch.optim.Adam(
    filter(lambda p: p.requires_grad, model.parameters()),
    lr=1e-4
)

注意事项

  • 如果模型存在分支结构,钩子会如实记录分支模块的执行顺序,需根据业务需求判断是否属于“backbone之后”的可训练部分。
  • 若只需追踪顶层大模块(如backbone、neck、head),可在forward_hook中添加判断,比如仅保留name不包含.的模块。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.25 13:55:09