| 54 | |
| 55 | |
| 56 | class Encoder4Editing(nn.Module): |
| 57 | def __init__(self, num_layers, mode='ir', stylegan_size=1024): |
| 58 | super(Encoder4Editing, self).__init__() |
| 59 | assert num_layers in [50, 100, 152], 'num_layers should be 50,100, or 152' |
| 60 | assert mode in ['ir', 'ir_se'], 'mode should be ir or ir_se' |
| 61 | blocks = get_blocks(num_layers) |
| 62 | if mode == 'ir': |
| 63 | unit_module = bottleneck_IR |
| 64 | elif mode == 'ir_se': |
| 65 | unit_module = bottleneck_IR_SE |
| 66 | self.input_layer = nn.Sequential(nn.Conv2d(3, 64, (3, 3), 1, 1, bias=False), |
| 67 | nn.BatchNorm2d(64), |
| 68 | nn.PReLU(64)) |
| 69 | modules = [] |
| 70 | for block in blocks: |
| 71 | for bottleneck in block: |
| 72 | modules.append(unit_module(bottleneck.in_channel, |
| 73 | bottleneck.depth, |
| 74 | bottleneck.stride)) |
| 75 | self.body = nn.Sequential(*modules) |
| 76 | |
| 77 | self.styles = nn.ModuleList() |
| 78 | log_size = int(math.log(stylegan_size, 2)) |
| 79 | self.style_count = 2 * log_size - 2 |
| 80 | self.coarse_ind = 3 |
| 81 | self.middle_ind = 7 |
| 82 | |
| 83 | for i in range(self.style_count): |
| 84 | if i < self.coarse_ind: |
| 85 | style = GradualStyleBlock(512, 512, 16) |
| 86 | elif i < self.middle_ind: |
| 87 | style = GradualStyleBlock(512, 512, 32) |
| 88 | else: |
| 89 | style = GradualStyleBlock(512, 512, 64) |
| 90 | self.styles.append(style) |
| 91 | |
| 92 | self.latlayer1 = nn.Conv2d(256, 512, kernel_size=1, stride=1, padding=0) |
| 93 | self.latlayer2 = nn.Conv2d(128, 512, kernel_size=1, stride=1, padding=0) |
| 94 | |
| 95 | self.progressive_stage = ProgressiveStage.Inference |
| 96 | |
| 97 | def get_deltas_starting_dimensions(self): |
| 98 | """ Get a list of the initial dimension of every delta from which it is applied """ |
| 99 | return list(range(self.style_count)) # Each dimension has a delta applied to it |
| 100 | |
| 101 | def set_progressive_stage(self, new_stage: ProgressiveStage): |
| 102 | self.progressive_stage = new_stage |
| 103 | print('Changed progressive stage to: ', new_stage) |
| 104 | |
| 105 | def forward(self, x): |
| 106 | x = self.input_layer(x) |
| 107 | |
| 108 | modulelist = list(self.body._modules.values()) |
| 109 | for i, l in enumerate(modulelist): |
| 110 | x = l(x) |
| 111 | if i == 6: |
| 112 | c1 = x |
| 113 | elif i == 20: |