CoolFace
Apppublic

gulabpatel/First-Order-Motion

sourceHugging Faceupdated 5y agoView on Hugging Face
0likes
logger.py209 linesDownload Raw Back to root
1import numpy as np2import torch3import torch.nn.functional as F4import imageio5 6import os7from skimage.draw import circle8 9import matplotlib.pyplot as plt10import collections11 12 13class Logger:14    def __init__(self, log_dir, checkpoint_freq=100, visualizer_params=None, zfill_num=8, log_file_name='log.txt'):15 16        self.loss_list = []17        self.cpk_dir = log_dir18        self.visualizations_dir = os.path.join(log_dir, 'train-vis')19        if not os.path.exists(self.visualizations_dir):20            os.makedirs(self.visualizations_dir)21        self.log_file = open(os.path.join(log_dir, log_file_name), 'a')22        self.zfill_num = zfill_num23        self.visualizer = Visualizer(**visualizer_params)24        self.checkpoint_freq = checkpoint_freq25        self.epoch = 026        self.best_loss = float('inf')27        self.names = None28 29    def log_scores(self, loss_names):30        loss_mean = np.array(self.loss_list).mean(axis=0)31 32        loss_string = "; ".join(["%s - %.5f" % (name, value) for name, value in zip(loss_names, loss_mean)])33        loss_string = str(self.epoch).zfill(self.zfill_num) + ") " + loss_string34 35        print(loss_string, file=self.log_file)36        self.loss_list = []37        self.log_file.flush()38 39    def visualize_rec(self, inp, out):40        image = self.visualizer.visualize(inp['driving'], inp['source'], out)41        imageio.imsave(os.path.join(self.visualizations_dir, "%s-rec.png" % str(self.epoch).zfill(self.zfill_num)), image)42 43    def save_cpk(self, emergent=False):44        cpk = {k: v.state_dict() for k, v in self.models.items()}45        cpk['epoch'] = self.epoch46        cpk_path = os.path.join(self.cpk_dir, '%s-checkpoint.pth.tar' % str(self.epoch).zfill(self.zfill_num)) 47        if not (os.path.exists(cpk_path) and emergent):48            torch.save(cpk, cpk_path)49 50    @staticmethod51    def load_cpk(checkpoint_path, generator=None, discriminator=None, kp_detector=None,52                 optimizer_generator=None, optimizer_discriminator=None, optimizer_kp_detector=None):53        checkpoint = torch.load(checkpoint_path)54        if generator is not None:55            generator.load_state_dict(checkpoint['generator'])56        if kp_detector is not None:57            kp_detector.load_state_dict(checkpoint['kp_detector'])58        if discriminator is not None:59            try:60               discriminator.load_state_dict(checkpoint['discriminator'])61            except:62               print ('No discriminator in the state-dict. Dicriminator will be randomly initialized')63        if optimizer_generator is not None:64            optimizer_generator.load_state_dict(checkpoint['optimizer_generator'])65        if optimizer_discriminator is not None:66            try:67                optimizer_discriminator.load_state_dict(checkpoint['optimizer_discriminator'])68            except RuntimeError as e:69                print ('No discriminator optimizer in the state-dict. Optimizer will be not initialized')70        if optimizer_kp_detector is not None:71            optimizer_kp_detector.load_state_dict(checkpoint['optimizer_kp_detector'])72 73        return checkpoint['epoch']74 75    def __enter__(self):76        return self77 78    def __exit__(self, exc_type, exc_val, exc_tb):79        if 'models' in self.__dict__:80            self.save_cpk()81        self.log_file.close()82 83    def log_iter(self, losses):84        losses = collections.OrderedDict(losses.items())85        if self.names is None:86            self.names = list(losses.keys())87        self.loss_list.append(list(losses.values()))88 89    def log_epoch(self, epoch, models, inp, out):90        self.epoch = epoch91        self.models = models92        if (self.epoch + 1) % self.checkpoint_freq == 0:93            self.save_cpk()94        self.log_scores(self.names)95        self.visualize_rec(inp, out)96 97 98class Visualizer:99    def __init__(self, kp_size=5, draw_border=False, colormap='gist_rainbow'):100        self.kp_size = kp_size101        self.draw_border = draw_border102        self.colormap = plt.get_cmap(colormap)103 104    def draw_image_with_kp(self, image, kp_array):105        image = np.copy(image)106        spatial_size = np.array(image.shape[:2][::-1])[np.newaxis]107        kp_array = spatial_size * (kp_array + 1) / 2108        num_kp = kp_array.shape[0]109        for kp_ind, kp in enumerate(kp_array):110            rr, cc = circle(kp[1], kp[0], self.kp_size, shape=image.shape[:2])111            image[rr, cc] = np.array(self.colormap(kp_ind / num_kp))[:3]112        return image113 114    def create_image_column_with_kp(self, images, kp):115        image_array = np.array([self.draw_image_with_kp(v, k) for v, k in zip(images, kp)])116        return self.create_image_column(image_array)117 118    def create_image_column(self, images):119        if self.draw_border:120            images = np.copy(images)121            images[:, :, [0, -1]] = (1, 1, 1)122            images[:, :, [0, -1]] = (1, 1, 1)123        return np.concatenate(list(images), axis=0)124 125    def create_image_grid(self, *args):126        out = []127        for arg in args:128            if type(arg) == tuple:129                out.append(self.create_image_column_with_kp(arg[0], arg[1]))130            else:131                out.append(self.create_image_column(arg))132        return np.concatenate(out, axis=1)133 134    def visualize(self, driving, source, out):135        images = []136 137        # Source image with keypoints138        source = source.data.cpu()139        kp_source = out['kp_source']['value'].data.cpu().numpy()140        source = np.transpose(source, [0, 2, 3, 1])141        images.append((source, kp_source))142 143        # Equivariance visualization144        if 'transformed_frame' in out:145            transformed = out['transformed_frame'].data.cpu().numpy()146            transformed = np.transpose(transformed, [0, 2, 3, 1])147            transformed_kp = out['transformed_kp']['value'].data.cpu().numpy()148            images.append((transformed, transformed_kp))149 150        # Driving image with keypoints151        kp_driving = out['kp_driving']['value'].data.cpu().numpy()152        driving = driving.data.cpu().numpy()153        driving = np.transpose(driving, [0, 2, 3, 1])154        images.append((driving, kp_driving))155 156        # Deformed image157        if 'deformed' in out:158            deformed = out['deformed'].data.cpu().numpy()159            deformed = np.transpose(deformed, [0, 2, 3, 1])160            images.append(deformed)161 162        # Result with and without keypoints163        prediction = out['prediction'].data.cpu().numpy()164        prediction = np.transpose(prediction, [0, 2, 3, 1])165        if 'kp_norm' in out:166            kp_norm = out['kp_norm']['value'].data.cpu().numpy()167            images.append((prediction, kp_norm))168        images.append(prediction)169 170 171        ## Occlusion map172        if 'occlusion_map' in out:173            occlusion_map = out['occlusion_map'].data.cpu().repeat(1, 3, 1, 1)174            occlusion_map = F.interpolate(occlusion_map, size=source.shape[1:3]).numpy()175            occlusion_map = np.transpose(occlusion_map, [0, 2, 3, 1])176            images.append(occlusion_map)177 178        # Deformed images according to each individual transform179        if 'sparse_deformed' in out:180            full_mask = []181            for i in range(out['sparse_deformed'].shape[1]):182                image = out['sparse_deformed'][:, i].data.cpu()183                image = F.interpolate(image, size=source.shape[1:3])184                mask = out['mask'][:, i:(i+1)].data.cpu().repeat(1, 3, 1, 1)185                mask = F.interpolate(mask, size=source.shape[1:3])186                image = np.transpose(image.numpy(), (0, 2, 3, 1))187                mask = np.transpose(mask.numpy(), (0, 2, 3, 1))188 189                if i != 0:190                    color = np.array(self.colormap((i - 1) / (out['sparse_deformed'].shape[1] - 1)))[:3]191                else:192                    color = np.array((0, 0, 0))193 194                color = color.reshape((1, 1, 1, 3))195 196                images.append(image)197                if i != 0:198                    images.append(mask * color)199                else:200                    images.append(mask)201 202                full_mask.append(mask * color)203 204            images.append(sum(full_mask))205 206        image = self.create_image_grid(*images)207        image = (255 * image).astype(np.uint8)208        return image209