Team Ai
Apppublic

radames/Text2Human-API

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