MCPcopy Create free account
hub / github.com/SLDGroup/EMCAD / EMCADNet

Class EMCADNet

lib/networks.py:10–115  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

8
9
10class 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:

Callers 5

train_synapse.pyFile · 0.90
train_polyp.pyFile · 0.90
test_synapse.pyFile · 0.90
test_polyp.pyFile · 0.90
networks.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected