Team Ai
Apppublic

gulabpatel/First-Order-Motion

sourceHugging Faceupdated 5y agoView on Hugging Face
0likes
model.py260 linesDownload Raw Back to modules
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