如何在PyTorch中对同架构的两个模型权重取平均值?
如何在PyTorch中计算两个同架构模型的权重平均值?
嗨,这个需求在模型融合场景里特别常见——不管是整合同一数据集训练的不同模型版本,还是融合不同训练策略得到的模型,权重平均都是个简单有效的提升泛化能力的方法。下面我给你一步步拆解实现过程:
前提准备:确保模型架构一致
首先必须确认model1和model2的架构完全相同(包括层的数量、参数形状、命名,甚至缓冲区如BN层的running均值/方差),否则后续操作会直接报错。
步骤1:加载并准备模型
先把两个模型的权重加载好,并且切换到评估模式(避免训练模式下的梯度计算或BN层更新干扰):
import torch from your_model_module import MyModel # 替换成你的模型类路径 # 初始化模型实例 model1 = MyModel() model2 = MyModel() # 加载预训练权重 model1.load_state_dict(torch.load("model1_weights.pth")) model2.load_state_dict(torch.load("model2_weights.pth")) # 切换到评估模式,禁用梯度计算 model1.eval() model2.eval() # 如果用GPU,统一设备(可选但推荐) device = torch.device("cuda" if torch.cuda.is_available() else "cpu") model1.to(device) model2.to(device)
步骤2:创建目标模型并计算权重平均
这里有两种常用实现方式,推荐第一种(覆盖所有参数和缓冲区):
方式1:基于State_dict的完整平均(推荐)
State_dict包含了模型所有可学习参数和缓冲区(比如BN层的running stats),能保证模型的所有状态都被平均:
# 初始化目标模型 target_model = MyModel().to(device) # 获取两个模型的state_dict state_dict1 = model1.state_dict() state_dict2 = model2.state_dict() # 计算平均后的state_dict avg_state_dict = {} for key in state_dict1.keys(): # 架构一致的情况下,两个模型的state_dict key完全匹配 avg_state_dict[key] = (state_dict1[key] + state_dict2[key]) / 2.0 # 将平均后的权重加载到目标模型 target_model.load_state_dict(avg_state_dict)
方式2:仅平均可学习参数
如果你的模型没有缓冲区(比如没有BN层),或者只需要平均可学习参数,可以直接遍历模型参数:
# 初始化目标模型 target_model = MyModel().to(device) # 遍历参数并赋值平均结果 for param_target, param1, param2 in zip(target_model.parameters(), model1.parameters(), model2.parameters()): param_target.data.copy_((param1.data + param2.data) / 2.0)
扩展:加权平均(非50:50比例)
如果想给其中一个模型更高的权重(比如给效果更好的模型分配0.7权重),只需要修改平均公式:
weight1 = 0.7 weight2 = 0.3 avg_state_dict[key] = weight1 * state_dict1[key] + weight2 * state_dict2[key]
注意事项
- 务必保证两个模型架构完全一致,否则会出现key不匹配或形状不兼容的错误。
- 如果模型在不同设备(比如一个CPU一个GPU)上,先统一到同一设备再进行平均操作。
- 平均后的模型记得保存权重:
torch.save(target_model.state_dict(), "avg_model_weights.pth")
内容的提问来源于stack exchange,提问作者patapouf_ai
相关产品推荐
相关产品推荐

