gulabpatel/First-Order-Motion
0
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 