| 133 | |
| 134 | |
| 135 | class E4eEncoder(nn.Module): |
| 136 | def __init__(self, latent_avg): |
| 137 | super(E4eEncoder, self).__init__() |
| 138 | self.encoder = Encoder4Editing(50, 'ir_se', stylegan_size=1024) |
| 139 | self.latent_avg = latent_avg |
| 140 | |
| 141 | def forward(self, x): |
| 142 | codes = self.encoder(x) |
| 143 | # normalize with respect to the center of an average face |
| 144 | if codes.ndim == 2: |
| 145 | w_latent = codes + self.latent_avg.repeat(codes.shape[0], 1, 1)[:, 0, :] |
| 146 | else: |
| 147 | w_latent = codes + self.latent_avg.repeat(codes.shape[0], 1, 1) |
| 148 | return w_latent |