mingyuan/MotionDiffuse
69
1import os2import torch3import argparse4 5import utils.paramUtil as paramUtil6from torch.utils.data import DataLoader7from utils.plot_script import *8 9from utils.utils import *10from utils.motion_process import recover_from_ric11 12 13def plot_t2m(opt, data, result_path, caption):14 joint = recover_from_ric(torch.from_numpy(data).float(), opt.joints_num).numpy()15 # joint = motion_temporal_filter(joint, sigma=1)16 plot_3d_motion(result_path, paramUtil.t2m_kinematic_chain, joint, title=caption, fps=20)17 18 19def process(trainer, opt, device, mean, std, text, motion_length, result_path):20 21 result_dict = {}22 with torch.no_grad():23 if motion_length != -1:24 caption = [text]25 m_lens = torch.LongTensor([motion_length]).to(device)26 pred_motions = trainer.generate(caption, m_lens, opt.dim_pose)27 motion = pred_motions[0].cpu().numpy()28 motion = motion * std + mean29 title = text + " #%d" % motion.shape[0]30 plot_t2m(opt, motion, result_path, title)31 