radames/Text2Human-API
1
1import math2import sys3from collections import OrderedDict4 5sys.path.append('..')6import lpips7import torch8import torch.nn.functional as F9from torchvision.utils import save_image10 11from models.archs.vqgan_arch import (Decoder, DecoderRes, Discriminator,12 Encoder,13 VectorQuantizerSpatialTextureAware,14 VectorQuantizerTexture)15from models.losses.vqgan_loss import (DiffAugment, adopt_weight,16 calculate_adaptive_weight, hinge_d_loss)17 18 19class HierarchyVQSpatialTextureAwareModel():20 21 def __init__(self, opt):22 self.opt = opt23 self.device = torch.device('cuda')24 self.top_encoder = Encoder(25 ch=opt['top_ch'],26 num_res_blocks=opt['top_num_res_blocks'],27 attn_resolutions=opt['top_attn_resolutions'],28 ch_mult=opt['top_ch_mult'],29 in_channels=opt['top_in_channels'],30 resolution=opt['top_resolution'],31 z_channels=opt['top_z_channels'],32 double_z=opt['top_double_z'],33 dropout=opt['top_dropout']).to(self.device)34 self.decoder = Decoder(35 in_channels=opt['top_in_channels'],36 resolution=opt['top_resolution'],37 z_channels=opt['top_z_channels'],38 ch=opt['top_ch'],39 out_ch=opt['top_out_ch'],40 num_res_blocks=opt['top_num_res_blocks'],41 attn_resolutions=opt['top_attn_resolutions'],42 ch_mult=opt['top_ch_mult'],43 dropout=opt['top_dropout'],44 resamp_with_conv=True,45 give_pre_end=False).to(self.device)46 self.top_quantize = VectorQuantizerTexture(47 1024, opt['embed_dim'], beta=0.25).to(self.device)48 self.top_quant_conv = torch.nn.Conv2d(opt["top_z_channels"],49 opt['embed_dim'],50 1).to(self.device)51 self.top_post_quant_conv = torch.nn.Conv2d(opt['embed_dim'],52 opt["top_z_channels"],53 1).to(self.device)54 self.load_top_pretrain_models()55 56 self.bot_encoder = Encoder(57 ch=opt['bot_ch'],58 num_res_blocks=opt['bot_num_res_blocks'],59 attn_resolutions=opt['bot_attn_resolutions'],60 ch_mult=opt['bot_ch_mult'],61 in_channels=opt['bot_in_channels'],62 resolution=opt['bot_resolution'],63 z_channels=opt['bot_z_channels'],64 double_z=opt['bot_double_z'],65 dropout=opt['bot_dropout']).to(self.device)66 self.bot_decoder_res = DecoderRes(67 in_channels=opt['bot_in_channels'],68 resolution=opt['bot_resolution'],69 z_channels=opt['bot_z_channels'],70 ch=opt['bot_ch'],71 num_res_blocks=opt['bot_num_res_blocks'],72 ch_mult=opt['bot_ch_mult'],73 dropout=opt['bot_dropout'],74 give_pre_end=False).to(self.device)75 self.bot_quantize = VectorQuantizerSpatialTextureAware(76 opt['bot_n_embed'],77 opt['embed_dim'],78 beta=0.25,79 spatial_size=opt['codebook_spatial_size']).to(self.device)80 self.bot_quant_conv = torch.nn.Conv2d(opt["bot_z_channels"],81 opt['embed_dim'],82 1).to(self.device)83 self.bot_post_quant_conv = torch.nn.Conv2d(opt['embed_dim'],84 opt["bot_z_channels"],85 1).to(self.device)86 87 self.disc = Discriminator(88 opt['n_channels'], opt['ndf'],89 n_layers=opt['disc_layers']).to(self.device)90 self.perceptual = lpips.LPIPS(net="vgg").to(self.device)91 self.perceptual_weight = opt['perceptual_weight']92 self.disc_start_step = opt['disc_start_step']93 self.disc_weight_max = opt['disc_weight_max']94 self.diff_aug = opt['diff_aug']95 self.policy = "color,translation"96 97 self.load_discriminator_models()98 99 self.disc.train()100 101 self.fix_decoder = opt['fix_decoder']102 103 self.init_training_settings()104 105 def load_top_pretrain_models(self):106 # load pretrained vqgan for segmentation mask107 top_vae_checkpoint = torch.load(self.opt['top_vae_path'])108 self.top_encoder.load_state_dict(109 top_vae_checkpoint['encoder'], strict=True)110 self.decoder.load_state_dict(111 top_vae_checkpoint['decoder'], strict=True)112 self.top_quantize.load_state_dict(113 top_vae_checkpoint['quantize'], strict=True)114 self.top_quant_conv.load_state_dict(115 top_vae_checkpoint['quant_conv'], strict=True)116 self.top_post_quant_conv.load_state_dict(117 top_vae_checkpoint['post_quant_conv'], strict=True)118 self.top_encoder.eval()119 self.top_quantize.eval()120 self.top_quant_conv.eval()121 self.top_post_quant_conv.eval()122 123 def init_training_settings(self):124 self.log_dict = OrderedDict()125 self.configure_optimizers()126 127 def configure_optimizers(self):128 optim_params = []129 for v in self.bot_encoder.parameters():130 if v.requires_grad:131 optim_params.append(v)132 for v in self.bot_decoder_res.parameters():133 if v.requires_grad:134 optim_params.append(v)135 for v in self.bot_quantize.parameters():136 if v.requires_grad:137 optim_params.append(v)138 for v in self.bot_quant_conv.parameters():139 if v.requires_grad:140 optim_params.append(v)141 for v in self.bot_post_quant_conv.parameters():142 if v.requires_grad:143 optim_params.append(v)144 if not self.fix_decoder:145 for name, v in self.decoder.named_parameters():146 if v.requires_grad:147 if 'up.0' in name:148 optim_params.append(v)149 if 'up.1' in name:150 optim_params.append(v)151 if 'up.2' in name:152 optim_params.append(v)153 if 'up.3' in name:154 optim_params.append(v)155 156 self.optimizer = torch.optim.Adam(optim_params, lr=self.opt['lr'])157 158 self.disc_optimizer = torch.optim.Adam(159 self.disc.parameters(), lr=self.opt['lr'])160 161 def load_discriminator_models(self):162 # load pretrained vqgan for segmentation mask163 top_vae_checkpoint = torch.load(self.opt['top_vae_path'])164 self.disc.load_state_dict(165 top_vae_checkpoint['discriminator'], strict=True)166 167 def save_network(self, save_path):168 """Save networks.169 """170 171 save_dict = {}172 save_dict['bot_encoder'] = self.bot_encoder.state_dict()173 save_dict['bot_decoder_res'] = self.bot_decoder_res.state_dict()174 save_dict['decoder'] = self.decoder.state_dict()175 save_dict['bot_quantize'] = self.bot_quantize.state_dict()176 save_dict['bot_quant_conv'] = self.bot_quant_conv.state_dict()177 save_dict['bot_post_quant_conv'] = self.bot_post_quant_conv.state_dict(178 )179 save_dict['discriminator'] = self.disc.state_dict()180 torch.save(save_dict, save_path)181 182 def load_network(self):183 checkpoint = torch.load(self.opt['pretrained_models'])184 self.bot_encoder.load_state_dict(185 checkpoint['bot_encoder'], strict=True)186 self.bot_decoder_res.load_state_dict(187 checkpoint['bot_decoder_res'], strict=True)188 self.decoder.load_state_dict(checkpoint['decoder'], strict=True)189 self.bot_quantize.load_state_dict(190 checkpoint['bot_quantize'], strict=True)191 self.bot_quant_conv.load_state_dict(192 checkpoint['bot_quant_conv'], strict=True)193 self.bot_post_quant_conv.load_state_dict(194 checkpoint['bot_post_quant_conv'], strict=True)195 196 def optimize_parameters(self, data, step):197 self.bot_encoder.train()198 self.bot_decoder_res.train()199 if not self.fix_decoder:200 self.decoder.train()201 self.bot_quantize.train()202 self.bot_quant_conv.train()203 self.bot_post_quant_conv.train()204 205 loss, d_loss = self.training_step(data, step)206 self.optimizer.zero_grad()207 loss.backward()208 self.optimizer.step()209 210 if step > self.disc_start_step:211 self.disc_optimizer.zero_grad()212 d_loss.backward()213 self.disc_optimizer.step()214 215 def top_encode(self, x, mask):216 h = self.top_encoder(x)217 h = self.top_quant_conv(h)218 quant, _, _ = self.top_quantize(h, mask)219 quant = self.top_post_quant_conv(quant)220 return quant221 222 def bot_encode(self, x, mask):223 h = self.bot_encoder(x)224 h = self.bot_quant_conv(h)225 quant, emb_loss, info = self.bot_quantize(h, mask)226 quant = self.bot_post_quant_conv(quant)227 bot_dec_res = self.bot_decoder_res(quant)228 return bot_dec_res, emb_loss, info229 230 def decode(self, quant_top, bot_dec_res):231 dec = self.decoder(quant_top, bot_h=bot_dec_res)232 return dec233 234 def forward_step(self, input, mask):235 with torch.no_grad():236 quant_top = self.top_encode(input, mask)237 bot_dec_res, diff, _ = self.bot_encode(input, mask)238 dec = self.decode(quant_top, bot_dec_res)239 return dec, diff240 241 def feed_data(self, data):242 x = data['image'].float().to(self.device)243 mask = data['texture_mask'].float().to(self.device)244 245 return x, mask246 247 def training_step(self, data, step):248 x, mask = self.feed_data(data)249 xrec, codebook_loss = self.forward_step(x, mask)250 251 # get recon/perceptual loss252 recon_loss = torch.abs(x.contiguous() - xrec.contiguous())253 p_loss = self.perceptual(x.contiguous(), xrec.contiguous())254 nll_loss = recon_loss + self.perceptual_weight * p_loss255 nll_loss = torch.mean(nll_loss)256 257 # augment for input to discriminator258 if self.diff_aug:259 xrec = DiffAugment(xrec, policy=self.policy)260 261 # update generator262 logits_fake = self.disc(xrec)263 g_loss = -torch.mean(logits_fake)264 last_layer = self.decoder.conv_out.weight265 d_weight = calculate_adaptive_weight(nll_loss, g_loss, last_layer,266 self.disc_weight_max)267 d_weight *= adopt_weight(1, step, self.disc_start_step)268 loss = nll_loss + d_weight * g_loss + codebook_loss269 270 self.log_dict["loss"] = loss271 self.log_dict["l1"] = recon_loss.mean().item()272 self.log_dict["perceptual"] = p_loss.mean().item()273 self.log_dict["nll_loss"] = nll_loss.item()274 self.log_dict["g_loss"] = g_loss.item()275 self.log_dict["d_weight"] = d_weight276 self.log_dict["codebook_loss"] = codebook_loss.item()277 278 if step > self.disc_start_step:279 if self.diff_aug:280 logits_real = self.disc(281 DiffAugment(x.contiguous().detach(), policy=self.policy))282 else:283 logits_real = self.disc(x.contiguous().detach())284 logits_fake = self.disc(xrec.contiguous().detach(285 )) # detach so that generator isn"t also updated286 d_loss = hinge_d_loss(logits_real, logits_fake)287 self.log_dict["d_loss"] = d_loss288 else:289 d_loss = None290 291 return loss, d_loss292 293 @torch.no_grad()294 def inference(self, data_loader, save_dir):295 self.bot_encoder.eval()296 self.bot_decoder_res.eval()297 self.decoder.eval()298 self.bot_quantize.eval()299 self.bot_quant_conv.eval()300 self.bot_post_quant_conv.eval()301 302 loss_total = 0303 num = 0304 305 for _, data in enumerate(data_loader):306 img_name = data['img_name'][0]307 x, mask = self.feed_data(data)308 xrec, _ = self.forward_step(x, mask)309 310 recon_loss = torch.abs(x.contiguous() - xrec.contiguous())311 p_loss = self.perceptual(x.contiguous(), xrec.contiguous())312 nll_loss = recon_loss + self.perceptual_weight * p_loss313 nll_loss = torch.mean(nll_loss)314 loss_total += nll_loss315 316 num += x.size(0)317 318 if x.shape[1] > 3:319 # colorize with random projection320 assert xrec.shape[1] > 3321 # convert logits to indices322 xrec = torch.argmax(xrec, dim=1, keepdim=True)323 xrec = F.one_hot(xrec, num_classes=x.shape[1])324 xrec = xrec.squeeze(1).permute(0, 3, 1, 2).float()325 x = self.to_rgb(x)326 xrec = self.to_rgb(xrec)327 328 img_cat = torch.cat([x, xrec], dim=3).detach()329 img_cat = ((img_cat + 1) / 2)330 img_cat = img_cat.clamp_(0, 1)331 save_image(332 img_cat, f'{save_dir}/{img_name}.png', nrow=1, padding=4)333 334 return (loss_total / num).item()335 336 def get_current_log(self):337 return self.log_dict338 339 def update_learning_rate(self, epoch):340 """Update learning rate.341 342 Args:343 current_iter (int): Current iteration.344 warmup_iter (int): Warmup iter numbers. -1 for no warmup.345 Default: -1.346 """347 lr = self.optimizer.param_groups[0]['lr']348 349 if self.opt['lr_decay'] == 'step':350 lr = self.opt['lr'] * (351 self.opt['gamma']**(epoch // self.opt['step']))352 elif self.opt['lr_decay'] == 'cos':353 lr = self.opt['lr'] * (354 1 + math.cos(math.pi * epoch / self.opt['num_epochs'])) / 2355 elif self.opt['lr_decay'] == 'linear':356 lr = self.opt['lr'] * (1 - epoch / self.opt['num_epochs'])357 elif self.opt['lr_decay'] == 'linear2exp':358 if epoch < self.opt['turning_point'] + 1:359 # learning rate decay as 95%360 # at the turning point (1 / 95% = 1.0526)361 lr = self.opt['lr'] * (362 1 - epoch / int(self.opt['turning_point'] * 1.0526))363 else:364 lr *= self.opt['gamma']365 elif self.opt['lr_decay'] == 'schedule':366 if epoch in self.opt['schedule']:367 lr *= self.opt['gamma']368 else:369 raise ValueError('Unknown lr mode {}'.format(self.opt['lr_decay']))370 # set learning rate371 for param_group in self.optimizer.param_groups:372 param_group['lr'] = lr373 374 return lr375 