mingyuan/MotionDiffuse
69
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 