如何让两个不同模块中结构相同的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
相关产品推荐
相关产品推荐

