Team Ai
Apppublic

radames/Text2Human-API

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
hierarchy_inference_model.py364 linesDownload Raw Back to models
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