CoolFace
Datasetpublic

H-Liu1997/tango_cached_utils

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes230downloads
motion_encoder.py199 linesDownload Raw Back to emage
1import torch.nn as nn2import torch3import numpy as np4from .skeleton import ResidualBlock, SkeletonResidual, residual_ratio, SkeletonConv, SkeletonPool, find_neighbor, build_edge_topology5 6class LocalEncoder(nn.Module):7    def __init__(self, args, topology):8        super(LocalEncoder, self).__init__()9        args.channel_base = 610        args.activation = "tanh"11        args.use_residual_blocks=True12        args.z_dim=102413        args.temporal_scale=814        args.kernel_size=415        args.num_layers=args.vae_layer16        args.skeleton_dist=217        args.extra_conv=018        # check how to reflect in 1d19        args.padding_mode="constant"20        args.skeleton_pool="mean"21        args.upsampling="linear"22 23 24        self.topologies = [topology]25        self.channel_base = [args.channel_base]26 27        self.channel_list = []28        self.edge_num = [len(topology)]29        self.pooling_list = []30        self.layers = nn.ModuleList()31        self.args = args32        # self.convs = []33 34        kernel_size = args.kernel_size35        kernel_even = False if kernel_size % 2 else True36        padding = (kernel_size - 1) // 237        bias = True38        self.grow = args.vae_grow39        for i in range(args.num_layers):40            self.channel_base.append(self.channel_base[-1]*self.grow[i])41 42        for i in range(args.num_layers):43            seq = []44            neighbour_list = find_neighbor(self.topologies[i], args.skeleton_dist)45            in_channels = self.channel_base[i] * self.edge_num[i]46            out_channels = self.channel_base[i + 1] * self.edge_num[i]47            if i == 0:48                self.channel_list.append(in_channels)49            self.channel_list.append(out_channels)50            last_pool = True if i == args.num_layers - 1 else False51 52            # (T, J, D) => (T, J', D)53            pool = SkeletonPool(edges=self.topologies[i], pooling_mode=args.skeleton_pool,54                                channels_per_edge=out_channels // len(neighbour_list), last_pool=last_pool)55 56            if args.use_residual_blocks:57                # (T, J, D) => (T/2, J', 2D)58                seq.append(SkeletonResidual(self.topologies[i], neighbour_list, joint_num=self.edge_num[i], in_channels=in_channels, out_channels=out_channels,59                                            kernel_size=kernel_size, stride=2, padding=padding, padding_mode=args.padding_mode, bias=bias,60                                            extra_conv=args.extra_conv, pooling_mode=args.skeleton_pool, activation=args.activation, last_pool=last_pool))61            else:62                for _ in range(args.extra_conv):63                    # (T, J, D) => (T, J, D)64                    seq.append(SkeletonConv(neighbour_list, in_channels=in_channels, out_channels=in_channels,65                                            joint_num=self.edge_num[i], kernel_size=kernel_size - 1 if kernel_even else kernel_size,66                                            stride=1,67                                            padding=padding, padding_mode=args.padding_mode, bias=bias))68                    seq.append(nn.PReLU() if args.activation == 'relu' else nn.Tanh())69                # (T, J, D) => (T/2, J, 2D)70                seq.append(SkeletonConv(neighbour_list, in_channels=in_channels, out_channels=out_channels,71                                        joint_num=self.edge_num[i], kernel_size=kernel_size, stride=2,72                                        padding=padding, padding_mode=args.padding_mode, bias=bias, add_offset=False,73                                        in_offset_channel=3 * self.channel_base[i] // self.channel_base[0]))74                # self.convs.append(seq[-1])75 76                seq.append(pool)77                seq.append(nn.PReLU() if args.activation == 'relu' else nn.Tanh())78            self.layers.append(nn.Sequential(*seq))79 80            self.topologies.append(pool.new_edges)81            self.pooling_list.append(pool.pooling_list)82            self.edge_num.append(len(self.topologies[-1]))83 84        # in_features = self.channel_base[-1] * len(self.pooling_list[-1])85        # in_features *= int(args.temporal_scale / 2) 86        # self.reduce = nn.Linear(in_features, args.z_dim)87        # self.mu = nn.Linear(in_features, args.z_dim)88        # self.logvar = nn.Linear(in_features, args.z_dim)89 90    def forward(self, input):91        #bs, n, c = input.shape[0], input.shape[1], input.shape[2]92        output = input.permute(0, 2, 1)#input.reshape(bs, n, -1, 6)93        for layer in self.layers:94            output = layer(output)95        #output = output.view(output.shape[0], -1)96        output = output.permute(0, 2, 1)97        return output98 99class ResBlock(nn.Module):100    def __init__(self, channel):101        super(ResBlock, self).__init__()102        self.model = nn.Sequential(103            nn.Conv1d(channel, channel, kernel_size=3, stride=1, padding=1),104            nn.LeakyReLU(0.2, inplace=True),105            nn.Conv1d(channel, channel, kernel_size=3, stride=1, padding=1),106        )107 108    def forward(self, x):109        residual = x110        out = self.model(x)111        out += residual112        return out113    114class VQDecoderV3(nn.Module):115    def __init__(self, args):116        super(VQDecoderV3, self).__init__()117        n_up = args.vae_layer118        channels = []119        for i in range(n_up-1):120            channels.append(args.vae_length)121        channels.append(args.vae_length)122        channels.append(args.vae_test_dim)123        input_size = args.vae_length124        n_resblk = 2125        assert len(channels) == n_up + 1126        if input_size == channels[0]:127            layers = []128        else:129            layers = [nn.Conv1d(input_size, channels[0], kernel_size=3, stride=1, padding=1)]130 131        for i in range(n_resblk):132            layers += [ResBlock(channels[0])]133        # channels = channels134        for i in range(n_up):135            layers += [136                nn.Upsample(scale_factor=2, mode='nearest'),137                nn.Conv1d(channels[i], channels[i+1], kernel_size=3, stride=1, padding=1),138                nn.LeakyReLU(0.2, inplace=True)139            ]140        layers += [nn.Conv1d(channels[-1], channels[-1], kernel_size=3, stride=1, padding=1)]141        self.main = nn.Sequential(*layers)142        # self.main.apply(init_weight)143 144    def forward(self, inputs):145        inputs = inputs.permute(0, 2, 1)146        outputs = self.main(inputs).permute(0, 2, 1)147        return outputs148    149def reparameterize(mu, logvar):150    std = torch.exp(0.5 * logvar)151    eps = torch.randn_like(std)152    return mu + eps * std153 154class VAEConv(nn.Module):155    def __init__(self, args):156        super(VAEConv, self).__init__()157        # self.encoder = VQEncoderV3(args)158        # self.decoder = VQDecoderV3(args)159        self.fc_mu = nn.Linear(args.vae_length, args.vae_length)160        self.fc_logvar = nn.Linear(args.vae_length, args.vae_length)161        self.variational = args.variational162        163    def forward(self, inputs):164        pre_latent = self.encoder(inputs)165        mu, logvar = None, None166        if self.variational:167            mu = self.fc_mu(pre_latent)168            logvar = self.fc_logvar(pre_latent)169            pre_latent = reparameterize(mu, logvar)170        rec_pose = self.decoder(pre_latent)171        return {172            "poses_feat":pre_latent,173            "rec_pose": rec_pose,174            "pose_mu": mu,175            "pose_logvar": logvar,176            }177    178    def map2latent(self, inputs):179        pre_latent = self.encoder(inputs)180        if self.variational:181            mu = self.fc_mu(pre_latent)182            logvar = self.fc_logvar(pre_latent)183            pre_latent = reparameterize(mu, logvar)184        return pre_latent185    186    def decode(self, pre_latent):187        rec_pose = self.decoder(pre_latent)188        return rec_pose189 190class VAESKConv(VAEConv):191    def __init__(self, args, model_save_path="./emage/"):192        # args = args()193        super(VAESKConv, self).__init__(args)194        smpl_fname = model_save_path +'smplx_models/smplx/SMPLX_NEUTRAL_2020.npz'195        smpl_data = np.load(smpl_fname, encoding='latin1')196        parents = smpl_data['kintree_table'][0].astype(np.int32)197        edges = build_edge_topology(parents)198        self.encoder = LocalEncoder(args, edges)199        self.decoder = VQDecoderV3(args)