Team Ai
Apppublic

mingyuan/MotionDiffuse

sourceHugging Facemitupdated 3y agoView on Hugging Face
69likes
evaluator_models.py439 linesDownload Raw Back to datasets
1import torch2import torch.nn as nn3import numpy as np4import time5import math6from torch.nn.utils.rnn import pack_padded_sequence, pad_packed_sequence7# from networks.layers import *8import torch.nn.functional as F9 10 11class ContrastiveLoss(torch.nn.Module):12    """13    Contrastive loss function.14    Based on: http://yann.lecun.com/exdb/publis/pdf/hadsell-chopra-lecun-06.pdf15    """16    def __init__(self, margin=3.0):17        super(ContrastiveLoss, self).__init__()18        self.margin = margin19 20    def forward(self, output1, output2, label):21        euclidean_distance = F.pairwise_distance(output1, output2, keepdim=True)22        loss_contrastive = torch.mean((1-label) * torch.pow(euclidean_distance, 2) +23                                      (label) * torch.pow(torch.clamp(self.margin - euclidean_distance, min=0.0), 2))24        return loss_contrastive25 26 27def init_weight(m):28    if isinstance(m, nn.Conv1d) or isinstance(m, nn.Linear) or isinstance(m, nn.ConvTranspose1d):29        nn.init.xavier_normal_(m.weight)30        # m.bias.data.fill_(0.01)31        if m.bias is not None:32            nn.init.constant_(m.bias, 0)33 34 35def reparameterize(mu, logvar):36    s_var = logvar.mul(0.5).exp_()37    eps = s_var.data.new(s_var.size()).normal_()38    return eps.mul(s_var).add_(mu)39 40 41# batch_size, dimension and position42# output: (batch_size, dim)43def positional_encoding(batch_size, dim, pos):44    assert batch_size == pos.shape[0]45    positions_enc = np.array([46        [pos[j] / np.power(10000, (i-i%2)/dim) for i in range(dim)]47        for j in range(batch_size)48    ], dtype=np.float32)49    positions_enc[:, 0::2] = np.sin(positions_enc[:, 0::2])50    positions_enc[:, 1::2] = np.cos(positions_enc[:, 1::2])51    return torch.from_numpy(positions_enc).float()52 53 54def get_padding_mask(batch_size, seq_len, cap_lens):55    cap_lens = cap_lens.data.tolist()56    mask_2d = torch.ones((batch_size, seq_len, seq_len), dtype=torch.float32)57    for i, cap_len in enumerate(cap_lens):58        mask_2d[i, :, :cap_len] = 059    return mask_2d.bool(), 1 - mask_2d[:, :, 0].clone()60 61 62class PositionalEncoding(nn.Module):63 64    def __init__(self, d_model, max_len=300):65        super(PositionalEncoding, self).__init__()66 67        pe = torch.zeros(max_len, d_model)68        position = torch.arange(0, max_len, dtype=torch.float).unsqueeze(1)69        div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))70        pe[:, 0::2] = torch.sin(position * div_term)71        pe[:, 1::2] = torch.cos(position * div_term)72        # pe = pe.unsqueeze(0).transpose(0, 1)73        self.register_buffer('pe', pe)74 75    def forward(self, pos):76        return self.pe[pos]77 78 79class MovementConvEncoder(nn.Module):80    def __init__(self, input_size, hidden_size, output_size):81        super(MovementConvEncoder, self).__init__()82        self.main = nn.Sequential(83            nn.Conv1d(input_size, hidden_size, 4, 2, 1),84            nn.Dropout(0.2, inplace=True),85            nn.LeakyReLU(0.2, inplace=True),86            nn.Conv1d(hidden_size, output_size, 4, 2, 1),87            nn.Dropout(0.2, inplace=True),88            nn.LeakyReLU(0.2, inplace=True),89        )90        self.out_net = nn.Linear(output_size, output_size)91        self.main.apply(init_weight)92        self.out_net.apply(init_weight)93 94    def forward(self, inputs):95        inputs = inputs.permute(0, 2, 1)96        outputs = self.main(inputs).permute(0, 2, 1)97        # print(outputs.shape)98        return self.out_net(outputs)99 100 101class MovementConvDecoder(nn.Module):102    def __init__(self, input_size, hidden_size, output_size):103        super(MovementConvDecoder, self).__init__()104        self.main = nn.Sequential(105            nn.ConvTranspose1d(input_size, hidden_size, 4, 2, 1),106            # nn.Dropout(0.2, inplace=True),107            nn.LeakyReLU(0.2, inplace=True),108            nn.ConvTranspose1d(hidden_size, output_size, 4, 2, 1),109            # nn.Dropout(0.2, inplace=True),110            nn.LeakyReLU(0.2, inplace=True),111        )112        self.out_net = nn.Linear(output_size, output_size)113 114        self.main.apply(init_weight)115        self.out_net.apply(init_weight)116 117    def forward(self, inputs):118        inputs = inputs.permute(0, 2, 1)119        outputs = self.main(inputs).permute(0, 2, 1)120        return self.out_net(outputs)121 122 123class TextVAEDecoder(nn.Module):124    def __init__(self, text_size, input_size, output_size, hidden_size, n_layers):125        super(TextVAEDecoder, self).__init__()126        self.input_size = input_size127        self.output_size = output_size128        self.hidden_size = hidden_size129        self.n_layers = n_layers130        self.emb = nn.Sequential(131            nn.Linear(input_size, hidden_size),132            nn.LayerNorm(hidden_size),133            nn.LeakyReLU(0.2, inplace=True))134 135        self.z2init = nn.Linear(text_size, hidden_size * n_layers)136        self.gru = nn.ModuleList([nn.GRUCell(hidden_size, hidden_size) for i in range(self.n_layers)])137        self.positional_encoder = PositionalEncoding(hidden_size)138 139 140        self.output = nn.Sequential(141            nn.Linear(hidden_size, hidden_size),142            nn.LayerNorm(hidden_size),143            nn.LeakyReLU(0.2, inplace=True),144            nn.Linear(hidden_size, output_size)145        )146 147        #148        # self.output = nn.Sequential(149        #     nn.Linear(hidden_size, hidden_size),150        #     nn.LayerNorm(hidden_size),151        #     nn.LeakyReLU(0.2, inplace=True),152        #     nn.Linear(hidden_size, output_size-4)153        # )154 155        # self.contact_net = nn.Sequential(156        #     nn.Linear(output_size-4, 64),157        #     nn.LayerNorm(64),158        #     nn.LeakyReLU(0.2, inplace=True),159        #     nn.Linear(64, 4)160        # )161 162        self.output.apply(init_weight)163        self.emb.apply(init_weight)164        self.z2init.apply(init_weight)165        # self.contact_net.apply(init_weight)166 167    def get_init_hidden(self, latent):168        hidden = self.z2init(latent)169        hidden = torch.split(hidden, self.hidden_size, dim=-1)170        return list(hidden)171 172    def forward(self, inputs, last_pred, hidden, p):173        h_in = self.emb(inputs)174        pos_enc = self.positional_encoder(p).to(inputs.device).detach()175        h_in = h_in + pos_enc176        for i in range(self.n_layers):177            # print(h_in.shape)178            hidden[i] = self.gru[i](h_in, hidden[i])179            h_in = hidden[i]180        pose_pred = self.output(h_in)181        # pose_pred = self.output(h_in) + last_pred.detach()182        # contact = self.contact_net(pose_pred)183        # return torch.cat([pose_pred, contact], dim=-1), hidden184        return pose_pred, hidden185 186 187class TextDecoder(nn.Module):188    def __init__(self, text_size, input_size, output_size, hidden_size, n_layers):189        super(TextDecoder, self).__init__()190        self.input_size = input_size191        self.output_size = output_size192        self.hidden_size = hidden_size193        self.n_layers = n_layers194        self.emb = nn.Sequential(195            nn.Linear(input_size, hidden_size),196            nn.LayerNorm(hidden_size),197            nn.LeakyReLU(0.2, inplace=True))198 199        self.gru = nn.ModuleList([nn.GRUCell(hidden_size, hidden_size) for i in range(self.n_layers)])200        self.z2init = nn.Linear(text_size, hidden_size * n_layers)201        self.positional_encoder = PositionalEncoding(hidden_size)202 203        self.mu_net = nn.Linear(hidden_size, output_size)204        self.logvar_net = nn.Linear(hidden_size, output_size)205 206        self.emb.apply(init_weight)207        self.z2init.apply(init_weight)208        self.mu_net.apply(init_weight)209        self.logvar_net.apply(init_weight)210 211    def get_init_hidden(self, latent):212 213        hidden = self.z2init(latent)214        hidden = torch.split(hidden, self.hidden_size, dim=-1)215 216        return list(hidden)217 218    def forward(self, inputs, hidden, p):219        # print(inputs.shape)220        x_in = self.emb(inputs)221        pos_enc = self.positional_encoder(p).to(inputs.device).detach()222        x_in = x_in + pos_enc223 224        for i in range(self.n_layers):225            hidden[i] = self.gru[i](x_in, hidden[i])226            h_in = hidden[i]227        mu = self.mu_net(h_in)228        logvar = self.logvar_net(h_in)229        z = reparameterize(mu, logvar)230        return z, mu, logvar, hidden231 232class AttLayer(nn.Module):233    def __init__(self, query_dim, key_dim, value_dim):234        super(AttLayer, self).__init__()235        self.W_q = nn.Linear(query_dim, value_dim)236        self.W_k = nn.Linear(key_dim, value_dim, bias=False)237        self.W_v = nn.Linear(key_dim, value_dim)238 239        self.softmax = nn.Softmax(dim=1)240        self.dim = value_dim241 242        self.W_q.apply(init_weight)243        self.W_k.apply(init_weight)244        self.W_v.apply(init_weight)245 246    def forward(self, query, key_mat):247        '''248        query (batch, query_dim)249        key (batch, seq_len, key_dim)250        '''251        # print(query.shape)252        query_vec = self.W_q(query).unsqueeze(-1)       # (batch, value_dim, 1)253        val_set = self.W_v(key_mat)                     # (batch, seq_len, value_dim)254        key_set = self.W_k(key_mat)                     # (batch, seq_len, value_dim)255 256        weights = torch.matmul(key_set, query_vec) / np.sqrt(self.dim)257 258        co_weights = self.softmax(weights)              # (batch, seq_len, 1)259        values = val_set * co_weights                   # (batch, seq_len, value_dim)260        pred = values.sum(dim=1)                        # (batch, value_dim)261        return pred, co_weights262 263    def short_cut(self, querys, keys):264        return self.W_q(querys), self.W_k(keys)265 266 267class TextEncoderBiGRU(nn.Module):268    def __init__(self, word_size, pos_size, hidden_size, device):269        super(TextEncoderBiGRU, self).__init__()270        self.device = device271 272        self.pos_emb = nn.Linear(pos_size, word_size)273        self.input_emb = nn.Linear(word_size, hidden_size)274        self.gru = nn.GRU(hidden_size, hidden_size, batch_first=True, bidirectional=True)275        # self.linear2 = nn.Linear(hidden_size, output_size)276 277        self.input_emb.apply(init_weight)278        self.pos_emb.apply(init_weight)279        # self.linear2.apply(init_weight)280        # self.batch_size = batch_size281        self.hidden_size = hidden_size282        self.hidden = nn.Parameter(torch.randn((2, 1, self.hidden_size), requires_grad=True))283 284    # input(batch_size, seq_len, dim)285    def forward(self, word_embs, pos_onehot, cap_lens):286        num_samples = word_embs.shape[0]287 288        pos_embs = self.pos_emb(pos_onehot)289        inputs = word_embs + pos_embs290        input_embs = self.input_emb(inputs)291        hidden = self.hidden.repeat(1, num_samples, 1)292 293        cap_lens = cap_lens.data.tolist()294        emb = pack_padded_sequence(input_embs, cap_lens, batch_first=True)295 296        gru_seq, gru_last = self.gru(emb, hidden)297 298        gru_last = torch.cat([gru_last[0], gru_last[1]], dim=-1)299        gru_seq = pad_packed_sequence(gru_seq, batch_first=True)[0]300        forward_seq = gru_seq[..., :self.hidden_size]301        backward_seq = gru_seq[..., self.hidden_size:].clone()302 303        # Concate the forward and backward word embeddings304        for i, length in enumerate(cap_lens):305            backward_seq[i:i+1, :length] = torch.flip(backward_seq[i:i+1, :length].clone(), dims=[1])306        gru_seq = torch.cat([forward_seq, backward_seq], dim=-1)307 308        return gru_seq, gru_last309 310 311class TextEncoderBiGRUCo(nn.Module):312    def __init__(self, word_size, pos_size, hidden_size, output_size, device):313        super(TextEncoderBiGRUCo, self).__init__()314        self.device = device315 316        self.pos_emb = nn.Linear(pos_size, word_size)317        self.input_emb = nn.Linear(word_size, hidden_size)318        self.gru = nn.GRU(hidden_size, hidden_size, batch_first=True, bidirectional=True)319        self.output_net = nn.Sequential(320            nn.Linear(hidden_size * 2, hidden_size),321            nn.LayerNorm(hidden_size),322            nn.LeakyReLU(0.2, inplace=True),323            nn.Linear(hidden_size, output_size)324        )325 326        self.input_emb.apply(init_weight)327        self.pos_emb.apply(init_weight)328        self.output_net.apply(init_weight)329        # self.linear2.apply(init_weight)330        # self.batch_size = batch_size331        self.hidden_size = hidden_size332        self.hidden = nn.Parameter(torch.randn((2, 1, self.hidden_size), requires_grad=True))333 334    # input(batch_size, seq_len, dim)335    def forward(self, word_embs, pos_onehot, cap_lens):336        num_samples = word_embs.shape[0]337 338        pos_embs = self.pos_emb(pos_onehot)339        inputs = word_embs + pos_embs340        input_embs = self.input_emb(inputs)341        hidden = self.hidden.repeat(1, num_samples, 1)342 343        cap_lens = cap_lens.data.tolist()344        emb = pack_padded_sequence(input_embs, cap_lens, batch_first=True)345 346        gru_seq, gru_last = self.gru(emb, hidden)347 348        gru_last = torch.cat([gru_last[0], gru_last[1]], dim=-1)349 350        return self.output_net(gru_last)351 352 353class MotionEncoderBiGRUCo(nn.Module):354    def __init__(self, input_size, hidden_size, output_size, device):355        super(MotionEncoderBiGRUCo, self).__init__()356        self.device = device357 358        self.input_emb = nn.Linear(input_size, hidden_size)359        self.gru = nn.GRU(hidden_size, hidden_size, batch_first=True, bidirectional=True)360        self.output_net = nn.Sequential(361            nn.Linear(hidden_size*2, hidden_size),362            nn.LayerNorm(hidden_size),363            nn.LeakyReLU(0.2, inplace=True),364            nn.Linear(hidden_size, output_size)365        )366 367        self.input_emb.apply(init_weight)368        self.output_net.apply(init_weight)369        self.hidden_size = hidden_size370        self.hidden = nn.Parameter(torch.randn((2, 1, self.hidden_size), requires_grad=True))371 372    # input(batch_size, seq_len, dim)373    def forward(self, inputs, m_lens):374        num_samples = inputs.shape[0]375 376        input_embs = self.input_emb(inputs)377        hidden = self.hidden.repeat(1, num_samples, 1)378 379        cap_lens = m_lens.data.tolist()380        emb = pack_padded_sequence(input_embs, cap_lens, batch_first=True)381 382        gru_seq, gru_last = self.gru(emb, hidden)383 384        gru_last = torch.cat([gru_last[0], gru_last[1]], dim=-1)385 386        return self.output_net(gru_last)387 388 389class MotionLenEstimatorBiGRU(nn.Module):390    def __init__(self, word_size, pos_size, hidden_size, output_size):391        super(MotionLenEstimatorBiGRU, self).__init__()392 393        self.pos_emb = nn.Linear(pos_size, word_size)394        self.input_emb = nn.Linear(word_size, hidden_size)395        self.gru = nn.GRU(hidden_size, hidden_size, batch_first=True, bidirectional=True)396        nd = 512397        self.output = nn.Sequential(398            nn.Linear(hidden_size*2, nd),399            nn.LayerNorm(nd),400            nn.LeakyReLU(0.2, inplace=True),401 402            nn.Linear(nd, nd // 2),403            nn.LayerNorm(nd // 2),404            nn.LeakyReLU(0.2, inplace=True),405 406            nn.Linear(nd // 2, nd // 4),407            nn.LayerNorm(nd // 4),408            nn.LeakyReLU(0.2, inplace=True),409 410            nn.Linear(nd // 4, output_size)411        )412        # self.linear2 = nn.Linear(hidden_size, output_size)413 414        self.input_emb.apply(init_weight)415        self.pos_emb.apply(init_weight)416        self.output.apply(init_weight)417        # self.linear2.apply(init_weight)418        # self.batch_size = batch_size419        self.hidden_size = hidden_size420        self.hidden = nn.Parameter(torch.randn((2, 1, self.hidden_size), requires_grad=True))421 422    # input(batch_size, seq_len, dim)423    def forward(self, word_embs, pos_onehot, cap_lens):424        num_samples = word_embs.shape[0]425 426        pos_embs = self.pos_emb(pos_onehot)427        inputs = word_embs + pos_embs428        input_embs = self.input_emb(inputs)429        hidden = self.hidden.repeat(1, num_samples, 1)430 431        cap_lens = cap_lens.data.tolist()432        emb = pack_padded_sequence(input_embs, cap_lens, batch_first=True)433 434        gru_seq, gru_last = self.gru(emb, hidden)435 436        gru_last = torch.cat([gru_last[0], gru_last[1]], dim=-1)437 438        return self.output(gru_last)439