Team Ai
Apppublic

radames/Text2Human-API

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
transformer_model.py483 linesDownload Raw Back to models
1import logging2import math3from collections import OrderedDict4 5import numpy as np6import torch7import torch.distributions as dists8import torch.nn.functional as F9from torchvision.utils import save_image10 11from models.archs.transformer_arch import TransformerMultiHead12from models.archs.vqgan_arch import (Decoder, Encoder, VectorQuantizer,13                                     VectorQuantizerTexture)14 15logger = logging.getLogger('base')16 17 18class TransformerTextureAwareModel():19    """Texture-Aware Diffusion based Transformer model.20    """21 22    def __init__(self, opt):23        self.opt = opt24        self.device = torch.device('cuda')25        self.is_train = opt['is_train']26 27        # VQVAE for image28        self.img_encoder = Encoder(29            ch=opt['img_ch'],30            num_res_blocks=opt['img_num_res_blocks'],31            attn_resolutions=opt['img_attn_resolutions'],32            ch_mult=opt['img_ch_mult'],33            in_channels=opt['img_in_channels'],34            resolution=opt['img_resolution'],35            z_channels=opt['img_z_channels'],36            double_z=opt['img_double_z'],37            dropout=opt['img_dropout']).to(self.device)38        self.img_decoder = Decoder(39            in_channels=opt['img_in_channels'],40            resolution=opt['img_resolution'],41            z_channels=opt['img_z_channels'],42            ch=opt['img_ch'],43            out_ch=opt['img_out_ch'],44            num_res_blocks=opt['img_num_res_blocks'],45            attn_resolutions=opt['img_attn_resolutions'],46            ch_mult=opt['img_ch_mult'],47            dropout=opt['img_dropout'],48            resamp_with_conv=True,49            give_pre_end=False).to(self.device)50        self.img_quantizer = VectorQuantizerTexture(51            opt['img_n_embed'], opt['img_embed_dim'],52            beta=0.25).to(self.device)53        self.img_quant_conv = torch.nn.Conv2d(opt["img_z_channels"],54                                              opt['img_embed_dim'],55                                              1).to(self.device)56        self.img_post_quant_conv = torch.nn.Conv2d(opt['img_embed_dim'],57                                                   opt["img_z_channels"],58                                                   1).to(self.device)59        self.load_pretrained_image_vae()60 61        # VAE for segmentation mask62        self.segm_encoder = Encoder(63            ch=opt['segm_ch'],64            num_res_blocks=opt['segm_num_res_blocks'],65            attn_resolutions=opt['segm_attn_resolutions'],66            ch_mult=opt['segm_ch_mult'],67            in_channels=opt['segm_in_channels'],68            resolution=opt['segm_resolution'],69            z_channels=opt['segm_z_channels'],70            double_z=opt['segm_double_z'],71            dropout=opt['segm_dropout']).to(self.device)72        self.segm_quantizer = VectorQuantizer(73            opt['segm_n_embed'],74            opt['segm_embed_dim'],75            beta=0.25,76            sane_index_shape=True).to(self.device)77        self.segm_quant_conv = torch.nn.Conv2d(opt["segm_z_channels"],78                                               opt['segm_embed_dim'],79                                               1).to(self.device)80        self.load_pretrained_segm_vae()81 82        # define sampler83        self._denoise_fn = TransformerMultiHead(84            codebook_size=opt['codebook_size'],85            segm_codebook_size=opt['segm_codebook_size'],86            texture_codebook_size=opt['texture_codebook_size'],87            bert_n_emb=opt['bert_n_emb'],88            bert_n_layers=opt['bert_n_layers'],89            bert_n_head=opt['bert_n_head'],90            block_size=opt['block_size'],91            latent_shape=opt['latent_shape'],92            embd_pdrop=opt['embd_pdrop'],93            resid_pdrop=opt['resid_pdrop'],94            attn_pdrop=opt['attn_pdrop'],95            num_head=opt['num_head']).to(self.device)96 97        self.num_classes = opt['codebook_size']98        self.shape = tuple(opt['latent_shape'])99        self.num_timesteps = 1000100 101        self.mask_id = opt['codebook_size']102        self.loss_type = opt['loss_type']103        self.mask_schedule = opt['mask_schedule']104 105        self.sample_steps = opt['sample_steps']106 107        self.init_training_settings()108 109    def load_pretrained_image_vae(self):110        # load pretrained vqgan for segmentation mask111        img_ae_checkpoint = torch.load(self.opt['img_ae_path'])112        self.img_encoder.load_state_dict(113            img_ae_checkpoint['encoder'], strict=True)114        self.img_decoder.load_state_dict(115            img_ae_checkpoint['decoder'], strict=True)116        self.img_quantizer.load_state_dict(117            img_ae_checkpoint['quantize'], strict=True)118        self.img_quant_conv.load_state_dict(119            img_ae_checkpoint['quant_conv'], strict=True)120        self.img_post_quant_conv.load_state_dict(121            img_ae_checkpoint['post_quant_conv'], strict=True)122        self.img_encoder.eval()123        self.img_decoder.eval()124        self.img_quantizer.eval()125        self.img_quant_conv.eval()126        self.img_post_quant_conv.eval()127 128    def load_pretrained_segm_vae(self):129        # load pretrained vqgan for segmentation mask130        segm_ae_checkpoint = torch.load(self.opt['segm_ae_path'])131        self.segm_encoder.load_state_dict(132            segm_ae_checkpoint['encoder'], strict=True)133        self.segm_quantizer.load_state_dict(134            segm_ae_checkpoint['quantize'], strict=True)135        self.segm_quant_conv.load_state_dict(136            segm_ae_checkpoint['quant_conv'], strict=True)137        self.segm_encoder.eval()138        self.segm_quantizer.eval()139        self.segm_quant_conv.eval()140 141    def init_training_settings(self):142        optim_params = []143        for v in self._denoise_fn.parameters():144            if v.requires_grad:145                optim_params.append(v)146        # set up optimizer147        self.optimizer = torch.optim.Adam(148            optim_params,149            self.opt['lr'],150            weight_decay=self.opt['weight_decay'])151        self.log_dict = OrderedDict()152 153    @torch.no_grad()154    def get_quantized_img(self, image, texture_mask):155        encoded_img = self.img_encoder(image)156        encoded_img = self.img_quant_conv(encoded_img)157 158        # img_tokens_input is the continual index for the input of transformer159        # img_tokens_gt_list is the index for 18 texture-aware codebooks respectively160        _, _, [_, img_tokens_input, img_tokens_gt_list161               ] = self.img_quantizer(encoded_img, texture_mask)162 163        # reshape the tokens164        b = image.size(0)165        img_tokens_input = img_tokens_input.view(b, -1)166        img_tokens_gt_return_list = [167            img_tokens_gt.view(b, -1) for img_tokens_gt in img_tokens_gt_list168        ]169 170        return img_tokens_input, img_tokens_gt_return_list171 172    @torch.no_grad()173    def decode(self, quant):174        quant = self.img_post_quant_conv(quant)175        dec = self.img_decoder(quant)176        return dec177 178    @torch.no_grad()179    def decode_image_indices(self, indices_list, texture_mask):180        quant = self.img_quantizer.get_codebook_entry(181            indices_list, texture_mask,182            (indices_list[0].size(0), self.shape[0], self.shape[1],183             self.opt["img_z_channels"]))184        dec = self.decode(quant)185 186        return dec187 188    def sample_time(self, b, device, method='uniform'):189        if method == 'importance':190            if not (self.Lt_count > 10).all():191                return self.sample_time(b, device, method='uniform')192 193            Lt_sqrt = torch.sqrt(self.Lt_history + 1e-10) + 0.0001194            Lt_sqrt[0] = Lt_sqrt[1]  # Overwrite decoder term with L1.195            pt_all = Lt_sqrt / Lt_sqrt.sum()196 197            t = torch.multinomial(pt_all, num_samples=b, replacement=True)198 199            pt = pt_all.gather(dim=0, index=t)200 201            return t, pt202 203        elif method == 'uniform':204            t = torch.randint(205                1, self.num_timesteps + 1, (b, ), device=device).long()206            pt = torch.ones_like(t).float() / self.num_timesteps207            return t, pt208 209        else:210            raise ValueError211 212    def q_sample(self, x_0, x_0_gt_list, t):213        # samples q(x_t | x_0)214        # randomly set token to mask with probability t/T215        # x_t, x_0_ignore = x_0.clone(), x_0.clone()216        x_t = x_0.clone()217 218        mask = torch.rand_like(x_t.float()) < (219            t.float().unsqueeze(-1) / self.num_timesteps)220        x_t[mask] = self.mask_id221        # x_0_ignore[torch.bitwise_not(mask)] = -1222 223        # for every gt token list, we also need to do the mask224        x_0_gt_ignore_list = []225        for x_0_gt in x_0_gt_list:226            x_0_gt_ignore = x_0_gt.clone()227            x_0_gt_ignore[torch.bitwise_not(mask)] = -1228            x_0_gt_ignore_list.append(x_0_gt_ignore)229 230        return x_t, x_0_gt_ignore_list, mask231 232    def _train_loss(self, x_0, x_0_gt_list):233        b, device = x_0.size(0), x_0.device234 235        # choose what time steps to compute loss at236        t, pt = self.sample_time(b, device, 'uniform')237 238        # make x noisy and denoise239        if self.mask_schedule == 'random':240            x_t, x_0_gt_ignore_list, mask = self.q_sample(241                x_0=x_0, x_0_gt_list=x_0_gt_list, t=t)242        else:243            raise NotImplementedError244 245        # sample p(x_0 | x_t)246        x_0_hat_logits_list = self._denoise_fn(247            x_t, self.segm_tokens, self.texture_tokens, t=t)248 249        # Always compute ELBO for comparison purposes250        cross_entropy_loss = 0251        for x_0_hat_logits, x_0_gt_ignore in zip(x_0_hat_logits_list,252                                                 x_0_gt_ignore_list):253            cross_entropy_loss += F.cross_entropy(254                x_0_hat_logits.permute(0, 2, 1),255                x_0_gt_ignore,256                ignore_index=-1,257                reduction='none').sum(1)258        vb_loss = cross_entropy_loss / t259        vb_loss = vb_loss / pt260        vb_loss = vb_loss / (math.log(2) * x_0.shape[1:].numel())261        if self.loss_type == 'elbo':262            loss = vb_loss263        elif self.loss_type == 'mlm':264            denom = mask.float().sum(1)265            denom[denom == 0] = 1  # prevent divide by 0 errors.266            loss = cross_entropy_loss / denom267        elif self.loss_type == 'reweighted_elbo':268            weight = (1 - (t / self.num_timesteps))269            loss = weight * cross_entropy_loss270            loss = loss / (math.log(2) * x_0.shape[1:].numel())271        else:272            raise ValueError273 274        return loss.mean(), vb_loss.mean()275 276    def feed_data(self, data):277        self.image = data['image'].to(self.device)278        self.segm = data['segm'].to(self.device)279        self.texture_mask = data['texture_mask'].to(self.device)280        self.input_indices, self.gt_indices_list = self.get_quantized_img(281            self.image, self.texture_mask)282 283        self.texture_tokens = F.interpolate(284            self.texture_mask, size=self.shape,285            mode='nearest').view(self.image.size(0), -1).long()286 287        self.segm_tokens = self.get_quantized_segm(self.segm)288        self.segm_tokens = self.segm_tokens.view(self.image.size(0), -1)289 290    def optimize_parameters(self):291        self._denoise_fn.train()292 293        loss, vb_loss = self._train_loss(self.input_indices,294                                         self.gt_indices_list)295 296        self.optimizer.zero_grad()297        loss.backward()298        self.optimizer.step()299 300        self.log_dict['loss'] = loss301        self.log_dict['vb_loss'] = vb_loss302 303        self._denoise_fn.eval()304 305    @torch.no_grad()306    def get_quantized_segm(self, segm):307        segm_one_hot = F.one_hot(308            segm.squeeze(1).long(),309            num_classes=self.opt['segm_num_segm_classes']).permute(310                0, 3, 1, 2).to(memory_format=torch.contiguous_format).float()311        encoded_segm_mask = self.segm_encoder(segm_one_hot)312        encoded_segm_mask = self.segm_quant_conv(encoded_segm_mask)313        _, _, [_, _, segm_tokens] = self.segm_quantizer(encoded_segm_mask)314 315        return segm_tokens316 317    def sample_fn(self, temp=1.0, sample_steps=None):318        self._denoise_fn.eval()319 320        b, device = self.image.size(0), 'cuda'321        x_t = torch.ones(322            (b, np.prod(self.shape)), device=device).long() * self.mask_id323        unmasked = torch.zeros_like(x_t, device=device).bool()324        sample_steps = list(range(1, sample_steps + 1))325 326        texture_mask_flatten = self.texture_tokens.view(-1)327 328        # min_encodings_indices_list would be used to visualize the image329        min_encodings_indices_list = [330            torch.full(331                texture_mask_flatten.size(),332                fill_value=-1,333                dtype=torch.long,334                device=texture_mask_flatten.device) for _ in range(18)335        ]336 337        for t in reversed(sample_steps):338            print(f'Sample timestep {t:4d}', end='\r')339            t = torch.full((b, ), t, device=device, dtype=torch.long)340 341            # where to unmask342            changes = torch.rand(343                x_t.shape, device=device) < 1 / t.float().unsqueeze(-1)344            # don't unmask somewhere already unmasked345            changes = torch.bitwise_xor(changes,346                                        torch.bitwise_and(changes, unmasked))347            # update mask with changes348            unmasked = torch.bitwise_or(unmasked, changes)349 350            x_0_logits_list = self._denoise_fn(351                x_t, self.segm_tokens, self.texture_tokens, t=t)352 353            changes_flatten = changes.view(-1)354            ori_shape = x_t.shape  # [b, h*w]355            x_t = x_t.view(-1)  # [b*h*w]356            for codebook_idx, x_0_logits in enumerate(x_0_logits_list):357                if torch.sum(texture_mask_flatten[changes_flatten] ==358                             codebook_idx) > 0:359                    # scale by temperature360                    x_0_logits = x_0_logits / temp361                    x_0_dist = dists.Categorical(logits=x_0_logits)362                    x_0_hat = x_0_dist.sample().long()363                    x_0_hat = x_0_hat.view(-1)364 365                    # only replace the changed indices with corresponding codebook_idx366                    changes_segm = torch.bitwise_and(367                        changes_flatten, texture_mask_flatten == codebook_idx)368 369                    # x_t would be the input to the transformer, so the index range should be continual one370                    x_t[changes_segm] = x_0_hat[371                        changes_segm] + 1024 * codebook_idx372                    min_encodings_indices_list[codebook_idx][373                        changes_segm] = x_0_hat[changes_segm]374 375            x_t = x_t.view(ori_shape)  # [b, h*w]376 377        min_encodings_indices_return_list = [378            min_encodings_indices.view(ori_shape)379            for min_encodings_indices in min_encodings_indices_list380        ]381 382        self._denoise_fn.train()383 384        return min_encodings_indices_return_list385 386    def get_vis(self, image, gt_indices, predicted_indices, texture_mask,387                save_path):388        # original image389        ori_img = self.decode_image_indices(gt_indices, texture_mask)390        # pred image391        pred_img = self.decode_image_indices(predicted_indices, texture_mask)392        img_cat = torch.cat([393            image,394            ori_img,395            pred_img,396        ], dim=3).detach()397        img_cat = ((img_cat + 1) / 2)398        img_cat = img_cat.clamp_(0, 1)399        save_image(img_cat, save_path, nrow=1, padding=4)400 401    def inference(self, data_loader, save_dir):402        self._denoise_fn.eval()403 404        for _, data in enumerate(data_loader):405            img_name = data['img_name']406            self.feed_data(data)407            b = self.image.size(0)408            with torch.no_grad():409                sampled_indices_list = self.sample_fn(410                    temp=1, sample_steps=self.sample_steps)411            for idx in range(b):412                self.get_vis(self.image[idx:idx + 1], [413                    gt_indices[idx:idx + 1]414                    for gt_indices in self.gt_indices_list415                ], [416                    sampled_indices[idx:idx + 1]417                    for sampled_indices in sampled_indices_list418                ], self.texture_mask[idx:idx + 1],419                             f'{save_dir}/{img_name[idx]}')420 421        self._denoise_fn.train()422 423    def get_current_log(self):424        return self.log_dict425 426    def update_learning_rate(self, epoch, iters=None):427        """Update learning rate.428 429        Args:430            current_iter (int): Current iteration.431            warmup_iter (int): Warmup iter numbers. -1 for no warmup.432                Default: -1.433        """434        lr = self.optimizer.param_groups[0]['lr']435 436        if self.opt['lr_decay'] == 'step':437            lr = self.opt['lr'] * (438                self.opt['gamma']**(epoch // self.opt['step']))439        elif self.opt['lr_decay'] == 'cos':440            lr = self.opt['lr'] * (441                1 + math.cos(math.pi * epoch / self.opt['num_epochs'])) / 2442        elif self.opt['lr_decay'] == 'linear':443            lr = self.opt['lr'] * (1 - epoch / self.opt['num_epochs'])444        elif self.opt['lr_decay'] == 'linear2exp':445            if epoch < self.opt['turning_point'] + 1:446                # learning rate decay as 95%447                # at the turning point (1 / 95% = 1.0526)448                lr = self.opt['lr'] * (449                    1 - epoch / int(self.opt['turning_point'] * 1.0526))450            else:451                lr *= self.opt['gamma']452        elif self.opt['lr_decay'] == 'schedule':453            if epoch in self.opt['schedule']:454                lr *= self.opt['gamma']455        elif self.opt['lr_decay'] == 'warm_up':456            if iters <= self.opt['warmup_iters']:457                lr = self.opt['lr'] * float(iters) / self.opt['warmup_iters']458            else:459                lr = self.opt['lr']460        else:461            raise ValueError('Unknown lr mode {}'.format(self.opt['lr_decay']))462        # set learning rate463        for param_group in self.optimizer.param_groups:464            param_group['lr'] = lr465 466        return lr467 468    def save_network(self, net, save_path):469        """Save networks.470 471        Args:472            net (nn.Module): Network to be saved.473            net_label (str): Network label.474            current_iter (int): Current iter number.475        """476        state_dict = net.state_dict()477        torch.save(state_dict, save_path)478 479    def load_network(self):480        checkpoint = torch.load(self.opt['pretrained_sampler'])481        self._denoise_fn.load_state_dict(checkpoint, strict=True)482        self._denoise_fn.eval()483