【Bug已解决】Freeze certain layers of an existing model in PyTorch 解决方案
问题描述
在迁移学习和微调中,经常需要冻结预训练模型的部分层,只训练特定层。开发者常遇到以下问题:
- 设置
requires_grad=False后,优化器仍为冻结参数维护状态 - 冻结 BatchNorm 层后 running 统计量仍在更新
- 不知道如何精确选择性地冻结不同层
- 冻结后显存没有明显减少
错误复现
场景一:优化器包含冻结参数
import torch import torch.nn as nn model = nn.Sequential(nn.Linear(10, 20), nn.ReLU(), nn.Linear(20, 5)) # 冻结第一层 for param in model[0].parameters(): param.requires_grad = False # 错误:优化器包含所有参数(含冻结的) optimizer = torch.optim.Adam(model.parameters(), lr=0.01) # 优化器仍为冻结参数分配 Adam 状态(动量等),浪费内存场景二:BatchNorm 冻结不完整
model = nn.Sequential( nn.Conv2d(3, 64, 3), nn.BatchNorm2d(64), nn.ReLU(), nn.Conv2d(64, 10, 3), nn.BatchNorm2d(10) ) # 冻结所有参数 for param in model.parameters(): param.requires_grad = False # 但忘记 eval(),BN 的 running_mean/var 仍在更新 model.train() # BN 仍会更新 running 统计量根因分析
1. requires_grad 的作用
requires_grad=False阻止梯度计算,参数不会被更新。但优化器如果在创建时包含了这些参数,仍会维护优化器状态(如 Adam 的动量向量),浪费内存。
2. BatchNorm 的特殊性
BatchNorm 有可训练参数(weight、bias)和非参数状态(running_mean、running_var):
train()模式:用当前 batch 更新 running 统计量eval()模式:使用固定的 running 统计量- 只冻结参数不调用
eval(),running 统计量仍会变化
3. 冻结层不减少前向传播开销
冻结只影响反向传播(不计算梯度),前向传播仍然执行所有计算。显存减少主要来自不保存中间激活值用于反向传播。
解决方案
方案一:冻结指定层并正确配置优化器
import torch import torch.nn as nn model = nn.Sequential(nn.Linear(10, 20), nn.ReLU(), nn.Linear(20, 5)) # 冻结第一层 for param in model[0].parameters(): param.requires_grad = False # 正确:只将可训练参数传入优化器 trainable_params = [p for p in model.parameters() if p.requires_grad] optimizer = torch.optim.Adam(trainable_params, lr=0.01)方案二:按名称冻结
from torchvision import models model = models.resnet18(pretrained=True) # 只训练最后一层 fc for name, param in model.named_parameters(): param.requires_grad = ('fc' in name)方案三:正确冻结 BatchNorm
def freeze_bn(model): for module in model.modules(): if isinstance(module, (nn.BatchNorm1d, nn.BatchNorm2d, nn.BatchNorm3d)): for param in module.parameters(): param.requires_grad = False module.eval() # 关键:停止更新 running 统计量方案四:渐进式解冻
# 阶段1:只训练 fc for name, param in model.named_parameters(): param.requires_grad = ('fc' in name) optimizer = torch.optim.Adam([p for p in model.parameters() if p.requires_grad], lr=0.001) # 训练若干 epoch 后... # 阶段2:解冻 layer4 for name, param in model.named_parameters(): if 'layer4' in name or 'fc' in name: param.requires_grad = True optimizer = torch.optim.Adam([p for p in model.parameters() if p.requires_grad], lr=0.0001)完整修复代码
import torch import torch.nn as nn import torch.optim as optim from torchvision import models def freeze_parameters(model, layer_names=None, freeze_bn=True): """冻结模型参数""" if layer_names is None: for param in model.parameters(): param.requires_grad = False else: for name, param in model.named_parameters(): for layer_name in layer_names: if layer_name in name: param.requires_grad = False break if freeze_bn: for module in model.modules(): if isinstance(module, (nn.BatchNorm1d, nn.BatchNorm2d, nn.BatchNorm3d)): module.eval() for param in module.parameters(): param.requires_grad = False return model def get_trainable_params(model): return [p for p in model.parameters() if p.requires_grad] def count_parameters(model): total = sum(p.numel() for p in model.parameters())  trainable = sum(p.numel() for p in model.parameters() if p.requires_grad) return total, trainable, total - trainable def freeze_bn(model): for m in model.modules(): if isinstance(m, (nn.BatchNorm1d, nn.BatchNorm2d, nn.BatchNorm3d)): m.eval() for p in m.parameters(): p.requires_grad = False def demo_basic_freeze(): print("=" * 60) print("基本冻结操作") print("=" * 60) model = nn.Sequential( nn.Linear(10, 20), nn.ReLU(), nn.Linear(20, 10), nn.ReLU(), nn.Linear(10, 5) ) for param in model[0].parameters(): param.requires_grad = False for param in model[2].parameters(): param.requires_grad = False total, trainable, frozen = count_parameters(model) print(f" 总参数: {total}, 可训练: {trainable}, 已冻结: {frozen}") optimizer = optim.Adam(get_trainable_params(model), lr=0.01) print(f" 优化器参数数量: {len(optimizer.param_groups[0]['params'])}") print() def demo_resnet_freeze(): print("=" * 60) print("ResNet18 冻结") print("=" * 60) model = models.resnet18(pretrained=True) for param in model.parameters(): param.requires_grad = False model.fc = nn.Linear(model.fc.in_features, 10) total, trainable, frozen = count_parameters(model) print(f" 总参数: {total:,}, 可训练: {trainable:,}, 已冻结: {frozen:,}") optimizer = optim.Adam(get_trainable_params(model), lr=0.001) print(f" 优化器参数数量: {len(optimizer.param_groups[0]['params'])}") print() def demo_bn_freeze(): print("=" * 60) print("BatchNorm 冻结") print("=" * 60) model = nn.Sequential( nn.Conv2d(3, 64, 3, padding=1), nn.BatchNorm2d(64), nn.ReLU(), nn.Conv2d(64, 128, 3, padding=1), nn.BatchNorm2d(128), nn.ReLU(), nn.AdaptiveAvgPool2d(1), nn.Flatten(), nn.Linear(128, 10) ) model.train() bn = model[1] initial_mean = bn.running_mean.clone() x = torch.randn(4, 3, 32, 32) y = torch.tensor([0, 1, 2, 3]) criterion = nn.CrossEntropyLoss() # 未冻结 BN optimizer = optim.Adam(model.parameters(), lr=0.01) optimizer.zero_grad() loss = criterion(model(x), y) loss.backward() optimizer.step() changed = not torch.equal(initial_mean, bn.running_mean) print(f" 未冻结 BN: running_mean 变化 = {changed}") # 冻结 BN model.train() initial_mean2 = bn.running_mean.clone() freeze_bn(model) optimizer.zero_grad() loss = criterion(model(x), y) loss.backward() optimizer.step() changed2 = not torch.equal(initial_mean2, bn.running_mean) print(f" 冻结 BN (eval): running_mean 变化 = {changed2}") print() def demo_progressive_unfreezing(): print("=" * 60) print("渐进式解冻") print("=" * 60) model = models.resnet18(pretrained=True) model.fc = nn.Linear(model.fc.in_features, 10) print("\n 阶段1:只训练 fc") for name, param in model.named_parameters(): param.requires_grad = ('fc' in name) _, t1, _ = count_parameters(model) print(f" 可训练参数: {t1:,}") print("\n 阶段2:解冻 layer4 + fc") for name, param in model.named_parameters(): if 'layer4' in name or 'fc' in name: param.requires_grad = True _, t2, _ = count_parameters(model) print(f" 可训练参数: {t2:,}") print("\n 阶段3:全部解冻") for param in model.parameters(): param.requires_grad = True _, t3, _ = count_parameters(model) print(f" 可训练参数: {t3:,}") print() def verify_freeze(): print("=" * 60) print("验证冻结生效") print("=" * 60) model = nn.Sequential(nn.Linear(10, 20), nn.Linear(20, 5)) for param in model[0].parameters(): param.requires_grad = False optimizer = optim.Adam(get_trainable_params(model), lr=0.01) w0_before = model[0].weight.data.clone() w1_before = model[1].weight.data.clone() x = torch.randn(4, 10) y = torch.tensor([0, 1, 2, 3]) for _ in range(5): optimizer.zero_grad() loss = nn.CrossEntropyLoss()(model(x), y) loss.backward() optimizer.step() print(f" 冻结层权重变化: {not torch.equal(w0_before, model[0].weight.data)}") print(f" 可训练层权重变化: {not torch.equal(w1_before, model[1].weight.data)}") print(f" 冻结层梯度: {model[0].weight.grad}") print() if __name__ == '__main__': demo_basic_freeze() demo_resnet_freeze() demo_bn_freeze() demo_progressive_unfreezing() verify_freeze() print("=" * 60) print("关键总结:") print("1. requires_grad=False 阻止梯度计算和参数更新") print("2. 只将可训练参数传入优化器") print("3. 冻结 BatchNorm 时必须同时调用 .eval()") print("4. 渐进式解冻:先训练最后一层,再逐步解冻")常见陷阱与注意事项
1. 优化器包含冻结参数
# 错误 optimizer = optim.Adam(model.parameters()) # 正确 optimizer = optim.Adam([p for p in model.parameters() if p.requires_grad])2. BatchNorm 忘记 eval
# 冻结 BN 参数后必须 eval(),否则 running 统计量仍在更新 for m in model.modules(): if isinstance(m, nn.BatchNorm2d): for p in m.parameters(): p.requires_grad = False m.eval() # 关键!3. 冻结后调用 model.train() 重置 BN
freeze_bn(model) # 后续调用 model.train() 会把 BN 切回 train 模式! model.train() # BN 又会更新 running 统计量 # 解决:freeze_bn 在 train() 之后调用4. 替换最后一层后默认可训练
model = models.resnet18(pretrained=True) for param in model.parameters(): param.requires_grad = False model.fc = nn.Linear(512, 10) # 新层默认 requires_grad=True5. 使用 torch.no_grad() 进一步节省内存
# 对冻结的部分使用 no_grad 上下文 with torch.no_grad(): features = backbone(x) # 冻结的 backbone output = classifier(features) # 可训练的分类头6. 检查冻结是否生效
# 打印各层的 requires_grad 状态 for name, param in model.named_parameters(): print(f"{name}: requires_grad={param.requires_grad}")7. 差分学习率
# 冻结层用更小的学习率,新层用更大的学习率 optimizer = optim.Adam([ {'params': backbone_params, 'lr': 0.0001}, {'params': new_params, 'lr': 0.001} ])总结
冻结 PyTorch 模型层的关键要点:
- 设置 requires_grad=False:阻止梯度计算和参数更新
- 优化器只含可训练参数:
[p for p in model.parameters() if p.requires_grad],避免浪费内存 - BatchNorm 冻结要 eval():只冻结参数不够,必须调用
.eval()停止更新 running 统计量 - 注意 train() 的调用顺序:
model.train()会重置 BN 到训练模式,冻结 BN 的操作要在train()之后 - 渐进式微调:先冻结 backbone 只训练分类头,再逐步解冻深层网络
- 差分学习率:预训练层用小 lr,新层用大 lr,通过参数组实现
- 验证冻结生效:训练前后对比冻结层权重,确认没有变化
核心原则:冻结不仅仅是设置 requires_grad=False,还需要正确配置优化器、处理 BatchNorm 的 eval 模式、注意 train/eval 切换顺序。