gulabpatel/First-Order-Motion
0
1from torch import nn2import torch3import torch.nn.functional as F4from modules.util import AntiAliasInterpolation2d, make_coordinate_grid5from torchvision import models6import numpy as np7from torch.autograd import grad8 9 10class Vgg19(torch.nn.Module):11 """12 Vgg19 network for perceptual loss. See Sec 3.3.13 """14 def __init__(self, requires_grad=False):15 super(Vgg19, self).__init__()16 vgg_pretrained_features = models.vgg19(pretrained=True).features17 self.slice1 = torch.nn.Sequential()18 self.slice2 = torch.nn.Sequential()19 self.slice3 = torch.nn.Sequential()20 self.slice4 = torch.nn.Sequential()21 self.slice5 = torch.nn.Sequential()22 for x in range(2):23 self.slice1.add_module(str(x), vgg_pretrained_features[x])24 for x in range(2, 7):25 self.slice2.add_module(str(x), vgg_pretrained_features[x])26 for x in range(7, 12):27 self.slice3.add_module(str(x), vgg_pretrained_features[x])28 for x in range(12, 21):29 self.slice4.add_module(str(x), vgg_pretrained_features[x])30 for x in range(21, 30):31 self.slice5.add_module(str(x), vgg_pretrained_features[x])32 33 self.mean = torch.nn.Parameter(data=torch.Tensor(np.array([0.485, 0.456, 0.406]).reshape((1, 3, 1, 1))),34 requires_grad=False)35 self.std = torch.nn.Parameter(data=torch.Tensor(np.array([0.229, 0.224, 0.225]).reshape((1, 3, 1, 1))),36 requires_grad=False)37 38 if not requires_grad:39 for param in self.parameters():40 param.requires_grad = False41 42 def forward(self, X):43 X = (X - self.mean) / self.std44 h_relu1 = self.slice1(X)45 h_relu2 = self.slice2(h_relu1)46 h_relu3 = self.slice3(h_relu2)47 h_relu4 = self.slice4(h_relu3)48 h_relu5 = self.slice5(h_relu4)49 out = [h_relu1, h_relu2, h_relu3, h_relu4, h_relu5]50 return out51 52 53class ImagePyramide(torch.nn.Module):54 """55 Create image pyramide for computing pyramide perceptual loss. See Sec 3.356 """57 def __init__(self, scales, num_channels):58 super(ImagePyramide, self).__init__()59 downs = {}60 for scale in scales:61 downs[str(scale).replace('.', '-')] = AntiAliasInterpolation2d(num_channels, scale)62 self.downs = nn.ModuleDict(downs)63 64 def forward(self, x):65 out_dict = {}66 for scale, down_module in self.downs.items():67 out_dict['prediction_' + str(scale).replace('-', '.')] = down_module(x)68 return out_dict69 70 71class Transform:72 """73 Random tps transformation for equivariance constraints. See Sec 3.374 """75 def __init__(self, bs, **kwargs):76 noise = torch.normal(mean=0, std=kwargs['sigma_affine'] * torch.ones([bs, 2, 3]))77 self.theta = noise + torch.eye(2, 3).view(1, 2, 3)78 self.bs = bs79 80 if ('sigma_tps' in kwargs) and ('points_tps' in kwargs):81 self.tps = True82 self.control_points = make_coordinate_grid((kwargs['points_tps'], kwargs['points_tps']), type=noise.type())83 self.control_points = self.control_points.unsqueeze(0)84 self.control_params = torch.normal(mean=0,85 std=kwargs['sigma_tps'] * torch.ones([bs, 1, kwargs['points_tps'] ** 2]))86 else:87 self.tps = False88 89 def transform_frame(self, frame):90 grid = make_coordinate_grid(frame.shape[2:], type=frame.type()).unsqueeze(0)91 grid = grid.view(1, frame.shape[2] * frame.shape[3], 2)92 grid = self.warp_coordinates(grid).view(self.bs, frame.shape[2], frame.shape[3], 2)93 return F.grid_sample(frame, grid, padding_mode="reflection")94 95 def warp_coordinates(self, coordinates):96 theta = self.theta.type(coordinates.type())97 theta = theta.unsqueeze(1)98 transformed = torch.matmul(theta[:, :, :, :2], coordinates.unsqueeze(-1)) + theta[:, :, :, 2:]99 transformed = transformed.squeeze(-1)100 101 if self.tps:102 control_points = self.control_points.type(coordinates.type())103 control_params = self.control_params.type(coordinates.type())104 distances = coordinates.view(coordinates.shape[0], -1, 1, 2) - control_points.view(1, 1, -1, 2)105 distances = torch.abs(distances).sum(-1)106 107 result = distances ** 2108 result = result * torch.log(distances + 1e-6)109 result = result * control_params110 result = result.sum(dim=2).view(self.bs, coordinates.shape[1], 1)111 transformed = transformed + result112 113 return transformed114 115 def jacobian(self, coordinates):116 new_coordinates = self.warp_coordinates(coordinates)117 grad_x = grad(new_coordinates[..., 0].sum(), coordinates, create_graph=True)118 grad_y = grad(new_coordinates[..., 1].sum(), coordinates, create_graph=True)119 jacobian = torch.cat([grad_x[0].unsqueeze(-2), grad_y[0].unsqueeze(-2)], dim=-2)120 return jacobian121 122 123def detach_kp(kp):124 return {key: value.detach() for key, value in kp.items()}125 126 127class GeneratorFullModel(torch.nn.Module):128 """129 Merge all generator related updates into single model for better multi-gpu usage130 """131 132 def __init__(self, kp_extractor, generator, discriminator, train_params):133 super(GeneratorFullModel, self).__init__()134 self.kp_extractor = kp_extractor135 self.generator = generator136 self.discriminator = discriminator137 self.train_params = train_params138 self.scales = train_params['scales']139 self.disc_scales = self.discriminator.scales140 self.pyramid = ImagePyramide(self.scales, generator.num_channels)141 if torch.cuda.is_available():142 self.pyramid = self.pyramid.cuda()143 144 self.loss_weights = train_params['loss_weights']145 146 if sum(self.loss_weights['perceptual']) != 0:147 self.vgg = Vgg19()148 if torch.cuda.is_available():149 self.vgg = self.vgg.cuda()150 151 def forward(self, x):152 kp_source = self.kp_extractor(x['source'])153 kp_driving = self.kp_extractor(x['driving'])154 155 generated = self.generator(x['source'], kp_source=kp_source, kp_driving=kp_driving)156 generated.update({'kp_source': kp_source, 'kp_driving': kp_driving})157 158 loss_values = {}159 160 pyramide_real = self.pyramid(x['driving'])161 pyramide_generated = self.pyramid(generated['prediction'])162 163 if sum(self.loss_weights['perceptual']) != 0:164 value_total = 0165 for scale in self.scales:166 x_vgg = self.vgg(pyramide_generated['prediction_' + str(scale)])167 y_vgg = self.vgg(pyramide_real['prediction_' + str(scale)])168 169 for i, weight in enumerate(self.loss_weights['perceptual']):170 value = torch.abs(x_vgg[i] - y_vgg[i].detach()).mean()171 value_total += self.loss_weights['perceptual'][i] * value172 loss_values['perceptual'] = value_total173 174 if self.loss_weights['generator_gan'] != 0:175 discriminator_maps_generated = self.discriminator(pyramide_generated, kp=detach_kp(kp_driving))176 discriminator_maps_real = self.discriminator(pyramide_real, kp=detach_kp(kp_driving))177 value_total = 0178 for scale in self.disc_scales:179 key = 'prediction_map_%s' % scale180 value = ((1 - discriminator_maps_generated[key]) ** 2).mean()181 value_total += self.loss_weights['generator_gan'] * value182 loss_values['gen_gan'] = value_total183 184 if sum(self.loss_weights['feature_matching']) != 0:185 value_total = 0186 for scale in self.disc_scales:187 key = 'feature_maps_%s' % scale188 for i, (a, b) in enumerate(zip(discriminator_maps_real[key], discriminator_maps_generated[key])):189 if self.loss_weights['feature_matching'][i] == 0:190 continue191 value = torch.abs(a - b).mean()192 value_total += self.loss_weights['feature_matching'][i] * value193 loss_values['feature_matching'] = value_total194 195 if (self.loss_weights['equivariance_value'] + self.loss_weights['equivariance_jacobian']) != 0:196 transform = Transform(x['driving'].shape[0], **self.train_params['transform_params'])197 transformed_frame = transform.transform_frame(x['driving'])198 transformed_kp = self.kp_extractor(transformed_frame)199 200 generated['transformed_frame'] = transformed_frame201 generated['transformed_kp'] = transformed_kp202 203 ## Value loss part204 if self.loss_weights['equivariance_value'] != 0:205 value = torch.abs(kp_driving['value'] - transform.warp_coordinates(transformed_kp['value'])).mean()206 loss_values['equivariance_value'] = self.loss_weights['equivariance_value'] * value207 208 ## jacobian loss part209 if self.loss_weights['equivariance_jacobian'] != 0:210 jacobian_transformed = torch.matmul(transform.jacobian(transformed_kp['value']),211 transformed_kp['jacobian'])212 213 normed_driving = torch.inverse(kp_driving['jacobian'])214 normed_transformed = jacobian_transformed215 value = torch.matmul(normed_driving, normed_transformed)216 217 eye = torch.eye(2).view(1, 1, 2, 2).type(value.type())218 219 value = torch.abs(eye - value).mean()220 loss_values['equivariance_jacobian'] = self.loss_weights['equivariance_jacobian'] * value221 222 return loss_values, generated223 224 225class DiscriminatorFullModel(torch.nn.Module):226 """227 Merge all discriminator related updates into single model for better multi-gpu usage228 """229 230 def __init__(self, kp_extractor, generator, discriminator, train_params):231 super(DiscriminatorFullModel, self).__init__()232 self.kp_extractor = kp_extractor233 self.generator = generator234 self.discriminator = discriminator235 self.train_params = train_params236 self.scales = self.discriminator.scales237 self.pyramid = ImagePyramide(self.scales, generator.num_channels)238 if torch.cuda.is_available():239 self.pyramid = self.pyramid.cuda()240 241 self.loss_weights = train_params['loss_weights']242 243 def forward(self, x, generated):244 pyramide_real = self.pyramid(x['driving'])245 pyramide_generated = self.pyramid(generated['prediction'].detach())246 247 kp_driving = generated['kp_driving']248 discriminator_maps_generated = self.discriminator(pyramide_generated, kp=detach_kp(kp_driving))249 discriminator_maps_real = self.discriminator(pyramide_real, kp=detach_kp(kp_driving))250 251 loss_values = {}252 value_total = 0253 for scale in self.scales:254 key = 'prediction_map_%s' % scale255 value = (1 - discriminator_maps_real[key]) ** 2 + discriminator_maps_generated[key] ** 2256 value_total += self.loss_weights['discriminator_gan'] * value.mean()257 loss_values['disc_gan'] = value_total258 259 return loss_values260 