模型搭建

模型搭建

1.残差块定义

结构:
image

class Residual(nn.Module):
    #构造函数进行初始化
    def __init__(self, input_channels, mid_channels, out_channels, use_oneConv= False, strides= 1):
        super(Residual, self).__init__()
        self.ReLU = nn.ReLU()
        self.conv1 = nn.Conv2d(in_channels=input_channels,out_channels=mid_channels,kernel_size=1,padding=0)
        self.conv2 = nn.Conv2d(in_channels=mid_channels,out_channels=mid_channels,kernel_size=3,padding=1,stride=strides)
        self.conv3 = nn.Conv2d(in_channels=mid_channels,out_channels=out_channels,kernel_size=1,padding=0)
        self.bn1 = nn.BatchNorm2d(mid_channels)
        self.bn2 = nn.BatchNorm2d(mid_channels)
        self.bn3 = nn.BatchNorm2d(out_channels)
        #是否使用 1*1卷积
        if use_oneConv:
            self.conv4 = nn.Conv2d(in_channels=input_channels,out_channels=out_channels,kernel_size=1,stride=strides)
        else:
            self.conv4 = None

    def forward(self, x):
        y = self.ReLU(self.bn1(self.conv1(x)))
        y = self.ReLU(self.bn2(self.conv2(y)))
        y = self.bn3(self.conv3(y))
        if self.conv4:  # 如果 1*1卷积不为空
            x = self.conv4(x)   # 先对x进行1*1卷积运算再进行相加
        y = self.ReLU(y + x)

        return y

  • 参照结构模型,在一个残差块内通道数会有三次变化,所以要定义input_channels, mid_channels, out_channels来分别表示输入通道数,中间通道数,输出通道数。
  • use_oneConv= False这个参数来定义是否有11卷积,如果此值为true,则会定义conv4
  • 输入strides默认为1,注意,在下列卷积定义中,只有第二个卷积和第四个卷积定义了strides的赋值,因为在ResNet50结构中,只有这两个卷积会变化,故需要修改这两个卷积的步幅即可

2.搭建残差块

  • 在搭建完残差块后,使用Sequential定义块来分类搭建不同的残差块即可
class ResNet50(nn.Module):
    def __init__(self, Residual):
        super(ResNet50, self).__init__()
        self.b1 = nn.Sequential(
            #输入通道数
            nn.Conv2d(in_channels=3, out_channels=64, kernel_size=7, stride=2, padding=3),
            nn.ReLU(),
            nn.BatchNorm2d(64),
            nn.MaxPool2d(kernel_size=3, stride=2, padding=1)
        )

        self.b2 = nn.Sequential(Residual(64, 64, 256, use_oneConv=True, strides=1),
                                Residual(256, 64, 256, use_oneConv=False, strides=1),
                                Residual(256, 64, 256, use_oneConv=False, strides=1)
                                )

        self.b3 = nn.Sequential(Residual(256, 128, 512,use_oneConv=True, strides=2),
                                Residual(512, 128, 512,use_oneConv=False, strides=1),
                                Residual(512, 128, 512,use_oneConv=False, strides=1),
                                Residual(512, 128, 512,use_oneConv=False, strides=1),
                                )

        self.b4 = nn.Sequential(Residual(512,256,1024,use_oneConv=True,strides=2),
                                Residual(1024,256,1024,use_oneConv=False,strides=1),
                                Residual(1024,256,1024,use_oneConv=False,strides=1),
                                Residual(1024,256,1024,use_oneConv=False,strides=1),
                                Residual(1024,256,1024,use_oneConv=False,strides=1),
                                Residual(1024,256,1024,use_oneConv=False,strides=1),
                                )

        self.b5 = nn.Sequential(Residual(1024,512,2048,use_oneConv=True,strides=2),
                                Residual(2048,512,2048,use_oneConv=False,strides=1),
                                Residual(2048,512,2048,use_oneConv=False,strides=1),
                                )

        self.b6 = nn.Sequential(nn.AdaptiveAvgPool2d((1, 1)),
                                nn.Flatten(),
                                #分类数定义
                                nn.Linear(2048, 5)
                                )

    def forward(self,x):
        x = self.b1(x)
        x = self.b2(x)
        x = self.b3(x)
        x = self.b4(x)
        x = self.b5(x)
        x = self.b6(x)
        return x

  • Sequential实质就是把输入输出通道数相同的残差块分类到一个块,按照官方给出的模型规范进行相应的搭建即可

3.模型测试

if __name__ == "__main__":
    device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
    model = ResNet50(Residual).to(device)
    print(summary(model, (3,224,224)))
posted @ 2023-11-12 12:06  我叫仨太阳  阅读(61)  评论(0)    收藏  举报