ARTICLE DETAIL

资讯详情

深耕郑州网站建设与运营推广的一线实战洞察。

【Bug已解决】Freeze certain layers of an existing model in PyTorch 解决方案

【Bug已解决】Freeze certain layers of an existing model in PyTorch 解决方案 【Bug已解决】Freeze certain layers of an existing model in PyTorch 解决方案问题描述在迁移学习和微调中经常需要冻结预训练模型的部分层只训练特定层。开发者常遇到以下问题设置requires_gradFalse后优化器仍为冻结参数维护状态冻结 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(), lr0.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_gradFalse阻止梯度计算参数不会被更新。但优化器如果在创建时包含了这些参数仍会维护优化器状态如 Adam 的动量向量浪费内存。2. BatchNorm 的特殊性BatchNorm 有可训练参数weight、bias和非参数状态running_mean、running_vartrain()模式用当前 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, lr0.01)方案二按名称冻结from torchvision import models model models.resnet18(pretrainedTrue) # 只训练最后一层 fc for name, param in model.named_parameters(): param.requires_grad (fc in name)方案三正确冻结 BatchNormdef 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], lr0.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], lr0.0001)完整修复代码import torch import torch.nn as nn import torch.optim as optim from torchvision import models def freeze_parameters(model, layer_namesNone, freeze_bnTrue): 冻结模型参数 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()) ![配图](https://i-blog.csdnimg.cn/img_convert/ca7329cf6834d8b53cb16e30dc625a13.png) 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), lr0.01) print(f 优化器参数数量: {len(optimizer.param_groups[0][params])}) print() def demo_resnet_freeze(): print( * 60) print(ResNet18 冻结) print( * 60) model models.resnet18(pretrainedTrue) 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), lr0.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, padding1), nn.BatchNorm2d(64), nn.ReLU(), nn.Conv2d(64, 128, 3, padding1), 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(), lr0.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(pretrainedTrue) 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), lr0.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_gradFalse 阻止梯度计算和参数更新) 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() 重置 BNfreeze_bn(model) # 后续调用 model.train() 会把 BN 切回 train 模式 model.train() # BN 又会更新 running 统计量 # 解决freeze_bn 在 train() 之后调用4. 替换最后一层后默认可训练model models.resnet18(pretrainedTrue) for param in model.parameters(): param.requires_grad False model.fc nn.Linear(512, 10) # 新层默认 requires_gradTrue5. 使用 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_gradFalse阻止梯度计算和参数更新优化器只含可训练参数[p for p in model.parameters() if p.requires_grad]避免浪费内存BatchNorm 冻结要 eval()只冻结参数不够必须调用.eval()停止更新 running 统计量注意 train() 的调用顺序model.train()会重置 BN 到训练模式冻结 BN 的操作要在train()之后渐进式微调先冻结 backbone 只训练分类头再逐步解冻深层网络差分学习率预训练层用小 lr新层用大 lr通过参数组实现验证冻结生效训练前后对比冻结层权重确认没有变化核心原则冻结不仅仅是设置 requires_gradFalse还需要正确配置优化器、处理 BatchNorm 的 eval 模式、注意 train/eval 切换顺序。
返回列表