Team Ai
Apppublic

radames/Text2Human-API

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
sample_model.py501 linesDownload Raw Back to models
1import logging2 3import numpy as np4import torch5import torch.distributions as dists6import torch.nn.functional as F7from torchvision.utils import save_image8 9from models.archs.fcn_arch import FCNHead, MultiHeadFCNHead10from models.archs.shape_attr_embedding_arch import ShapeAttrEmbedding11from models.archs.transformer_arch import TransformerMultiHead12from models.archs.unet_arch import ShapeUNet, UNet13from models.archs.vqgan_arch import (Decoder, DecoderRes, Encoder,14                                     VectorQuantizer,15                                     VectorQuantizerSpatialTextureAware,16                                     VectorQuantizerTexture)17 18logger = logging.getLogger('base')19 20 21class BaseSampleModel():22    """Base Model"""23 24    def __init__(self, opt):25        self.opt = opt26        self.device = torch.device('cuda')27 28        # hierarchical VQVAE29        self.decoder = Decoder(30            in_channels=opt['top_in_channels'],31            resolution=opt['top_resolution'],32            z_channels=opt['top_z_channels'],33            ch=opt['top_ch'],34            out_ch=opt['top_out_ch'],35            num_res_blocks=opt['top_num_res_blocks'],36            attn_resolutions=opt['top_attn_resolutions'],37            ch_mult=opt['top_ch_mult'],38            dropout=opt['top_dropout'],39            resamp_with_conv=True,40            give_pre_end=False).to(self.device)41        self.top_quantize = VectorQuantizerTexture(42            1024, opt['embed_dim'], beta=0.25).to(self.device)43        self.top_post_quant_conv = torch.nn.Conv2d(opt['embed_dim'],44                                                   opt["top_z_channels"],45                                                   1).to(self.device)46        self.load_top_pretrain_models()47 48        self.bot_decoder_res = DecoderRes(49            in_channels=opt['bot_in_channels'],50            resolution=opt['bot_resolution'],51            z_channels=opt['bot_z_channels'],52            ch=opt['bot_ch'],53            num_res_blocks=opt['bot_num_res_blocks'],54            ch_mult=opt['bot_ch_mult'],55            dropout=opt['bot_dropout'],56            give_pre_end=False).to(self.device)57        self.bot_quantize = VectorQuantizerSpatialTextureAware(58            opt['bot_n_embed'],59            opt['embed_dim'],60            beta=0.25,61            spatial_size=opt['bot_codebook_spatial_size']).to(self.device)62        self.bot_post_quant_conv = torch.nn.Conv2d(opt['embed_dim'],63                                                   opt["bot_z_channels"],64                                                   1).to(self.device)65        self.load_bot_pretrain_network()66 67        # top -> bot prediction68        self.index_pred_guidance_encoder = UNet(69            in_channels=opt['index_pred_encoder_in_channels']).to(self.device)70        self.index_pred_decoder = MultiHeadFCNHead(71            in_channels=opt['index_pred_fc_in_channels'],72            in_index=opt['index_pred_fc_in_index'],73            channels=opt['index_pred_fc_channels'],74            num_convs=opt['index_pred_fc_num_convs'],75            concat_input=opt['index_pred_fc_concat_input'],76            dropout_ratio=opt['index_pred_fc_dropout_ratio'],77            num_classes=opt['index_pred_fc_num_classes'],78            align_corners=opt['index_pred_fc_align_corners'],79            num_head=18).to(self.device)80        self.load_index_pred_network()81 82        # VAE for segmentation mask83        self.segm_encoder = Encoder(84            ch=opt['segm_ch'],85            num_res_blocks=opt['segm_num_res_blocks'],86            attn_resolutions=opt['segm_attn_resolutions'],87            ch_mult=opt['segm_ch_mult'],88            in_channels=opt['segm_in_channels'],89            resolution=opt['segm_resolution'],90            z_channels=opt['segm_z_channels'],91            double_z=opt['segm_double_z'],92            dropout=opt['segm_dropout']).to(self.device)93        self.segm_quantizer = VectorQuantizer(94            opt['segm_n_embed'],95            opt['segm_embed_dim'],96            beta=0.25,97            sane_index_shape=True).to(self.device)98        self.segm_quant_conv = torch.nn.Conv2d(opt["segm_z_channels"],99                                               opt['segm_embed_dim'],100                                               1).to(self.device)101        self.load_pretrained_segm_token()102 103        # define sampler104        self.sampler_fn = TransformerMultiHead(105            codebook_size=opt['codebook_size'],106            segm_codebook_size=opt['segm_codebook_size'],107            texture_codebook_size=opt['texture_codebook_size'],108            bert_n_emb=opt['bert_n_emb'],109            bert_n_layers=opt['bert_n_layers'],110            bert_n_head=opt['bert_n_head'],111            block_size=opt['block_size'],112            latent_shape=opt['latent_shape'],113            embd_pdrop=opt['embd_pdrop'],114            resid_pdrop=opt['resid_pdrop'],115            attn_pdrop=opt['attn_pdrop'],116            num_head=opt['num_head']).to(self.device)117        self.load_sampler_pretrained_network()118 119        self.shape = tuple(opt['latent_shape'])120 121        self.mask_id = opt['codebook_size']122        self.sample_steps = opt['sample_steps']123 124    def load_top_pretrain_models(self):125        # load pretrained vqgan126        top_vae_checkpoint = torch.load(self.opt['top_vae_path'])127 128        self.decoder.load_state_dict(129            top_vae_checkpoint['decoder'], strict=True)130        self.top_quantize.load_state_dict(131            top_vae_checkpoint['quantize'], strict=True)132        self.top_post_quant_conv.load_state_dict(133            top_vae_checkpoint['post_quant_conv'], strict=True)134 135        self.decoder.eval()136        self.top_quantize.eval()137        self.top_post_quant_conv.eval()138 139    def load_bot_pretrain_network(self):140        checkpoint = torch.load(self.opt['bot_vae_path'])141        self.bot_decoder_res.load_state_dict(142            checkpoint['bot_decoder_res'], strict=True)143        self.decoder.load_state_dict(checkpoint['decoder'], strict=True)144        self.bot_quantize.load_state_dict(145            checkpoint['bot_quantize'], strict=True)146        self.bot_post_quant_conv.load_state_dict(147            checkpoint['bot_post_quant_conv'], strict=True)148 149        self.bot_decoder_res.eval()150        self.decoder.eval()151        self.bot_quantize.eval()152        self.bot_post_quant_conv.eval()153 154    def load_pretrained_segm_token(self):155        # load pretrained vqgan for segmentation mask156        segm_token_checkpoint = torch.load(self.opt['segm_token_path'])157        self.segm_encoder.load_state_dict(158            segm_token_checkpoint['encoder'], strict=True)159        self.segm_quantizer.load_state_dict(160            segm_token_checkpoint['quantize'], strict=True)161        self.segm_quant_conv.load_state_dict(162            segm_token_checkpoint['quant_conv'], strict=True)163 164        self.segm_encoder.eval()165        self.segm_quantizer.eval()166        self.segm_quant_conv.eval()167 168    def load_index_pred_network(self):169        checkpoint = torch.load(self.opt['pretrained_index_network'])170        self.index_pred_guidance_encoder.load_state_dict(171            checkpoint['guidance_encoder'], strict=True)172        self.index_pred_decoder.load_state_dict(173            checkpoint['index_decoder'], strict=True)174 175        self.index_pred_guidance_encoder.eval()176        self.index_pred_decoder.eval()177 178    def load_sampler_pretrained_network(self):179        checkpoint = torch.load(self.opt['pretrained_sampler'])180        self.sampler_fn.load_state_dict(checkpoint, strict=True)181        self.sampler_fn.eval()182 183    def bot_index_prediction(self, feature_top, texture_mask):184        self.index_pred_guidance_encoder.eval()185        self.index_pred_decoder.eval()186 187        texture_tokens = F.interpolate(188            texture_mask, (32, 16), mode='nearest').view(self.batch_size,189                                                         -1).long()190 191        texture_mask_flatten = texture_tokens.view(-1)192        min_encodings_indices_list = [193            torch.full(194                texture_mask_flatten.size(),195                fill_value=-1,196                dtype=torch.long,197                device=texture_mask_flatten.device) for _ in range(18)198        ]199        with torch.no_grad():200            feature_enc = self.index_pred_guidance_encoder(feature_top)201            memory_logits_list = self.index_pred_decoder(feature_enc)202            for codebook_idx, memory_logits in enumerate(memory_logits_list):203                region_of_interest = texture_mask_flatten == codebook_idx204                if torch.sum(region_of_interest) > 0:205                    memory_indices_pred = memory_logits.argmax(dim=1).view(-1)206                    memory_indices_pred = memory_indices_pred207                    min_encodings_indices_list[codebook_idx][208                        region_of_interest] = memory_indices_pred[209                            region_of_interest]210            min_encodings_indices_return_list = [211                min_encodings_indices.view((1, 32, 16))212                for min_encodings_indices in min_encodings_indices_list213            ]214 215        return min_encodings_indices_return_list216 217    def sample_and_refine(self, save_dir=None, img_name=None):218        # sample 32x16 features indices219        sampled_top_indices_list = self.sample_fn(220            temp=1, sample_steps=self.sample_steps)221 222        for sample_idx in range(self.batch_size):223            sample_indices = [224                sampled_indices_cur[sample_idx:sample_idx + 1]225                for sampled_indices_cur in sampled_top_indices_list226            ]227            top_quant = self.top_quantize.get_codebook_entry(228                sample_indices, self.texture_mask[sample_idx:sample_idx + 1],229                (sample_indices[0].size(0), self.shape[0], self.shape[1],230                 self.opt["top_z_channels"]))231 232            top_quant = self.top_post_quant_conv(top_quant)233 234            bot_indices_list = self.bot_index_prediction(235                top_quant, self.texture_mask[sample_idx:sample_idx + 1])236 237            quant_bot = self.bot_quantize.get_codebook_entry(238                bot_indices_list, self.texture_mask[sample_idx:sample_idx + 1],239                (bot_indices_list[0].size(0), bot_indices_list[0].size(1),240                 bot_indices_list[0].size(2),241                 self.opt["bot_z_channels"]))  #.permute(0, 3, 1, 2)242            quant_bot = self.bot_post_quant_conv(quant_bot)243            bot_dec_res = self.bot_decoder_res(quant_bot)244 245            dec = self.decoder(top_quant, bot_h=bot_dec_res)246 247            dec = ((dec + 1) / 2)248            dec = dec.clamp_(0, 1)249            if save_dir is None and img_name is None:250                return dec251            else:252                save_image(253                    dec,254                    f'{save_dir}/{img_name[sample_idx]}',255                    nrow=1,256                    padding=4)257 258    def sample_fn(self, temp=1.0, sample_steps=None):259        self.sampler_fn.eval()260 261        x_t = torch.ones((self.batch_size, np.prod(self.shape)),262                         device=self.device).long() * self.mask_id263        unmasked = torch.zeros_like(x_t, device=self.device).bool()264        sample_steps = list(range(1, sample_steps + 1))265 266        texture_tokens = F.interpolate(267            self.texture_mask, (32, 16),268            mode='nearest').view(self.batch_size, -1).long()269 270        texture_mask_flatten = texture_tokens.view(-1)271 272        # min_encodings_indices_list would be used to visualize the image273        min_encodings_indices_list = [274            torch.full(275                texture_mask_flatten.size(),276                fill_value=-1,277                dtype=torch.long,278                device=texture_mask_flatten.device) for _ in range(18)279        ]280 281        for t in reversed(sample_steps):282            t = torch.full((self.batch_size, ),283                           t,284                           device=self.device,285                           dtype=torch.long)286 287            # where to unmask288            changes = torch.rand(289                x_t.shape, device=self.device) < 1 / t.float().unsqueeze(-1)290            # don't unmask somewhere already unmasked291            changes = torch.bitwise_xor(changes,292                                        torch.bitwise_and(changes, unmasked))293            # update mask with changes294            unmasked = torch.bitwise_or(unmasked, changes)295 296            x_0_logits_list = self.sampler_fn(297                x_t, self.segm_tokens, texture_tokens, t=t)298 299            changes_flatten = changes.view(-1)300            ori_shape = x_t.shape  # [b, h*w]301            x_t = x_t.view(-1)  # [b*h*w]302            for codebook_idx, x_0_logits in enumerate(x_0_logits_list):303                if torch.sum(texture_mask_flatten[changes_flatten] ==304                             codebook_idx) > 0:305                    # scale by temperature306                    x_0_logits = x_0_logits / temp307                    x_0_dist = dists.Categorical(logits=x_0_logits)308                    x_0_hat = x_0_dist.sample().long()309                    x_0_hat = x_0_hat.view(-1)310 311                    # only replace the changed indices with corresponding codebook_idx312                    changes_segm = torch.bitwise_and(313                        changes_flatten, texture_mask_flatten == codebook_idx)314 315                    # x_t would be the input to the transformer, so the index range should be continual one316                    x_t[changes_segm] = x_0_hat[317                        changes_segm] + 1024 * codebook_idx318                    min_encodings_indices_list[codebook_idx][319                        changes_segm] = x_0_hat[changes_segm]320 321            x_t = x_t.view(ori_shape)  # [b, h*w]322 323        min_encodings_indices_return_list = [324            min_encodings_indices.view(ori_shape)325            for min_encodings_indices in min_encodings_indices_list326        ]327 328        self.sampler_fn.train()329 330        return min_encodings_indices_return_list331 332    @torch.no_grad()333    def get_quantized_segm(self, segm):334        segm_one_hot = F.one_hot(335            segm.squeeze(1).long(),336            num_classes=self.opt['segm_num_segm_classes']).permute(337                0, 3, 1, 2).to(memory_format=torch.contiguous_format).float()338        encoded_segm_mask = self.segm_encoder(segm_one_hot)339        encoded_segm_mask = self.segm_quant_conv(encoded_segm_mask)340        _, _, [_, _, segm_tokens] = self.segm_quantizer(encoded_segm_mask)341 342        return segm_tokens343 344 345class SampleFromParsingModel(BaseSampleModel):346    """SampleFromParsing model.347    """348 349    def feed_data(self, data):350        self.segm = data['segm'].to(self.device)351        self.texture_mask = data['texture_mask'].to(self.device)352        self.batch_size = self.segm.size(0)353 354        self.segm_tokens = self.get_quantized_segm(self.segm)355        self.segm_tokens = self.segm_tokens.view(self.batch_size, -1)356 357    def inference(self, data_loader, save_dir):358        for _, data in enumerate(data_loader):359            img_name = data['img_name']360            self.feed_data(data)361            with torch.no_grad():362                self.sample_and_refine(save_dir, img_name)363 364 365class SampleFromPoseModel(BaseSampleModel):366    """SampleFromPose model.367    """368 369    def __init__(self, opt):370        super().__init__(opt)371        # pose-to-parsing372        self.shape_attr_embedder = ShapeAttrEmbedding(373            dim=opt['shape_embedder_dim'],374            out_dim=opt['shape_embedder_out_dim'],375            cls_num_list=opt['shape_attr_class_num']).to(self.device)376        self.shape_parsing_encoder = ShapeUNet(377            in_channels=opt['shape_encoder_in_channels']).to(self.device)378        self.shape_parsing_decoder = FCNHead(379            in_channels=opt['shape_fc_in_channels'],380            in_index=opt['shape_fc_in_index'],381            channels=opt['shape_fc_channels'],382            num_convs=opt['shape_fc_num_convs'],383            concat_input=opt['shape_fc_concat_input'],384            dropout_ratio=opt['shape_fc_dropout_ratio'],385            num_classes=opt['shape_fc_num_classes'],386            align_corners=opt['shape_fc_align_corners'],387        ).to(self.device)388        self.load_shape_generation_models()389 390        self.palette = [[0, 0, 0], [255, 250, 250], [220, 220, 220],391                        [250, 235, 215], [255, 250, 205], [211, 211, 211],392                        [70, 130, 180], [127, 255, 212], [0, 100, 0],393                        [50, 205, 50], [255, 255, 0], [245, 222, 179],394                        [255, 140, 0], [255, 0, 0], [16, 78, 139],395                        [144, 238, 144], [50, 205, 174], [50, 155, 250],396                        [160, 140, 88], [213, 140, 88], [90, 140, 90],397                        [185, 210, 205], [130, 165, 180], [225, 141, 151]]398 399    def load_shape_generation_models(self):400        checkpoint = torch.load(self.opt['pretrained_parsing_gen'])401 402        self.shape_attr_embedder.load_state_dict(403            checkpoint['embedder'], strict=True)404        self.shape_attr_embedder.eval()405 406        self.shape_parsing_encoder.load_state_dict(407            checkpoint['encoder'], strict=True)408        self.shape_parsing_encoder.eval()409 410        self.shape_parsing_decoder.load_state_dict(411            checkpoint['decoder'], strict=True)412        self.shape_parsing_decoder.eval()413 414    def feed_data(self, data):415        self.pose = data['densepose'].to(self.device)416        self.batch_size = self.pose.size(0)417 418        self.shape_attr = data['shape_attr'].to(self.device)419        self.upper_fused_attr = data['upper_fused_attr'].to(self.device)420        self.lower_fused_attr = data['lower_fused_attr'].to(self.device)421        self.outer_fused_attr = data['outer_fused_attr'].to(self.device)422 423    def inference(self, data_loader, save_dir):424        for _, data in enumerate(data_loader):425            img_name = data['img_name']426            self.feed_data(data)427            with torch.no_grad():428                self.generate_parsing_map()429                self.generate_quantized_segm()430                self.generate_texture_map()431                self.sample_and_refine(save_dir, img_name)432 433    def generate_parsing_map(self):434        with torch.no_grad():435            attr_embedding = self.shape_attr_embedder(self.shape_attr)436            pose_enc = self.shape_parsing_encoder(self.pose, attr_embedding)437            seg_logits = self.shape_parsing_decoder(pose_enc)438        self.segm = seg_logits.argmax(dim=1)439        self.segm = self.segm.unsqueeze(1)440 441    def generate_quantized_segm(self):442        self.segm_tokens = self.get_quantized_segm(self.segm)443        self.segm_tokens = self.segm_tokens.view(self.batch_size, -1)444 445    def generate_texture_map(self):446        upper_cls = [1., 4.]447        lower_cls = [3., 5., 21.]448        outer_cls = [2.]449 450        mask_batch = []451        for idx in range(self.batch_size):452            mask = torch.zeros_like(self.segm[idx])453            upper_fused_attr = self.upper_fused_attr[idx]454            lower_fused_attr = self.lower_fused_attr[idx]455            outer_fused_attr = self.outer_fused_attr[idx]456            if upper_fused_attr != 17:457                for cls in upper_cls:458                    mask[self.segm[idx] == cls] = upper_fused_attr + 1459 460            if lower_fused_attr != 17:461                for cls in lower_cls:462                    mask[self.segm[idx] == cls] = lower_fused_attr + 1463 464            if outer_fused_attr != 17:465                for cls in outer_cls:466                    mask[self.segm[idx] == cls] = outer_fused_attr + 1467 468            mask_batch.append(mask)469        self.texture_mask = torch.stack(mask_batch, dim=0).to(torch.float32)470 471    def feed_pose_data(self, pose_img):472        # for ui demo473 474        self.pose = pose_img.to(self.device)475        self.batch_size = self.pose.size(0)476 477    def feed_shape_attributes(self, shape_attr):478        # for ui demo479 480        self.shape_attr = shape_attr.to(self.device)481 482    def feed_texture_attributes(self, texture_attr):483        # for ui demo484 485        self.upper_fused_attr = texture_attr[0].unsqueeze(0).to(self.device)486        self.lower_fused_attr = texture_attr[1].unsqueeze(0).to(self.device)487        self.outer_fused_attr = texture_attr[2].unsqueeze(0).to(self.device)488 489    def palette_result(self, result):490 491        seg = result[0]492        palette = np.array(self.palette)493        assert palette.shape[1] == 3494        assert len(palette.shape) == 2495        color_seg = np.zeros((seg.shape[0], seg.shape[1], 3), dtype=np.uint8)496        for label, color in enumerate(palette):497            color_seg[seg == label, :] = color498        # convert to BGR499        # color_seg = color_seg[..., ::-1]500        return color_seg501