如何为自定义Unet编码器复用预训练权重?以ResNet50为例
自定义Unet编码器加载ImageNet预训练权重的方案
对于你给出的基于VGGBlock的自定义Unet,完全可以复用ResNet50这类成熟预训练编码器的权重,也能针对现有结构做预训练,下面分两种实用方案说明:
方案一:替换编码器为预训练ResNet50(最便捷高效)
从头训练自定义编码器耗时极长,直接复用ResNet50的预训练权重是最优解。只需要修改Unet结构,把原有的下采样编码器部分替换为ResNet50的特征提取层,保留原Unet的上采样和解码逻辑即可:
import torch import torch.nn as nn from torchvision import models class VGGBlock(nn.Module): def __init__(self, in_channels, middle_channels, out_channels): super().__init__() self.relu = nn.ReLU(inplace=True) self.conv1 = nn.Conv2d(in_channels, middle_channels, 3, padding=1) self.bn1 = nn.BatchNorm2d(middle_channels) self.conv2 = nn.Conv2d(middle_channels, out_channels, 3, padding=1) self.bn2 = nn.BatchNorm2d(out_channels) def forward(self, x): out = self.conv1(x) out = self.bn1(out) out = self.relu(out) out = self.conv2(out) out = self.bn2(out) out = self.relu(out) return out class ResNetUNet(nn.Module): def __init__(self, num_classes, input_channels=3, pretrained=True): super().__init__() # 加载预训练ResNet50 resnet = models.resnet50(pretrained=pretrained) # 提取ResNet特征层,对应Unet下采样阶段 self.conv0_0 = nn.Sequential(resnet.conv1, resnet.bn1, resnet.relu) # 输出通道64 self.pool = resnet.maxpool # 下采样层 self.conv1_0 = resnet.layer1 # 输出通道256 self.conv2_0 = resnet.layer2 # 输出通道512 self.conv3_0 = resnet.layer3 # 输出通道1024 self.conv4_0 = resnet.layer4 # 输出通道2048 # 调整解码层输入通道,适配ResNet输出维度 self.up = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True) self.conv3_1 = VGGBlock(1024 + 2048, 1024, 1024) self.conv2_2 = VGGBlock(512 + 1024, 512, 512) self.conv1_3 = VGGBlock(256 + 512, 256, 256) self.conv0_4 = VGGBlock(64 + 256, 64, 64) self.final = nn.Conv2d(64, num_classes, kernel_size=1) def forward(self, input): # 编码器部分:ResNet预训练特征提取 x0_0 = self.conv0_0(input) x1_0 = self.conv1_0(self.pool(x0_0)) x2_0 = self.conv2_0(x1_0) x3_0 = self.conv3_0(x2_0) x4_0 = self.conv4_0(x3_0) # 解码器部分:保留原Unet拼接逻辑 x3_1 = self.conv3_1(torch.cat([x3_0, self.up(x4_0)], 1)) x2_2 = self.conv2_2(torch.cat([x2_0, self.up(x3_1)], 1)) x1_3 = self.conv1_3(torch.cat([x1_0, self.up(x2_2)], 1)) x0_4 = self.conv0_4(torch.cat([x0_0, self.up(x1_3)], 1)) output = self.final(x0_4) return output
修改后,编码器直接复用ResNet50的ImageNet预训练权重,解码层随机初始化后微调即可,训练效率和模型表现都会显著提升。
方案二:为现有自定义编码器添加ImageNet预训练(适合必须保留原结构的场景)
如果一定要保留原有的VGGBlock编码器结构,可以先在ImageNet上预训练该编码器,再将权重迁移到Unet中:
1. 拆分编码器为分类模型
把自定义Unet的下采样部分单独拆出来,添加分类头,用于ImageNet分类训练:
class EncoderClassifier(nn.Module): def __init__(self, num_classes=1000): super().__init__() nb_filter = [32, 64, 128, 256, 512] self.pool = nn.MaxPool2d(2, 2) self.conv0_0 = VGGBlock(3, nb_filter[0], nb_filter[0]) self.conv1_0 = VGGBlock(nb_filter[0], nb_filter[1], nb_filter[1]) self.conv2_0 = VGGBlock(nb_filter[1], nb_filter[2], nb_filter[2]) self.conv3_0 = VGGBlock(nb_filter[2], nb_filter[3], nb_filter[3]) self.conv4_0 = VGGBlock(nb_filter[3], nb_filter[4], nb_filter[4]) # 分类头适配ImageNet 1000类 self.global_avg_pool = nn.AdaptiveAvgPool2d(1) self.fc = nn.Linear(nb_filter[4], num_classes) def forward(self, x): x = self.conv0_0(x) x = self.pool(x) x = self.conv1_0(x) x = self.pool(x) x = self.conv2_0(x) x = self.pool(x) x = self.conv3_0(x) x = self.pool(x) x = self.conv4_0(x) x = self.global_avg_pool(x) x = torch.flatten(x, 1) x = self.fc(x) return x
2. 在ImageNet上训练分类模型
编写常规的ImageNet训练逻辑(数据加载、优化器、损失函数等),训练完成后保存权重:
classifier = EncoderClassifier() # 此处编写ImageNet训练代码(省略数据加载、循环训练等逻辑) torch.save(classifier.state_dict(), 'encoder_pretrained.pth')
3. 迁移权重到原Unet
加载预训练的编码器权重,匹配原Unet的对应层:
# 实例化原Unet unet = UNet(num_classes=你的目标类别数) # 加载分类模型权重 pretrained_weights = torch.load('encoder_pretrained.pth') # 提取编码器层权重,覆盖Unet对应参数 unet_weights = unet.state_dict() for k in unet_weights.keys(): if k in pretrained_weights and 'conv' in k: unet_weights[k] = pretrained_weights[k] unet.load_state_dict(unet_weights)
这种方法需要消耗大量资源在ImageNet上训练自定义编码器,仅在必须保留原结构时使用。
内容的提问来源于stack exchange,提问作者user836026
相关产品推荐
相关产品推荐

