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

如何让两个不同模块中结构相同的Encoder Block共享权重?

实现两个模型的Encoder Block权重共享

在PyTorch中,要实现不同模块内的Encoder Block权重共享,核心是让两个模型的encoder属性指向同一个nn.Module实例——因为PyTorch的参数和缓冲区(比如BatchNorm的running均值/方差)都存储在模块实例中,多个引用指向同一对象时,就会共享权重,训练时参数更新也会同步。

方法一:提前创建共享Encoder实例,传入模型构造函数

先定义好共享的Encoder模块,再将其作为参数传给Model1和Model2的构造函数,确保两个模型的encoder引用同一实例:

import torch
import torch.nn as nn

class Model1(nn.Module):
    def __init__(self, encoder):
        super().__init__()
        self.encoder = encoder  # 直接使用传入的共享encoder

    def forward(self, x):
        x = self.encoder(x)
        return x

class Model2(nn.Module):
    def __init__(self, encoder):
        super().__init__()
        self.encoder = encoder  # 直接使用传入的共享encoder

    def forward(self, x):
        x = self.encoder(x)
        return x

# 创建共享的Encoder Block实例
shared_encoder = nn.Sequential(
    nn.Conv2d(3, 64, kernel_size=3, padding=1),
    nn.BatchNorm2d(64),
    nn.ReLU(inplace=True),
)

# 实例化两个模型,传入同一个encoder
model1 = Model1(shared_encoder)
model2 = Model2(shared_encoder)

方法二:实例化模型后,手动替换encoder引用

如果不想修改模型的构造函数,可以在实例化两个模型后,将其中一个的encoder赋值给另一个,强制让它们指向同一实例:

import torch
import torch.nn as nn

class Model1(nn.Module):
    def __init__(self, n_channels=3):
        super().__init__()
        self.encoder = nn.Sequential(
            nn.Conv2d(n_channels, 64, kernel_size=3, padding=1),
            nn.BatchNorm2d(64),
            nn.ReLU(inplace=True),
        )

    def forward(self, x):
        x = self.encoder(x)
        return x

class Model2(nn.Module):
    def __init__(self, n_channels=3):
        super().__init__()
        self.encoder = nn.Sequential(
            nn.Conv2d(n_channels, 64, kernel_size=3, padding=1),
            nn.BatchNorm2d(64),
            nn.ReLU(inplace=True),
        )

    def forward(self, x):
        x = self.encoder(x)
        return x

# 实例化两个模型
model1 = Model1()
model2 = Model2()

# 将model2的encoder替换为model1的encoder实例,实现权重共享
model2.encoder = model1.encoder

验证权重共享

可以通过检查参数的内存地址是否一致,或者观察梯度更新是否同步来验证:

# 检查Conv层权重是否为同一张量
print(model1.encoder[0].weight is model2.encoder[0].weight)  # 输出True,说明权重共享

# 测试梯度同步更新
x1 = torch.randn(1, 3, 224, 224)
x2 = torch.randn(1, 3, 224, 224)

# 模型1前向传播+反向传播
out1 = model1(x1)
loss1 = out1.sum()
loss1.backward()
grad_sum1 = model1.encoder[0].weight.grad.sum().item()

# 模型2前向传播+反向传播
out2 = model2(x2)
loss2 = out2.sum()
loss2.backward()
grad_sum2 = model1.encoder[0].weight.grad.sum().item()

print(f"第一次梯度和: {grad_sum1}, 第二次梯度和: {grad_sum2}")
# 输出的第二次梯度和是第一次的两倍,说明两个模型的梯度累加在同一参数上

内容的提问来源于stack exchange,提问作者Sora

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 03:38:21