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

如何为自定义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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 13:56:57