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, Discriminator, Encoder,12 VectorQuantizer, VectorQuantizerTexture)13from models.losses.segmentation_loss import BCELossWithQuant14from models.losses.vqgan_loss import (DiffAugment, adopt_weight,15 calculate_adaptive_weight, hinge_d_loss)16 17 18class VQModel():19 20 def __init__(self, opt):21 super().__init__()22 self.opt = opt23 self.device = torch.device('cuda')24 self.encoder = Encoder(25 ch=opt['ch'],26 num_res_blocks=opt['num_res_blocks'],27 attn_resolutions=opt['attn_resolutions'],28 ch_mult=opt['ch_mult'],29 in_channels=opt['in_channels'],30 resolution=opt['resolution'],31 z_channels=opt['z_channels'],32 double_z=opt['double_z'],33 dropout=opt['dropout']).to(self.device)34 self.decoder = Decoder(35 in_channels=opt['in_channels'],36 resolution=opt['resolution'],37 z_channels=opt['z_channels'],38 ch=opt['ch'],39 out_ch=opt['out_ch'],40 num_res_blocks=opt['num_res_blocks'],41 attn_resolutions=opt['attn_resolutions'],42 ch_mult=opt['ch_mult'],43 dropout=opt['dropout'],44 resamp_with_conv=True,45 give_pre_end=False).to(self.device)46 self.quantize = VectorQuantizer(47 opt['n_embed'], opt['embed_dim'], beta=0.25).to(self.device)48 self.quant_conv = torch.nn.Conv2d(opt["z_channels"], opt['embed_dim'],49 1).to(self.device)50 self.post_quant_conv = torch.nn.Conv2d(opt['embed_dim'],51 opt["z_channels"],52 1).to(self.device)53 54 def init_training_settings(self):55 self.loss = BCELossWithQuant()56 self.log_dict = OrderedDict()57 self.configure_optimizers()58 59 def save_network(self, save_path):60 """Save networks.61 62 Args:63 net (nn.Module): Network to be saved.64 net_label (str): Network label.65 current_iter (int): Current iter number.66 """67 68 save_dict = {}69 save_dict['encoder'] = self.encoder.state_dict()70 save_dict['decoder'] = self.decoder.state_dict()71 save_dict['quantize'] = self.quantize.state_dict()72 save_dict['quant_conv'] = self.quant_conv.state_dict()73 save_dict['post_quant_conv'] = self.post_quant_conv.state_dict()74 save_dict['discriminator'] = self.disc.state_dict()75 torch.save(save_dict, save_path)76 77 def load_network(self):78 checkpoint = torch.load(self.opt['pretrained_models'])79 self.encoder.load_state_dict(checkpoint['encoder'], strict=True)80 self.decoder.load_state_dict(checkpoint['decoder'], strict=True)81 self.quantize.load_state_dict(checkpoint['quantize'], strict=True)82 self.quant_conv.load_state_dict(checkpoint['quant_conv'], strict=True)83 self.post_quant_conv.load_state_dict(84 checkpoint['post_quant_conv'], strict=True)85 86 def optimize_parameters(self, data, current_iter):87 self.encoder.train()88 self.decoder.train()89 self.quantize.train()90 self.quant_conv.train()91 self.post_quant_conv.train()92 93 loss = self.training_step(data)94 self.optimizer.zero_grad()95 loss.backward()96 self.optimizer.step()97 98 def encode(self, x):99 h = self.encoder(x)100 h = self.quant_conv(h)101 quant, emb_loss, info = self.quantize(h)102 return quant, emb_loss, info103 104 def decode(self, quant):105 quant = self.post_quant_conv(quant)106 dec = self.decoder(quant)107 return dec108 109 def decode_code(self, code_b):110 quant_b = self.quantize.embed_code(code_b)111 dec = self.decode(quant_b)112 return dec113 114 def forward_step(self, input):115 quant, diff, _ = self.encode(input)116 dec = self.decode(quant)117 return dec, diff118 119 def feed_data(self, data):120 x = data['segm']121 x = F.one_hot(x, num_classes=self.opt['num_segm_classes'])122 123 if len(x.shape) == 3:124 x = x[..., None]125 x = x.permute(0, 3, 1, 2).to(memory_format=torch.contiguous_format)126 return x.float().to(self.device)127 128 def get_current_log(self):129 return self.log_dict130 131 def update_learning_rate(self, epoch):132 """Update learning rate.133 134 Args:135 current_iter (int): Current iteration.136 warmup_iter (int): Warmup iter numbers. -1 for no warmup.137 Default: -1.138 """139 lr = self.optimizer.param_groups[0]['lr']140 141 if self.opt['lr_decay'] == 'step':142 lr = self.opt['lr'] * (143 self.opt['gamma']**(epoch // self.opt['step']))144 elif self.opt['lr_decay'] == 'cos':145 lr = self.opt['lr'] * (146 1 + math.cos(math.pi * epoch / self.opt['num_epochs'])) / 2147 elif self.opt['lr_decay'] == 'linear':148 lr = self.opt['lr'] * (1 - epoch / self.opt['num_epochs'])149 elif self.opt['lr_decay'] == 'linear2exp':150 if epoch < self.opt['turning_point'] + 1:151 # learning rate decay as 95%152 # at the turning point (1 / 95% = 1.0526)153 lr = self.opt['lr'] * (154 1 - epoch / int(self.opt['turning_point'] * 1.0526))155 else:156 lr *= self.opt['gamma']157 elif self.opt['lr_decay'] == 'schedule':158 if epoch in self.opt['schedule']:159 lr *= self.opt['gamma']160 else:161 raise ValueError('Unknown lr mode {}'.format(self.opt['lr_decay']))162 # set learning rate163 for param_group in self.optimizer.param_groups:164 param_group['lr'] = lr165 166 return lr167 168 169class VQSegmentationModel(VQModel):170 171 def __init__(self, opt):172 super().__init__(opt)173 self.colorize = torch.randn(3, opt['num_segm_classes'], 1,174 1).to(self.device)175 176 self.init_training_settings()177 178 def configure_optimizers(self):179 self.optimizer = torch.optim.Adam(180 list(self.encoder.parameters()) + list(self.decoder.parameters()) +181 list(self.quantize.parameters()) +182 list(self.quant_conv.parameters()) +183 list(self.post_quant_conv.parameters()),184 lr=self.opt['lr'],185 betas=(0.5, 0.9))186 187 def training_step(self, data):188 x = self.feed_data(data)189 xrec, qloss = self.forward_step(x)190 aeloss, log_dict_ae = self.loss(qloss, x, xrec, split="train")191 self.log_dict.update(log_dict_ae)192 return aeloss193 194 def to_rgb(self, x):195 x = F.conv2d(x, weight=self.colorize)196 x = 2. * (x - x.min()) / (x.max() - x.min()) - 1.197 return x198 199 @torch.no_grad()200 def inference(self, data_loader, save_dir):201 self.encoder.eval()202 self.decoder.eval()203 self.quantize.eval()204 self.quant_conv.eval()205 self.post_quant_conv.eval()206 207 loss_total = 0208 loss_bce = 0209 loss_quant = 0210 num = 0211 212 for _, data in enumerate(data_loader):213 img_name = data['img_name'][0]214 x = self.feed_data(data)215 xrec, qloss = self.forward_step(x)216 _, log_dict_ae = self.loss(qloss, x, xrec, split="val")217 218 loss_total += log_dict_ae['val/total_loss']219 loss_bce += log_dict_ae['val/bce_loss']220 loss_quant += log_dict_ae['val/quant_loss']221 222 num += x.size(0)223 224 if x.shape[1] > 3:225 # colorize with random projection226 assert xrec.shape[1] > 3227 # convert logits to indices228 xrec = torch.argmax(xrec, dim=1, keepdim=True)229 xrec = F.one_hot(xrec, num_classes=x.shape[1])230 xrec = xrec.squeeze(1).permute(0, 3, 1, 2).float()231 x = self.to_rgb(x)232 xrec = self.to_rgb(xrec)233 234 img_cat = torch.cat([x, xrec], dim=3).detach()235 img_cat = ((img_cat + 1) / 2)236 img_cat = img_cat.clamp_(0, 1)237 save_image(238 img_cat, f'{save_dir}/{img_name}.png', nrow=1, padding=4)239 240 return (loss_total / num).item(), (loss_bce /241 num).item(), (loss_quant /242 num).item()243 244 245class VQImageModel(VQModel):246 247 def __init__(self, opt):248 super().__init__(opt)249 self.disc = Discriminator(250 opt['n_channels'], opt['ndf'],251 n_layers=opt['disc_layers']).to(self.device)252 self.perceptual = lpips.LPIPS(net="vgg").to(self.device)253 self.perceptual_weight = opt['perceptual_weight']254 self.disc_start_step = opt['disc_start_step']255 self.disc_weight_max = opt['disc_weight_max']256 self.diff_aug = opt['diff_aug']257 self.policy = "color,translation"258 259 self.disc.train()260 261 self.init_training_settings()262 263 def feed_data(self, data):264 x = data['image']265 266 return x.float().to(self.device)267 268 def init_training_settings(self):269 self.log_dict = OrderedDict()270 self.configure_optimizers()271 272 def configure_optimizers(self):273 self.optimizer = torch.optim.Adam(274 list(self.encoder.parameters()) + list(self.decoder.parameters()) +275 list(self.quantize.parameters()) +276 list(self.quant_conv.parameters()) +277 list(self.post_quant_conv.parameters()),278 lr=self.opt['lr'])279 280 self.disc_optimizer = torch.optim.Adam(281 self.disc.parameters(), lr=self.opt['lr'])282 283 def training_step(self, data, step):284 x = self.feed_data(data)285 xrec, codebook_loss = self.forward_step(x)286 287 # get recon/perceptual loss288 recon_loss = torch.abs(x.contiguous() - xrec.contiguous())289 p_loss = self.perceptual(x.contiguous(), xrec.contiguous())290 nll_loss = recon_loss + self.perceptual_weight * p_loss291 nll_loss = torch.mean(nll_loss)292 293 # augment for input to discriminator294 if self.diff_aug:295 xrec = DiffAugment(xrec, policy=self.policy)296 297 # update generator298 logits_fake = self.disc(xrec)299 g_loss = -torch.mean(logits_fake)300 last_layer = self.decoder.conv_out.weight301 d_weight = calculate_adaptive_weight(nll_loss, g_loss, last_layer,302 self.disc_weight_max)303 d_weight *= adopt_weight(1, step, self.disc_start_step)304 loss = nll_loss + d_weight * g_loss + codebook_loss305 306 self.log_dict["loss"] = loss307 self.log_dict["l1"] = recon_loss.mean().item()308 self.log_dict["perceptual"] = p_loss.mean().item()309 self.log_dict["nll_loss"] = nll_loss.item()310 self.log_dict["g_loss"] = g_loss.item()311 self.log_dict["d_weight"] = d_weight312 self.log_dict["codebook_loss"] = codebook_loss.item()313 314 if step > self.disc_start_step:315 if self.diff_aug:316 logits_real = self.disc(317 DiffAugment(x.contiguous().detach(), policy=self.policy))318 else:319 logits_real = self.disc(x.contiguous().detach())320 logits_fake = self.disc(xrec.contiguous().detach(321 )) # detach so that generator isn"t also updated322 d_loss = hinge_d_loss(logits_real, logits_fake)323 self.log_dict["d_loss"] = d_loss324 else:325 d_loss = None326 327 return loss, d_loss328 329 def optimize_parameters(self, data, step):330 self.encoder.train()331 self.decoder.train()332 self.quantize.train()333 self.quant_conv.train()334 self.post_quant_conv.train()335 336 loss, d_loss = self.training_step(data, step)337 self.optimizer.zero_grad()338 loss.backward()339 self.optimizer.step()340 341 if step > self.disc_start_step:342 self.disc_optimizer.zero_grad()343 d_loss.backward()344 self.disc_optimizer.step()345 346 @torch.no_grad()347 def inference(self, data_loader, save_dir):348 self.encoder.eval()349 self.decoder.eval()350 self.quantize.eval()351 self.quant_conv.eval()352 self.post_quant_conv.eval()353 354 loss_total = 0355 num = 0356 357 for _, data in enumerate(data_loader):358 img_name = data['img_name'][0]359 x = self.feed_data(data)360 xrec, _ = self.forward_step(x)361 362 recon_loss = torch.abs(x.contiguous() - xrec.contiguous())363 p_loss = self.perceptual(x.contiguous(), xrec.contiguous())364 nll_loss = recon_loss + self.perceptual_weight * p_loss365 nll_loss = torch.mean(nll_loss)366 loss_total += nll_loss367 368 num += x.size(0)369 370 if x.shape[1] > 3:371 # colorize with random projection372 assert xrec.shape[1] > 3373 # convert logits to indices374 xrec = torch.argmax(xrec, dim=1, keepdim=True)375 xrec = F.one_hot(xrec, num_classes=x.shape[1])376 xrec = xrec.squeeze(1).permute(0, 3, 1, 2).float()377 x = self.to_rgb(x)378 xrec = self.to_rgb(xrec)379 380 img_cat = torch.cat([x, xrec], dim=3).detach()381 img_cat = ((img_cat + 1) / 2)382 img_cat = img_cat.clamp_(0, 1)383 save_image(384 img_cat, f'{save_dir}/{img_name}.png', nrow=1, padding=4)385 386 return (loss_total / num).item()387 388 389class VQImageSegmTextureModel(VQImageModel):390 391 def __init__(self, opt):392 self.opt = opt393 self.device = torch.device('cuda')394 self.encoder = Encoder(395 ch=opt['ch'],396 num_res_blocks=opt['num_res_blocks'],397 attn_resolutions=opt['attn_resolutions'],398 ch_mult=opt['ch_mult'],399 in_channels=opt['in_channels'],400 resolution=opt['resolution'],401 z_channels=opt['z_channels'],402 double_z=opt['double_z'],403 dropout=opt['dropout']).to(self.device)404 self.decoder = Decoder(405 in_channels=opt['in_channels'],406 resolution=opt['resolution'],407 z_channels=opt['z_channels'],408 ch=opt['ch'],409 out_ch=opt['out_ch'],410 num_res_blocks=opt['num_res_blocks'],411 attn_resolutions=opt['attn_resolutions'],412 ch_mult=opt['ch_mult'],413 dropout=opt['dropout'],414 resamp_with_conv=True,415 give_pre_end=False).to(self.device)416 self.quantize = VectorQuantizerTexture(417 opt['n_embed'], opt['embed_dim'], beta=0.25).to(self.device)418 self.quant_conv = torch.nn.Conv2d(opt["z_channels"], opt['embed_dim'],419 1).to(self.device)420 self.post_quant_conv = torch.nn.Conv2d(opt['embed_dim'],421 opt["z_channels"],422 1).to(self.device)423 424 self.disc = Discriminator(425 opt['n_channels'], opt['ndf'],426 n_layers=opt['disc_layers']).to(self.device)427 self.perceptual = lpips.LPIPS(net="vgg").to(self.device)428 self.perceptual_weight = opt['perceptual_weight']429 self.disc_start_step = opt['disc_start_step']430 self.disc_weight_max = opt['disc_weight_max']431 self.diff_aug = opt['diff_aug']432 self.policy = "color,translation"433 434 self.disc.train()435 436 self.init_training_settings()437 438 def feed_data(self, data):439 x = data['image'].float().to(self.device)440 mask = data['texture_mask'].float().to(self.device)441 442 return x, mask443 444 def training_step(self, data, step):445 x, mask = self.feed_data(data)446 xrec, codebook_loss = self.forward_step(x, mask)447 448 # get recon/perceptual loss449 recon_loss = torch.abs(x.contiguous() - xrec.contiguous())450 p_loss = self.perceptual(x.contiguous(), xrec.contiguous())451 nll_loss = recon_loss + self.perceptual_weight * p_loss452 nll_loss = torch.mean(nll_loss)453 454 # augment for input to discriminator455 if self.diff_aug:456 xrec = DiffAugment(xrec, policy=self.policy)457 458 # update generator459 logits_fake = self.disc(xrec)460 g_loss = -torch.mean(logits_fake)461 last_layer = self.decoder.conv_out.weight462 d_weight = calculate_adaptive_weight(nll_loss, g_loss, last_layer,463 self.disc_weight_max)464 d_weight *= adopt_weight(1, step, self.disc_start_step)465 loss = nll_loss + d_weight * g_loss + codebook_loss466 467 self.log_dict["loss"] = loss468 self.log_dict["l1"] = recon_loss.mean().item()469 self.log_dict["perceptual"] = p_loss.mean().item()470 self.log_dict["nll_loss"] = nll_loss.item()471 self.log_dict["g_loss"] = g_loss.item()472 self.log_dict["d_weight"] = d_weight473 self.log_dict["codebook_loss"] = codebook_loss.item()474 475 if step > self.disc_start_step:476 if self.diff_aug:477 logits_real = self.disc(478 DiffAugment(x.contiguous().detach(), policy=self.policy))479 else:480 logits_real = self.disc(x.contiguous().detach())481 logits_fake = self.disc(xrec.contiguous().detach(482 )) # detach so that generator isn"t also updated483 d_loss = hinge_d_loss(logits_real, logits_fake)484 self.log_dict["d_loss"] = d_loss485 else:486 d_loss = None487 488 return loss, d_loss489 490 @torch.no_grad()491 def inference(self, data_loader, save_dir):492 self.encoder.eval()493 self.decoder.eval()494 self.quantize.eval()495 self.quant_conv.eval()496 self.post_quant_conv.eval()497 498 loss_total = 0499 num = 0500 501 for _, data in enumerate(data_loader):502 img_name = data['img_name'][0]503 x, mask = self.feed_data(data)504 xrec, _ = self.forward_step(x, mask)505 506 recon_loss = torch.abs(x.contiguous() - xrec.contiguous())507 p_loss = self.perceptual(x.contiguous(), xrec.contiguous())508 nll_loss = recon_loss + self.perceptual_weight * p_loss509 nll_loss = torch.mean(nll_loss)510 loss_total += nll_loss511 512 num += x.size(0)513 514 if x.shape[1] > 3:515 # colorize with random projection516 assert xrec.shape[1] > 3517 # convert logits to indices518 xrec = torch.argmax(xrec, dim=1, keepdim=True)519 xrec = F.one_hot(xrec, num_classes=x.shape[1])520 xrec = xrec.squeeze(1).permute(0, 3, 1, 2).float()521 x = self.to_rgb(x)522 xrec = self.to_rgb(xrec)523 524 img_cat = torch.cat([x, xrec], dim=3).detach()525 img_cat = ((img_cat + 1) / 2)526 img_cat = img_cat.clamp_(0, 1)527 save_image(528 img_cat, f'{save_dir}/{img_name}.png', nrow=1, padding=4)529 530 return (loss_total / num).item()531 532 def encode(self, x, mask):533 h = self.encoder(x)534 h = self.quant_conv(h)535 quant, emb_loss, info = self.quantize(h, mask)536 return quant, emb_loss, info537 538 def decode(self, quant):539 quant = self.post_quant_conv(quant)540 dec = self.decoder(quant)541 return dec542 543 def decode_code(self, code_b):544 quant_b = self.quantize.embed_code(code_b)545 dec = self.decode(quant_b)546 return dec547 548 def forward_step(self, input, mask):549 quant, diff, _ = self.encode(input, mask)550 dec = self.decode(quant)551 return dec, diff552 