MCPcopy Create free account
hub / github.com/AtlasAnalyticsLab/AdaFisher / DenseNet

Class DenseNet

Image_Classification/src/models/densenet_cifar.py:41–92  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

39
40
41class DenseNet(nn.Module):
42 def __init__(self, block, nblocks, growth_rate=12, reduction=0.5, num_classes=10):
43 super(DenseNet, self).__init__()
44 self.growth_rate = growth_rate
45
46 num_planes = 2*growth_rate
47 self.conv1 = nn.Conv2d(
48 3, num_planes, kernel_size=3, padding=1, bias=False)
49
50 self.dense1 = self._make_dense_layers(block, num_planes, nblocks[0])
51 num_planes += nblocks[0]*growth_rate
52 out_planes = int(math.floor(num_planes*reduction))
53 self.trans1 = Transition(num_planes, out_planes)
54 num_planes = out_planes
55
56 self.dense2 = self._make_dense_layers(block, num_planes, nblocks[1])
57 num_planes += nblocks[1]*growth_rate
58 out_planes = int(math.floor(num_planes*reduction))
59 self.trans2 = Transition(num_planes, out_planes)
60 num_planes = out_planes
61
62 self.dense3 = self._make_dense_layers(block, num_planes, nblocks[2])
63 num_planes += nblocks[2]*growth_rate
64 out_planes = int(math.floor(num_planes*reduction))
65 self.trans3 = Transition(num_planes, out_planes)
66 num_planes = out_planes
67
68 self.dense4 = self._make_dense_layers(block, num_planes, nblocks[3])
69 num_planes += nblocks[3]*growth_rate
70
71 self.bn = nn.BatchNorm2d(num_planes)
72 self.linear = nn.Linear(num_planes, num_classes)
73 self.relu = nn.ReLU(inplace=False)
74 self.avgpool = nn.AvgPool2d(4)
75
76 def _make_dense_layers(self, block, in_planes, nblock):
77 layers = []
78 for i in range(nblock):
79 layers.append(block(in_planes, self.growth_rate))
80 in_planes += self.growth_rate
81 return nn.Sequential(*layers)
82
83 def forward(self, x):
84 out = self.conv1(x)
85 out = self.trans1(self.dense1(out))
86 out = self.trans2(self.dense2(out))
87 out = self.trans3(self.dense3(out))
88 out = self.dense4(out)
89 out = self.avgpool(self.relu(self.bn(out)))
90 out = out.view(out.size(0), -1)
91 out = self.linear(out)
92 return out
93
94
95def DenseNet121(num_classes: int = 10):

Callers 4

DenseNet121Function · 0.70
DenseNet169Function · 0.70
DenseNet201Function · 0.70
DenseNet161Function · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected