| 8 | |
| 9 | |
| 10 | class EMCADNet(nn.Module): |
| 11 | def __init__(self, num_classes=1, kernel_sizes=[1,3,5], expansion_factor=2, dw_parallel=True, add=True, lgag_ks=3, activation='relu', encoder='pvt_v2_b2', pretrain=True, pretrained_dir='./pretrained_pth/pvt/'): |
| 12 | super(EMCADNet, self).__init__() |
| 13 | |
| 14 | # conv block to convert single channel to 3 channels |
| 15 | self.conv = nn.Sequential( |
| 16 | nn.Conv2d(1, 3, kernel_size=1), |
| 17 | nn.BatchNorm2d(3), |
| 18 | nn.ReLU(inplace=True) |
| 19 | ) |
| 20 | |
| 21 | # backbone network initialization with pretrained weight |
| 22 | if encoder == 'pvt_v2_b0': |
| 23 | self.backbone = pvt_v2_b0() |
| 24 | path = pretrained_dir + '/pvt_v2_b0.pth' |
| 25 | channels=[256, 160, 64, 32] |
| 26 | elif encoder == 'pvt_v2_b1': |
| 27 | self.backbone = pvt_v2_b1() |
| 28 | path = pretrained_dir + '/pvt_v2_b1.pth' |
| 29 | channels=[512, 320, 128, 64] |
| 30 | elif encoder == 'pvt_v2_b2': |
| 31 | self.backbone = pvt_v2_b2() |
| 32 | path = pretrained_dir + '/pvt_v2_b2.pth' |
| 33 | channels=[512, 320, 128, 64] |
| 34 | elif encoder == 'pvt_v2_b3': |
| 35 | self.backbone = pvt_v2_b3() |
| 36 | path = pretrained_dir + '/pvt_v2_b3.pth' |
| 37 | channels=[512, 320, 128, 64] |
| 38 | elif encoder == 'pvt_v2_b4': |
| 39 | self.backbone = pvt_v2_b4() |
| 40 | path = pretrained_dir + '/pvt_v2_b4.pth' |
| 41 | channels=[512, 320, 128, 64] |
| 42 | elif encoder == 'pvt_v2_b5': |
| 43 | self.backbone = pvt_v2_b5() |
| 44 | path = pretrained_dir + '/pvt_v2_b5.pth' |
| 45 | channels=[512, 320, 128, 64] |
| 46 | elif encoder == 'resnet18': |
| 47 | self.backbone = resnet18(pretrained=pretrain) |
| 48 | channels=[512, 256, 128, 64] |
| 49 | elif encoder == 'resnet34': |
| 50 | self.backbone = resnet34(pretrained=pretrain) |
| 51 | channels=[512, 256, 128, 64] |
| 52 | elif encoder == 'resnet50': |
| 53 | self.backbone = resnet50(pretrained=pretrain) |
| 54 | channels=[2048, 1024, 512, 256] |
| 55 | elif encoder == 'resnet101': |
| 56 | self.backbone = resnet101(pretrained=pretrain) |
| 57 | channels=[2048, 1024, 512, 256] |
| 58 | elif encoder == 'resnet152': |
| 59 | self.backbone = resnet152(pretrained=pretrain) |
| 60 | channels=[2048, 1024, 512, 256] |
| 61 | else: |
| 62 | print('Encoder not implemented! Continuing with default encoder pvt_v2_b2.') |
| 63 | self.backbone = pvt_v2_b2() |
| 64 | path = pretrained_dir + '/pvt_v2_b2.pth' |
| 65 | channels=[512, 320, 128, 64] |
| 66 | |
| 67 | if pretrain==True and 'pvt_v2' in encoder: |
no outgoing calls
no test coverage detected