radames/Text2Human-API
1
1import logging2import math3from collections import OrderedDict4 5import torch6import torch.nn.functional as F7from torchvision.utils import save_image8 9from models.archs.fcn_arch import MultiHeadFCNHead10from models.archs.unet_arch import UNet11from models.archs.vqgan_arch import (Decoder, DecoderRes, Encoder,12 VectorQuantizerSpatialTextureAware,13 VectorQuantizerTexture)14from models.losses.accuracy import accuracy15from models.losses.cross_entropy_loss import CrossEntropyLoss16 17logger = logging.getLogger('base')18 19 20class VQGANTextureAwareSpatialHierarchyInferenceModel():21 22 def __init__(self, opt):23 self.opt = opt24 self.device = torch.device('cuda')25 self.is_train = opt['is_train']26 27 self.top_encoder = Encoder(28 ch=opt['top_ch'],29 num_res_blocks=opt['top_num_res_blocks'],30 attn_resolutions=opt['top_attn_resolutions'],31 ch_mult=opt['top_ch_mult'],32 in_channels=opt['top_in_channels'],33 resolution=opt['top_resolution'],34 z_channels=opt['top_z_channels'],35 double_z=opt['top_double_z'],36 dropout=opt['top_dropout']).to(self.device)37 self.decoder = Decoder(38 in_channels=opt['top_in_channels'],39 resolution=opt['top_resolution'],40 z_channels=opt['top_z_channels'],41 ch=opt['top_ch'],42 out_ch=opt['top_out_ch'],43 num_res_blocks=opt['top_num_res_blocks'],44 attn_resolutions=opt['top_attn_resolutions'],45 ch_mult=opt['top_ch_mult'],46 dropout=opt['top_dropout'],47 resamp_with_conv=True,48 give_pre_end=False).to(self.device)49 self.top_quantize = VectorQuantizerTexture(50 1024, opt['embed_dim'], beta=0.25).to(self.device)51 self.top_quant_conv = torch.nn.Conv2d(opt["top_z_channels"],52 opt['embed_dim'],53 1).to(self.device)54 self.top_post_quant_conv = torch.nn.Conv2d(opt['embed_dim'],55 opt["top_z_channels"],56 1).to(self.device)57 self.load_top_pretrain_models()58 59 self.bot_encoder = Encoder(60 ch=opt['bot_ch'],61 num_res_blocks=opt['bot_num_res_blocks'],62 attn_resolutions=opt['bot_attn_resolutions'],63 ch_mult=opt['bot_ch_mult'],64 in_channels=opt['bot_in_channels'],65 resolution=opt['bot_resolution'],66 z_channels=opt['bot_z_channels'],67 double_z=opt['bot_double_z'],68 dropout=opt['bot_dropout']).to(self.device)69 self.bot_decoder_res = DecoderRes(70 in_channels=opt['bot_in_channels'],71 resolution=opt['bot_resolution'],72 z_channels=opt['bot_z_channels'],73 ch=opt['bot_ch'],74 num_res_blocks=opt['bot_num_res_blocks'],75 ch_mult=opt['bot_ch_mult'],76 dropout=opt['bot_dropout'],77 give_pre_end=False).to(self.device)78 self.bot_quantize = VectorQuantizerSpatialTextureAware(79 opt['bot_n_embed'],80 opt['embed_dim'],81 beta=0.25,82 spatial_size=opt['codebook_spatial_size']).to(self.device)83 self.bot_quant_conv = torch.nn.Conv2d(opt["bot_z_channels"],84 opt['embed_dim'],85 1).to(self.device)86 self.bot_post_quant_conv = torch.nn.Conv2d(opt['embed_dim'],87 opt["bot_z_channels"],88 1).to(self.device)89 90 self.load_bot_pretrain_network()91 92 self.guidance_encoder = UNet(93 in_channels=opt['encoder_in_channels']).to(self.device)94 self.index_decoder = MultiHeadFCNHead(95 in_channels=opt['fc_in_channels'],96 in_index=opt['fc_in_index'],97 channels=opt['fc_channels'],98 num_convs=opt['fc_num_convs'],99 concat_input=opt['fc_concat_input'],100 dropout_ratio=opt['fc_dropout_ratio'],101 num_classes=opt['fc_num_classes'],102 align_corners=opt['fc_align_corners'],103 num_head=18).to(self.device)104 105 self.init_training_settings()106 107 def init_training_settings(self):108 optim_params = []109 for v in self.guidance_encoder.parameters():110 if v.requires_grad:111 optim_params.append(v)112 for v in self.index_decoder.parameters():113 if v.requires_grad:114 optim_params.append(v)115 # set up optimizers116 if self.opt['optimizer'] == 'Adam':117 self.optimizer = torch.optim.Adam(118 optim_params,119 self.opt['lr'],120 weight_decay=self.opt['weight_decay'])121 elif self.opt['optimizer'] == 'SGD':122 self.optimizer = torch.optim.SGD(123 optim_params,124 self.opt['lr'],125 momentum=self.opt['momentum'],126 weight_decay=self.opt['weight_decay'])127 self.log_dict = OrderedDict()128 if self.opt['loss_function'] == 'cross_entropy':129 self.loss_func = CrossEntropyLoss().to(self.device)130 131 def load_top_pretrain_models(self):132 # load pretrained vqgan for segmentation mask133 top_vae_checkpoint = torch.load(self.opt['top_vae_path'])134 self.top_encoder.load_state_dict(135 top_vae_checkpoint['encoder'], strict=True)136 self.decoder.load_state_dict(137 top_vae_checkpoint['decoder'], strict=True)138 self.top_quantize.load_state_dict(139 top_vae_checkpoint['quantize'], strict=True)140 self.top_quant_conv.load_state_dict(141 top_vae_checkpoint['quant_conv'], strict=True)142 self.top_post_quant_conv.load_state_dict(143 top_vae_checkpoint['post_quant_conv'], strict=True)144 self.top_encoder.eval()145 self.top_quantize.eval()146 self.top_quant_conv.eval()147 self.top_post_quant_conv.eval()148 149 def load_bot_pretrain_network(self):150 checkpoint = torch.load(self.opt['bot_vae_path'])151 self.bot_encoder.load_state_dict(152 checkpoint['bot_encoder'], strict=True)153 self.bot_decoder_res.load_state_dict(154 checkpoint['bot_decoder_res'], strict=True)155 self.decoder.load_state_dict(checkpoint['decoder'], strict=True)156 self.bot_quantize.load_state_dict(157 checkpoint['bot_quantize'], strict=True)158 self.bot_quant_conv.load_state_dict(159 checkpoint['bot_quant_conv'], strict=True)160 self.bot_post_quant_conv.load_state_dict(161 checkpoint['bot_post_quant_conv'], strict=True)162 163 self.bot_encoder.eval()164 self.bot_decoder_res.eval()165 self.decoder.eval()166 self.bot_quantize.eval()167 self.bot_quant_conv.eval()168 self.bot_post_quant_conv.eval()169 170 def top_encode(self, x, mask):171 h = self.top_encoder(x)172 h = self.top_quant_conv(h)173 quant, _, _ = self.top_quantize(h, mask)174 quant = self.top_post_quant_conv(quant)175 176 return quant, quant177 178 def feed_data(self, data):179 self.image = data['image'].to(self.device)180 self.texture_mask = data['texture_mask'].float().to(self.device)181 self.get_gt_indices()182 183 self.texture_tokens = F.interpolate(184 self.texture_mask, size=(32, 16),185 mode='nearest').view(self.image.size(0), -1).long()186 187 def bot_encode(self, x, mask):188 h = self.bot_encoder(x)189 h = self.bot_quant_conv(h)190 _, _, (_, _, indices_list) = self.bot_quantize(h, mask)191 192 return indices_list193 194 def get_gt_indices(self):195 self.quant_t, self.feature_t = self.top_encode(self.image,196 self.texture_mask)197 self.gt_indices_list = self.bot_encode(self.image, self.texture_mask)198 199 def index_to_image(self, index_bottom_list, texture_mask):200 quant_b = self.bot_quantize.get_codebook_entry(201 index_bottom_list, texture_mask,202 (index_bottom_list[0].size(0), index_bottom_list[0].size(1),203 index_bottom_list[0].size(2),204 self.opt["bot_z_channels"])) #.permute(0, 3, 1, 2)205 quant_b = self.bot_post_quant_conv(quant_b)206 bot_dec_res = self.bot_decoder_res(quant_b)207 208 dec = self.decoder(self.quant_t, bot_h=bot_dec_res)209 210 return dec211 212 def get_vis(self, pred_img_index, rec_img_index, texture_mask, save_path):213 rec_img = self.index_to_image(rec_img_index, texture_mask)214 pred_img = self.index_to_image(pred_img_index, texture_mask)215 216 base_img = self.decoder(self.quant_t)217 img_cat = torch.cat([218 self.image,219 rec_img,220 base_img,221 pred_img,222 ], dim=3).detach()223 img_cat = ((img_cat + 1) / 2)224 img_cat = img_cat.clamp_(0, 1)225 save_image(img_cat, save_path, nrow=1, padding=4)226 227 def optimize_parameters(self):228 self.guidance_encoder.train()229 self.index_decoder.train()230 231 self.feature_enc = self.guidance_encoder(self.feature_t)232 self.memory_logits_list = self.index_decoder(self.feature_enc)233 234 loss = 0235 for i in range(18):236 loss += self.loss_func(237 self.memory_logits_list[i],238 self.gt_indices_list[i],239 ignore_index=-1)240 241 self.optimizer.zero_grad()242 loss.backward()243 self.optimizer.step()244 245 self.log_dict['loss_total'] = loss246 247 def inference(self, data_loader, save_dir):248 self.guidance_encoder.eval()249 self.index_decoder.eval()250 251 acc = 0252 num = 0253 254 for _, data in enumerate(data_loader):255 self.feed_data(data)256 img_name = data['img_name']257 258 num += self.image.size(0)259 260 texture_mask_flatten = self.texture_tokens.view(-1)261 min_encodings_indices_list = [262 torch.full(263 texture_mask_flatten.size(),264 fill_value=-1,265 dtype=torch.long,266 device=texture_mask_flatten.device) for _ in range(18)267 ]268 with torch.no_grad():269 self.feature_enc = self.guidance_encoder(self.feature_t)270 memory_logits_list = self.index_decoder(self.feature_enc)271 # memory_indices_pred = memory_logits.argmax(dim=1)272 batch_acc = 0273 for codebook_idx, memory_logits in enumerate(memory_logits_list):274 region_of_interest = texture_mask_flatten == codebook_idx275 if torch.sum(region_of_interest) > 0:276 memory_indices_pred = memory_logits.argmax(dim=1).view(-1)277 batch_acc += torch.sum(278 memory_indices_pred[region_of_interest] ==279 self.gt_indices_list[codebook_idx].view(280 -1)[region_of_interest])281 memory_indices_pred = memory_indices_pred282 min_encodings_indices_list[codebook_idx][283 region_of_interest] = memory_indices_pred[284 region_of_interest]285 min_encodings_indices_return_list = [286 min_encodings_indices.view(self.gt_indices_list[0].size())287 for min_encodings_indices in min_encodings_indices_list288 ]289 batch_acc = batch_acc / self.gt_indices_list[codebook_idx].numel(290 ) * self.image.size(0)291 acc += batch_acc292 self.get_vis(min_encodings_indices_return_list,293 self.gt_indices_list, self.texture_mask,294 f'{save_dir}/{img_name[0]}')295 296 self.guidance_encoder.train()297 self.index_decoder.train()298 return (acc / num).item()299 300 def load_network(self):301 checkpoint = torch.load(self.opt['pretrained_models'])302 self.guidance_encoder.load_state_dict(303 checkpoint['guidance_encoder'], strict=True)304 self.guidance_encoder.eval()305 306 self.index_decoder.load_state_dict(307 checkpoint['index_decoder'], strict=True)308 self.index_decoder.eval()309 310 def save_network(self, save_path):311 """Save networks.312 313 Args:314 net (nn.Module): Network to be saved.315 net_label (str): Network label.316 current_iter (int): Current iter number.317 """318 319 save_dict = {}320 save_dict['guidance_encoder'] = self.guidance_encoder.state_dict()321 save_dict['index_decoder'] = self.index_decoder.state_dict()322 323 torch.save(save_dict, save_path)324 325 def update_learning_rate(self, epoch):326 """Update learning rate.327 328 Args:329 current_iter (int): Current iteration.330 warmup_iter (int): Warmup iter numbers. -1 for no warmup.331 Default: -1.332 """333 lr = self.optimizer.param_groups[0]['lr']334 335 if self.opt['lr_decay'] == 'step':336 lr = self.opt['lr'] * (337 self.opt['gamma']**(epoch // self.opt['step']))338 elif self.opt['lr_decay'] == 'cos':339 lr = self.opt['lr'] * (340 1 + math.cos(math.pi * epoch / self.opt['num_epochs'])) / 2341 elif self.opt['lr_decay'] == 'linear':342 lr = self.opt['lr'] * (1 - epoch / self.opt['num_epochs'])343 elif self.opt['lr_decay'] == 'linear2exp':344 if epoch < self.opt['turning_point'] + 1:345 # learning rate decay as 95%346 # at the turning point (1 / 95% = 1.0526)347 lr = self.opt['lr'] * (348 1 - epoch / int(self.opt['turning_point'] * 1.0526))349 else:350 lr *= self.opt['gamma']351 elif self.opt['lr_decay'] == 'schedule':352 if epoch in self.opt['schedule']:353 lr *= self.opt['gamma']354 else:355 raise ValueError('Unknown lr mode {}'.format(self.opt['lr_decay']))356 # set learning rate357 for param_group in self.optimizer.param_groups:358 param_group['lr'] = lr359 360 return lr361 362 def get_current_log(self):363 return self.log_dict364 