class ResNet101_UNet(nn.Module): def init(self, in_channels, num_classes, pretrained = False, bilinear = True, base_c = 64): super(ResNet101_UNet, self).init()

    # self.in_channels = in_channels
    # self.num_classes = num_classes
    # self.bilinear = bilinear

    self.resnet = models.resnet50(pretrained = pretrained)
    # resnet最后一层avgpool用卷积层代替
    self.avgpool = nn.Conv2d(2048, 1024, kernel_size=1)

    self.layer0 = nn.Sequential(
        self.resnet.conv1,
        self.resnet.bn1,
        self.resnet.relu,
        self.resnet.maxpool
    )
    # encode
    self.encoder1 = self.resnet.layer1
    self.encoder2 = self.resnet.layer2
    self.encoder3 = self.resnet.layer3
    self.encoder4 = self.resnet.layer4


    # decode
    self.decoder4 = DecoderBlock(2048 + 1024, 1024, base_c * 8)
    self.decoder3 = DecoderBlock(1024 + 512, 512, base_c * 4)
    self.decoder2 = DecoderBlock(512 + 256, 256, base_c * 2)
    self.decoder1 = DecoderBlock(256 + 64, 64, base_c)

    # final conv (without any concat)
    self.final_conv = nn.Conv2d(base_c, num_classes, kernel_size=1)

    # bilinear upsampling
    if bilinear:
        self.up = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)
    else:
        self.up = nn.ConvTranspose2d(base_c, base_c, kernel_size=2, stride=2)

def forward(self, x):

    # encode
    x = self.layer0(x)
    e1 = self.encoder1(x)
    e2 = self.encoder2(e1)
    e3 = self.encoder3(e2)
    e4 = self.encoder4(e3)

    # resnet最后一层avgpool用卷积层代替
    e5 = self.avgpool(e4)

    # decode
    d4 = self.decoder4(torch.cat([e5, e4], 1))
    d3 = self.decoder3(torch.cat([d4, e3], 1))
    d2 = self.decoder2(torch.cat([d3, e2], 1))
    d1 = self.decoder1(torch.cat([d2, e1], 1))

    # final conv
    out = self.final_conv(d1)

    # bilinear upsampling
    out = self.up(out)

    return out
ResNet101_UNet: A Deep Learning Model for Image Segmentation

原文地址: https://www.cveoy.top/t/topic/nXMM 著作权归作者所有。请勿转载和采集!

免费AI点我,无需注册和登录