CoolFace
Apppublic

RabbitRUI/ruispace

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
1import torch2import torch.nn.functional as F3from torch import nn4from src.audio2pose_models.res_unet import ResUnet5 6def class2onehot(idx, class_num):7 8    assert torch.max(idx).item() < class_num9    onehot = torch.zeros(idx.size(0), class_num).to(idx.device)10    onehot.scatter_(1, idx, 1)11    return onehot12 13class CVAE(nn.Module):14    def __init__(self, cfg):15        super().__init__()16        encoder_layer_sizes = cfg.MODEL.CVAE.ENCODER_LAYER_SIZES17        decoder_layer_sizes = cfg.MODEL.CVAE.DECODER_LAYER_SIZES18        latent_size = cfg.MODEL.CVAE.LATENT_SIZE19        num_classes = cfg.DATASET.NUM_CLASSES20        audio_emb_in_size = cfg.MODEL.CVAE.AUDIO_EMB_IN_SIZE21        audio_emb_out_size = cfg.MODEL.CVAE.AUDIO_EMB_OUT_SIZE22        seq_len = cfg.MODEL.CVAE.SEQ_LEN23 24        self.latent_size = latent_size25 26        self.encoder = ENCODER(encoder_layer_sizes, latent_size, num_classes,27                                audio_emb_in_size, audio_emb_out_size, seq_len)28        self.decoder = DECODER(decoder_layer_sizes, latent_size, num_classes,29                                audio_emb_in_size, audio_emb_out_size, seq_len)30    def reparameterize(self, mu, logvar):31        std = torch.exp(0.5 * logvar)32        eps = torch.randn_like(std)33        return mu + eps * std34 35    def forward(self, batch):36        batch = self.encoder(batch)37        mu = batch['mu']38        logvar = batch['logvar']39        z = self.reparameterize(mu, logvar)40        batch['z'] = z41        return self.decoder(batch)42 43    def test(self, batch):44        '''45        class_id = batch['class']46        z = torch.randn([class_id.size(0), self.latent_size]).to(class_id.device)47        batch['z'] = z48        '''49        return self.decoder(batch)50 51class ENCODER(nn.Module):52    def __init__(self, layer_sizes, latent_size, num_classes, 53                audio_emb_in_size, audio_emb_out_size, seq_len):54        super().__init__()55 56        self.resunet = ResUnet()57        self.num_classes = num_classes58        self.seq_len = seq_len59 60        self.MLP = nn.Sequential()61        layer_sizes[0] += latent_size + seq_len*audio_emb_out_size + 662        for i, (in_size, out_size) in enumerate(zip(layer_sizes[:-1], layer_sizes[1:])):63            self.MLP.add_module(64                name="L{:d}".format(i), module=nn.Linear(in_size, out_size))65            self.MLP.add_module(name="A{:d}".format(i), module=nn.ReLU())66 67        self.linear_means = nn.Linear(layer_sizes[-1], latent_size)68        self.linear_logvar = nn.Linear(layer_sizes[-1], latent_size)69        self.linear_audio = nn.Linear(audio_emb_in_size, audio_emb_out_size)70 71        self.classbias = nn.Parameter(torch.randn(self.num_classes, latent_size))72 73    def forward(self, batch):74        class_id = batch['class']75        pose_motion_gt = batch['pose_motion_gt']                             #bs seq_len 676        ref = batch['ref']                             #bs 677        bs = pose_motion_gt.shape[0]78        audio_in = batch['audio_emb']                          # bs seq_len audio_emb_in_size79 80        #pose encode81        pose_emb = self.resunet(pose_motion_gt.unsqueeze(1))          #bs 1 seq_len 6 82        pose_emb = pose_emb.reshape(bs, -1)                    #bs seq_len*683 84        #audio mapping85        print(audio_in.shape)86        audio_out = self.linear_audio(audio_in)                # bs seq_len audio_emb_out_size87        audio_out = audio_out.reshape(bs, -1)88 89        class_bias = self.classbias[class_id]                  #bs latent_size90        x_in = torch.cat([ref, pose_emb, audio_out, class_bias], dim=-1) #bs seq_len*(audio_emb_out_size+6)+latent_size91        x_out = self.MLP(x_in)92 93        mu = self.linear_means(x_out)94        logvar = self.linear_means(x_out)                      #bs latent_size 95 96        batch.update({'mu':mu, 'logvar':logvar})97        return batch98 99class DECODER(nn.Module):100    def __init__(self, layer_sizes, latent_size, num_classes, 101                audio_emb_in_size, audio_emb_out_size, seq_len):102        super().__init__()103 104        self.resunet = ResUnet()105        self.num_classes = num_classes106        self.seq_len = seq_len107 108        self.MLP = nn.Sequential()109        input_size = latent_size + seq_len*audio_emb_out_size + 6110        for i, (in_size, out_size) in enumerate(zip([input_size]+layer_sizes[:-1], layer_sizes)):111            self.MLP.add_module(112                name="L{:d}".format(i), module=nn.Linear(in_size, out_size))113            if i+1 < len(layer_sizes):114                self.MLP.add_module(name="A{:d}".format(i), module=nn.ReLU())115            else:116                self.MLP.add_module(name="sigmoid", module=nn.Sigmoid())117        118        self.pose_linear = nn.Linear(6, 6)119        self.linear_audio = nn.Linear(audio_emb_in_size, audio_emb_out_size)120 121        self.classbias = nn.Parameter(torch.randn(self.num_classes, latent_size))122 123    def forward(self, batch):124 125        z = batch['z']                                          #bs latent_size126        bs = z.shape[0]127        class_id = batch['class']128        ref = batch['ref']                             #bs 6129        audio_in = batch['audio_emb']                           # bs seq_len audio_emb_in_size130        #print('audio_in: ', audio_in[:, :, :10])131 132        audio_out = self.linear_audio(audio_in)                 # bs seq_len audio_emb_out_size133        #print('audio_out: ', audio_out[:, :, :10])134        audio_out = audio_out.reshape([bs, -1])                 # bs seq_len*audio_emb_out_size135        class_bias = self.classbias[class_id]                   #bs latent_size136 137        z = z + class_bias138        x_in = torch.cat([ref, z, audio_out], dim=-1)139        x_out = self.MLP(x_in)                                  # bs layer_sizes[-1]140        x_out = x_out.reshape((bs, self.seq_len, -1))141 142        #print('x_out: ', x_out)143 144        pose_emb = self.resunet(x_out.unsqueeze(1))             #bs 1 seq_len 6145 146        pose_motion_pred = self.pose_linear(pose_emb.squeeze(1))       #bs seq_len 6147 148        batch.update({'pose_motion_pred':pose_motion_pred})149        return batch150