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

如何在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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.19 03:13:01