RabbitRUI/ruispace
0
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 