MCPcopy Create free account
hub / github.com/csxmli2016/MARCONetPlusPlus / WEncoder

Class WEncoder

networks/w_encoder_arch.py:32–136  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

30 return out
31
32class WEncoder(nn.Module):
33 def __init__(self, block=BasicBlock, layers=[3, 4, 6, 6, 3], strides=[2,1,2,1,2]):
34 self.inplanes = 32
35 super(WEncoder, self).__init__()
36 self.conv1 = nn.Conv2d(3, 32, kernel_size=3, stride=1, padding=1,
37 bias=False)
38 self.relu = nn.LeakyReLU(0.2, inplace=True)
39
40 feature_out_dim = 512
41 self.layer1 = self._make_layer(block, 32, layers[0], stride=strides[0])
42 self.layer2 = self._make_layer(block, 64, layers[1], stride=strides[1])
43 self.layer3 = self._make_layer(block, 128, layers[2], stride=strides[2])
44 self.layer4 = self._make_layer(block, 256, layers[3], stride=strides[3])
45 self.layer5 = self._make_layer(block, feature_out_dim, layers[4], stride=strides[4])
46
47
48 self.down_h = 1
49 for stride in strides:
50 self.down_h *= stride
51 self.size_h = 32 // self.down_h
52
53
54 self.feature2w = nn.Sequential(
55 PixelNorm(),
56 EqualLinear(self.size_h*self.size_h*feature_out_dim, 512, bias=True, bias_init_val=0, lr_mul=1,
57 activation='fused_lrelu'),
58 EqualLinear(512, 512, bias=True, bias_init_val=0, lr_mul=1,
59 activation='fused_lrelu')
60 # EqualLinear(self.size_h*self.size_h*feature_out_dim, 512, bias=True),
61 # EqualLinear(512, 512, bias=True)
62 )
63
64 for m in self.modules():
65 if isinstance(m, nn.Conv2d):
66 n = m.kernel_size[0] * m.kernel_size[1] * m.out_channels
67 m.weight.data.normal_(0, math.sqrt(2. / n))
68
69
70 def _make_layer(self, block, planes, blocks, stride=1):
71 downsample = None
72 if stride != 1 or self.inplanes != planes:
73 downsample = nn.Sequential(
74 nn.Conv2d(self.inplanes, planes,
75 kernel_size=1, stride=stride, bias=False),
76 )
77 # GroupNorm(planes),
78
79 layers = []
80 layers.append(block(self.inplanes, planes, stride, downsample))
81 self.inplanes = planes
82 for i in range(1, blocks):
83 layers.append(block(self.inplanes, planes))
84
85 return nn.Sequential(*layers)
86
87 def _check_outliers(self, crop_feature, target_width):
88 _, _, H, W = crop_feature.size()
89 if W != target_width:

Callers 1

w_encoder_arch.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected