如何固定PyTorch预训练模型的running_mean、running_var与num_batches_tracked?
问题核心原因
requires_grad = False 仅会阻止反向传播更新模型的可训练参数(Parameter),而BatchNorm层的running_mean、running_var、num_batches_tracked属于模型缓冲区(Buffer),是在前向传播过程中就会自动更新的,不受requires_grad控制,这就是参数变动的根本原因。
解决方案
方法1:训练时将预训练模型固定为eval模式(最常用)
这是最简单高效的方案,BN层在eval模式下不会更新运行时统计量,会直接使用加载的固定值计算输出。你只需要在启动训练前,单独对预训练模型调用eval()方法即可,注意不要影响自定义模型的训练状态:
# 加载完预训练权重、设置requires_grad=False后执行 pretrained_model.eval() # 训练循环中每次迭代前额外确认状态,避免全局调用model.train()把预训练模型也切回训练模式 for epoch in range(epochs): for batch in dataloader: # 固定预训练模型状态 pretrained_model.eval() # 自定义模型保持训练模式 my_model.train() # 后续正常执行前向传播、损失计算、反向传播更新自定义模型即可
方法2:彻底关闭预训练模型BN层的统计量追踪
如果你需要完全杜绝预训练模型BN统计量被修改的可能,哪怕误将预训练模型切到train模式也不受影响,可以遍历预训练模型的所有BN层,关闭其track_running_stats属性:
import torch.nn as nn for module in pretrained_model.modules(): if isinstance(module, nn.BatchNorm2d): # 关闭运行时统计量追踪,不会再更新running_mean、running_var、num_batches_tracked module.track_running_stats = False
注意说明
两种方案二选一即可,无特殊需求优先选方法1,不需要改动模型结构,兼容性更好。如果你的训练逻辑中会频繁调用全局的train()/eval()切换,担心误改预训练模型的状态,可以选方法2做双重保险。
内容的提问来源于stack exchange,提问作者Nathan Wang
相关产品推荐
相关产品推荐

