| 30 | return out |
| 31 | |
| 32 | class 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: |