Team Ai
Apppublic

mingyuan/MotionDiffuse

sourceHugging Facemitupdated 3y agoView on Hugging Face
69likes
evaluator.py442 linesDownload Raw Back to datasets
1import torch2from utils.word_vectorizer import WordVectorizer, POS_enumerator3from utils.get_opt import get_opt4from models import MotionTransformer5from torch.utils.data import Dataset, DataLoader6from os.path import join as pjoin7from tqdm import tqdm8import numpy as np9from .evaluator_models import *10import os11import codecs as cs12import random13from torch.utils.data._utils.collate import default_collate14 15 16class EvaluationDataset(Dataset):17 18    def __init__(self, opt, trainer, dataset, w_vectorizer, mm_num_samples, mm_num_repeats):19        assert mm_num_samples < len(dataset)20        print(opt.model_dir)21 22        dataloader = DataLoader(dataset, batch_size=1, num_workers=1, shuffle=True)23        epoch, it = trainer.load(pjoin(opt.model_dir, opt.which_epoch + '.tar'))24 25        generated_motion = []26        min_mov_length = 10 if opt.dataset_name == 't2m' else 627 28        trainer.eval_mode()29        trainer.to(opt.device)30 31        # Pre-process all target captions32        mm_generated_motions = []33        mm_idxs = np.random.choice(len(dataset), mm_num_samples, replace=False)34        mm_idxs = np.sort(mm_idxs)35        all_caption = []36        all_m_lens = []37        all_data = []38        with torch.no_grad():39            for i, data in tqdm(enumerate(dataloader)):40                word_emb, pos_ohot, caption, cap_lens, motions, m_lens, tokens = data41                all_data.append(data)42                tokens = tokens[0].split('_')43                mm_num_now = len(mm_generated_motions)44                is_mm = True if ((mm_num_now < mm_num_samples) and (i == mm_idxs[mm_num_now])) else False45                repeat_times = mm_num_repeats if is_mm else 146                m_lens = max(m_lens // opt.unit_length * opt.unit_length, min_mov_length * opt.unit_length)47                m_lens = min(m_lens, opt.max_motion_length)48                if isinstance(m_lens, int):49                    m_lens = torch.LongTensor([m_lens]).to(opt.device)50                else:51                    m_lens = m_lens.to(opt.device)52                for t in range(repeat_times):53                    all_m_lens.append(m_lens)54                    all_caption.extend(caption)55                if is_mm:56                    mm_generated_motions.append(0)57        all_m_lens = torch.stack(all_m_lens)58        59        # Generate all sequences60        with torch.no_grad():61            all_pred_motions = trainer.generate(all_caption, all_m_lens, opt.dim_pose)62        63        cur_idx = 064        mm_generated_motions = []65        with torch.no_grad():66            for i, data_dummy in tqdm(enumerate(dataloader)):67                data = all_data[i]68                word_emb, pos_ohot, caption, cap_lens, motions, m_lens, tokens = data69                tokens = tokens[0].split('_')70                mm_num_now = len(mm_generated_motions)71                is_mm = True if ((mm_num_now < mm_num_samples) and (i == mm_idxs[mm_num_now])) else False72                repeat_times = mm_num_repeats if is_mm else 173                mm_motions = []74                m_lens = max(m_lens // opt.unit_length * opt.unit_length, min_mov_length * opt.unit_length)75                m_lens = min(m_lens, opt.max_motion_length)76                if isinstance(m_lens, int):77                    m_lens = torch.LongTensor([m_lens]).to(opt.device)78                else:79                    m_lens = m_lens.to(opt.device)80                for t in range(repeat_times):81                    m_len = m_lens[0].item()82                    pred_motions = all_pred_motions[cur_idx][:m_lens[0].item()]83                    assert pred_motions.shape[0] == m_lens[0].item()84                    cur_idx += 185                    if t == 0:86                        sub_dict = {'motion': pred_motions.cpu().numpy(),87                                    'length': pred_motions.shape[0],88                                    'caption': caption[0],89                                    'cap_len': cap_lens[0].item(),90                                    'tokens': tokens}91                        generated_motion.append(sub_dict)92 93                    if is_mm:94                        mm_motions.append({95                            'motion': pred_motions.cpu().numpy(),96                            'length': m_lens[0].item()97                        })98                if is_mm:99                    mm_generated_motions.append({'caption': caption[0],100                                                 'tokens': tokens,101                                                 'cap_len': cap_lens[0].item(),102                                                 'mm_motions': mm_motions})103        self.generated_motion = generated_motion104        self.mm_generated_motion = mm_generated_motions105        self.opt = opt106        self.w_vectorizer = w_vectorizer107 108 109    def __len__(self):110        return len(self.generated_motion)111 112 113    def __getitem__(self, item):114        data = self.generated_motion[item]115        motion, m_length, caption, tokens = data['motion'], data['length'], data['caption'], data['tokens']116        sent_len = data['cap_len']117        pos_one_hots = []118        word_embeddings = []119        for token in tokens:120            word_emb, pos_oh = self.w_vectorizer[token]121            pos_one_hots.append(pos_oh[None, :])122            word_embeddings.append(word_emb[None, :])123        pos_one_hots = np.concatenate(pos_one_hots, axis=0)124        word_embeddings = np.concatenate(word_embeddings, axis=0)125 126        if m_length < self.opt.max_motion_length:127            motion = np.concatenate([motion,128                                     np.zeros((self.opt.max_motion_length - m_length, motion.shape[1]))129                                     ], axis=0)130        return word_embeddings, pos_one_hots, caption, sent_len, motion, m_length, '_'.join(tokens)131 132 133def collate_fn(batch):134    batch.sort(key=lambda x: x[3], reverse=True)135    return default_collate(batch)136 137 138'''For use of training text motion matching model, and evaluations'''139class Text2MotionDatasetV2(Dataset):140    def __init__(self, opt, mean, std, split_file, w_vectorizer):141        self.opt = opt142        self.w_vectorizer = w_vectorizer143        self.max_length = 20144        self.pointer = 0145        self.max_motion_length = opt.max_motion_length146        min_motion_len = 40 if self.opt.dataset_name =='t2m' else 24147 148        data_dict = {}149        id_list = []150        with cs.open(split_file, 'r') as f:151            for line in f.readlines():152                id_list.append(line.strip())153 154        new_name_list = []155        length_list = []156        for name in tqdm(id_list):157            try:158                motion = np.load(pjoin(opt.motion_dir, name + '.npy'))159                if (len(motion)) < min_motion_len or (len(motion) >= 200):160                    continue161                text_data = []162                flag = False163                with cs.open(pjoin(opt.text_dir, name + '.txt')) as f:164                    for line in f.readlines():165                        text_dict = {}166                        line_split = line.strip().split('#')167                        caption = line_split[0]168                        tokens = line_split[1].split(' ')169                        f_tag = float(line_split[2])170                        to_tag = float(line_split[3])171                        f_tag = 0.0 if np.isnan(f_tag) else f_tag172                        to_tag = 0.0 if np.isnan(to_tag) else to_tag173 174                        text_dict['caption'] = caption175                        text_dict['tokens'] = tokens176                        if f_tag == 0.0 and to_tag == 0.0:177                            flag = True178                            text_data.append(text_dict)179                        else:180                            try:181                                n_motion = motion[int(f_tag*20) : int(to_tag*20)]182                                if (len(n_motion)) < min_motion_len or (len(n_motion) >= 200):183                                    continue184                                new_name = random.choice('ABCDEFGHIJKLMNOPQRSTUVW') + '_' + name185                                while new_name in data_dict:186                                    new_name = random.choice('ABCDEFGHIJKLMNOPQRSTUVW') + '_' + name187                                data_dict[new_name] = {'motion': n_motion,188                                                       'length': len(n_motion),189                                                       'text':[text_dict]}190                                new_name_list.append(new_name)191                                length_list.append(len(n_motion))192                            except:193                                print(line_split)194                                print(line_split[2], line_split[3], f_tag, to_tag, name)195                                # break196 197                if flag:198                    data_dict[name] = {'motion': motion,199                                       'length': len(motion),200                                       'text': text_data}201                    new_name_list.append(name)202                    length_list.append(len(motion))203            except:204                pass205 206        name_list, length_list = zip(*sorted(zip(new_name_list, length_list), key=lambda x: x[1]))207 208        self.mean = mean209        self.std = std210        self.length_arr = np.array(length_list)211        self.data_dict = data_dict212        self.name_list = name_list213        self.reset_max_len(self.max_length)214 215    def reset_max_len(self, length):216        assert length <= self.max_motion_length217        self.pointer = np.searchsorted(self.length_arr, length)218        print("Pointer Pointing at %d"%self.pointer)219        self.max_length = length220 221    def inv_transform(self, data):222        return data * self.std + self.mean223 224    def __len__(self):225        return len(self.data_dict) - self.pointer226 227    def __getitem__(self, item):228        idx = self.pointer + item229        data = self.data_dict[self.name_list[idx]]230        motion, m_length, text_list = data['motion'], data['length'], data['text']231        # Randomly select a caption232        text_data = random.choice(text_list)233        caption, tokens = text_data['caption'], text_data['tokens']234 235        if len(tokens) < self.opt.max_text_len:236            # pad with "unk"237            tokens = ['sos/OTHER'] + tokens + ['eos/OTHER']238            sent_len = len(tokens)239            tokens = tokens + ['unk/OTHER'] * (self.opt.max_text_len + 2 - sent_len)240        else:241            # crop242            tokens = tokens[:self.opt.max_text_len]243            tokens = ['sos/OTHER'] + tokens + ['eos/OTHER']244            sent_len = len(tokens)245        pos_one_hots = []246        word_embeddings = []247        for token in tokens:248            word_emb, pos_oh = self.w_vectorizer[token]249            pos_one_hots.append(pos_oh[None, :])250            word_embeddings.append(word_emb[None, :])251        pos_one_hots = np.concatenate(pos_one_hots, axis=0)252        word_embeddings = np.concatenate(word_embeddings, axis=0)253 254        # Crop the motions in to times of 4, and introduce small variations255        if self.opt.unit_length < 10:256            coin2 = np.random.choice(['single', 'single', 'double'])257        else:258            coin2 = 'single'259 260        if coin2 == 'double':261            m_length = (m_length // self.opt.unit_length - 1) * self.opt.unit_length262        elif coin2 == 'single':263            m_length = (m_length // self.opt.unit_length) * self.opt.unit_length264        idx = random.randint(0, len(motion) - m_length)265        motion = motion[idx:idx+m_length]266 267        "Z Normalization"268        motion = (motion - self.mean) / self.std269 270        if m_length < self.max_motion_length:271            motion = np.concatenate([motion,272                                     np.zeros((self.max_motion_length - m_length, motion.shape[1]))273                                     ], axis=0)274        return word_embeddings, pos_one_hots, caption, sent_len, motion, m_length, '_'.join(tokens)275 276 277def get_dataset_motion_loader(opt_path, batch_size, device):278    opt = get_opt(opt_path, device)279 280    # Configurations of T2M dataset and KIT dataset is almost the same281    if opt.dataset_name == 't2m' or opt.dataset_name == 'kit':282        print('Loading dataset %s ...' % opt.dataset_name)283 284        mean = np.load(pjoin(opt.meta_dir, 'mean.npy'))285        std = np.load(pjoin(opt.meta_dir, 'std.npy'))286 287        w_vectorizer = WordVectorizer('./data/glove', 'our_vab')288        split_file = pjoin(opt.data_root, 'test.txt')289        dataset = Text2MotionDatasetV2(opt, mean, std, split_file, w_vectorizer)290        dataloader = DataLoader(dataset, batch_size=batch_size, num_workers=4, drop_last=True,291                                collate_fn=collate_fn, shuffle=True)292    else:293        raise KeyError('Dataset not Recognized !!')294 295    print('Ground Truth Dataset Loading Completed!!!')296    return dataloader, dataset297 298 299class MMGeneratedDataset(Dataset):300    def __init__(self, opt, motion_dataset, w_vectorizer):301        self.opt = opt302        self.dataset = motion_dataset.mm_generated_motion303        self.w_vectorizer = w_vectorizer304 305    def __len__(self):306        return len(self.dataset)307 308    def __getitem__(self, item):309        data = self.dataset[item]310        mm_motions = data['mm_motions']311        m_lens = []312        motions = []313        for mm_motion in mm_motions:314            m_lens.append(mm_motion['length'])315            motion = mm_motion['motion']316            if len(motion) < self.opt.max_motion_length:317                motion = np.concatenate([motion,318                                         np.zeros((self.opt.max_motion_length - len(motion), motion.shape[1]))319                                         ], axis=0)320            motion = motion[None, :]321            motions.append(motion)322        m_lens = np.array(m_lens, dtype=np.int)323        motions = np.concatenate(motions, axis=0)324        sort_indx = np.argsort(m_lens)[::-1].copy()325        # print(m_lens)326        # print(sort_indx)327        # print(m_lens[sort_indx])328        m_lens = m_lens[sort_indx]329        motions = motions[sort_indx]330        return motions, m_lens331 332 333 334def get_motion_loader(opt, batch_size, trainer, ground_truth_dataset, mm_num_samples, mm_num_repeats):335 336    # Currently the configurations of two datasets are almost the same337    if opt.dataset_name == 't2m' or opt.dataset_name == 'kit':338        w_vectorizer = WordVectorizer('./data/glove', 'our_vab')339    else:340        raise KeyError('Dataset not recognized!!')341    print('Generating %s ...' % opt.name)342 343    dataset = EvaluationDataset(opt, trainer, ground_truth_dataset, w_vectorizer, mm_num_samples, mm_num_repeats)344    mm_dataset = MMGeneratedDataset(opt, dataset, w_vectorizer)345 346    motion_loader = DataLoader(dataset, batch_size=batch_size, collate_fn=collate_fn, drop_last=True, num_workers=4)347    mm_motion_loader = DataLoader(mm_dataset, batch_size=1, num_workers=1)348 349    print('Generated Dataset Loading Completed!!!')350 351    return motion_loader, mm_motion_loader352 353 354def build_models(opt):355    movement_enc = MovementConvEncoder(opt.dim_pose-4, opt.dim_movement_enc_hidden, opt.dim_movement_latent)356    text_enc = TextEncoderBiGRUCo(word_size=opt.dim_word,357                                  pos_size=opt.dim_pos_ohot,358                                  hidden_size=opt.dim_text_hidden,359                                  output_size=opt.dim_coemb_hidden,360                                  device=opt.device)361 362    motion_enc = MotionEncoderBiGRUCo(input_size=opt.dim_movement_latent,363                                      hidden_size=opt.dim_motion_hidden,364                                      output_size=opt.dim_coemb_hidden,365                                      device=opt.device)366 367    checkpoint = torch.load(pjoin('data/pretrained_models', opt.dataset_name, 'text_mot_match', 'model', 'finest.tar'),368                            map_location=opt.device)369    movement_enc.load_state_dict(checkpoint['movement_encoder'])370    text_enc.load_state_dict(checkpoint['text_encoder'])371    motion_enc.load_state_dict(checkpoint['motion_encoder'])372    print('Loading Evaluation Model Wrapper (Epoch %d) Completed!!' % (checkpoint['epoch']))373    return text_enc, motion_enc, movement_enc374 375 376class EvaluatorModelWrapper(object):377 378    def __init__(self, opt):379 380        if opt.dataset_name == 't2m':381            opt.dim_pose = 263382        elif opt.dataset_name == 'kit':383            opt.dim_pose = 251384        else:385            raise KeyError('Dataset not Recognized!!!')386 387        opt.dim_word = 300388        opt.max_motion_length = 196389        opt.dim_pos_ohot = len(POS_enumerator)390        opt.dim_motion_hidden = 1024391        opt.max_text_len = 20392        opt.dim_text_hidden = 512393        opt.dim_coemb_hidden = 512394 395        self.text_encoder, self.motion_encoder, self.movement_encoder = build_models(opt)396        self.opt = opt397        self.device = opt.device398 399        self.text_encoder.to(opt.device)400        self.motion_encoder.to(opt.device)401        self.movement_encoder.to(opt.device)402 403        self.text_encoder.eval()404        self.motion_encoder.eval()405        self.movement_encoder.eval()406 407    # Please note that the results does not following the order of inputs408    def get_co_embeddings(self, word_embs, pos_ohot, cap_lens, motions, m_lens):409        with torch.no_grad():410            word_embs = word_embs.detach().to(self.device).float()411            pos_ohot = pos_ohot.detach().to(self.device).float()412            motions = motions.detach().to(self.device).float()413 414            align_idx = np.argsort(m_lens.data.tolist())[::-1].copy()415            motions = motions[align_idx]416            m_lens = m_lens[align_idx]417 418            '''Movement Encoding'''419            movements = self.movement_encoder(motions[..., :-4]).detach()420            m_lens = m_lens // self.opt.unit_length421            motion_embedding = self.motion_encoder(movements, m_lens)422 423            '''Text Encoding'''424            text_embedding = self.text_encoder(word_embs, pos_ohot, cap_lens)425            text_embedding = text_embedding[align_idx]426        return text_embedding, motion_embedding427 428    # Please note that the results does not following the order of inputs429    def get_motion_embeddings(self, motions, m_lens):430        with torch.no_grad():431            motions = motions.detach().to(self.device).float()432 433            align_idx = np.argsort(m_lens.data.tolist())[::-1].copy()434            motions = motions[align_idx]435            m_lens = m_lens[align_idx]436 437            '''Movement Encoding'''438            movements = self.movement_encoder(motions[..., :-4]).detach()439            m_lens = m_lens // self.opt.unit_length440            motion_embedding = self.motion_encoder(movements, m_lens)441        return motion_embedding442