| 63 | |
| 64 | class ADAIN_Encoder(nn.Module): |
| 65 | def __init__(self, encoder, gpu_ids=[]): |
| 66 | super(ADAIN_Encoder, self).__init__() |
| 67 | enc_layers = list(encoder.children()) |
| 68 | self.enc_1 = nn.Sequential(*enc_layers[:4]) # input -> relu1_1 64 |
| 69 | self.enc_2 = nn.Sequential(*enc_layers[4:11]) # relu1_1 -> relu2_1 128 |
| 70 | self.enc_3 = nn.Sequential(*enc_layers[11:18]) # relu2_1 -> relu3_1 256 |
| 71 | self.enc_4 = nn.Sequential(*enc_layers[18:31]) # relu3_1 -> relu4_1 512 |
| 72 | |
| 73 | self.mse_loss = nn.MSELoss() |
| 74 | |
| 75 | # fix the encoder |
| 76 | for name in ['enc_1', 'enc_2', 'enc_3', 'enc_4']: |
| 77 | for param in getattr(self, name).parameters(): |
| 78 | param.requires_grad = False |
| 79 | |
| 80 | # extract relu1_1, relu2_1, relu3_1, relu4_1 from input image |
| 81 | def encode_with_intermediate(self, input): |