Team Ai
Apppublic

radames/Text2Human-API

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
vqgan_model.py552 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, 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