根据Resnet论文复现 Resnet50 Restnet101 Resnet152

Resnet的构成
在这里插入图片描述

代码

import torch
import torch.nn as nn
from torch.nn import (
    Module,
    Conv2d,
    MaxPool2d,
    AvgPool2d,
    Linear,
    Softmax,
    ReLU,
    BatchNorm2d,
    Sequential
)


def Conv1(in_channel=3, out_channel=64, kernel_size=(7, 7), stride=2, padding=3):
    return Sequential(
        Conv2d(in_channels=in_channel, out_channels=out_channel, kernel_size=kernel_size, stride=stride, padding=padding),
        MaxPool2d(kernel_size=(3, 3), stride=2, padding=1)
    )


class BottleNeck(Module):
    def __init__(self, in_channel, out_channel, stride=1, downsampling=False, expansion=4):
        super(BottleNeck, self).__init__()
        self.downsampling = downsampling
        self.expansion = expansion
        self.bottleneck = nn.Sequential(
            Conv2d(in_channels=in_channel,
                   out_channels=out_channel,
                   kernel_size=1
                   ),
            BatchNorm2d(out_channel),
            ReLU(),
            Conv2d(in_channels=out_channel,
                   out_channels=out_channel,
                   kernel_size=3,
                   stride=stride,
                   padding=1),
            BatchNorm2d(out_channel),
            ReLU(),
            Conv2d(in_channels=out_channel,
                   out_channels=out_channel * expansion,
                   kernel_size=1),
            BatchNorm2d(out_channel * expansion)
        )
        if downsampling:
            self.downsample = Sequential(
                Conv2d(in_channels=in_channel,
                       out_channels=out_channel*expansion,
                       kernel_size=1,
                       stride=stride),
                BatchNorm2d(out_channel * expansion)
            )
        self.relu = ReLU()

    def forward(self, x):
        residual = x
        x = self.bottleneck(x)
        if self.downsampling:
            residual = self.downsample(residual)
        x += residual
        return self.relu(x)


class Resnet(Module):
    def __init__(self, blocks, class_nums=100, downsampling=True):
        super(Resnet, self).__init__()
        self.class_nums = class_nums
        self.conv1 = Conv1()
        self.conv2 = self._make_layer(in_channel=64, out_channel=64, block=blocks[0], downsampling=downsampling, stride=1)
        self.conv3 = self._make_layer(in_channel=256, out_channel=128, block=blocks[1], downsampling=downsampling, stride=2)
        self.conv4 = self._make_layer(in_channel=512, out_channel=256, block=blocks[2], downsampling=downsampling, stride=2)
        self.conv5 =self._make_layer(in_channel=1024, out_channel=512, block=blocks[3], downsampling=downsampling, stride=2)
        self.avg_pool = AvgPool2d(kernel_size=7)
        self.fc = Linear(2048, class_nums)
        self.softmax = Softmax()

    def forward(self, x):
        x = self.conv1(x)
        x = self.conv2(x)
        x = self.conv3(x)
        x = self.conv4(x)
        x = self.conv5(x)
        x = self.avg_pool(x)
        x = x.view(x.size(0), -1)
        x = self.fc(x)
        output = self.softmax(x)
        return output

    def _make_layer(self, in_channel, out_channel, block, stride, downsampling=False, expansion=4):
        layers = []
        layers.append(BottleNeck(in_channel=in_channel, out_channel=out_channel, stride=stride, downsampling=downsampling))
        for i in range(1, block):
            layers.append(BottleNeck(in_channel=out_channel*expansion,
                                     out_channel=out_channel,
                                     downsampling=downsampling))
        return Sequential(*layers)


def get_Resnet50(blocks=[3, 4, 6, 3]):
    print("获取 Resnet50")
    return Resnet(blocks)


def get_Resnet101(blocks=[3, 4, 23, 3]):
    print("获取 Resnet101")
    return Resnet(blocks)


def get_Resnet152(blocks=[3, 8, 36, 3]):
    print("获取 Resnet152")
    return Resnet(blocks)


def main():
    pass


if __name__ == '__main__':
    main()

更多推荐