Team Ai
Apppublic

mingyuan/MotionDiffuse

sourceHugging Facemitupdated 3y agoView on Hugging Face
69likes
motion_process.py515 linesDownload Raw Back to utils
1from os.path import join as pjoin2 3import numpy as np4import os5from utils.quaternion import *6from utils.skeleton import Skeleton7from utils.paramUtil import *8 9import torch10from tqdm import tqdm11 12# positions (batch, joint_num, 3)13def uniform_skeleton(positions, target_offset):14    src_skel = Skeleton(n_raw_offsets, kinematic_chain, 'cpu')15    src_offset = src_skel.get_offsets_joints(torch.from_numpy(positions[0]))16    src_offset = src_offset.numpy()17    tgt_offset = target_offset.numpy()18    # print(src_offset)19    # print(tgt_offset)20    '''Calculate Scale Ratio as the ratio of legs'''21    src_leg_len = np.abs(src_offset[l_idx1]).max() + np.abs(src_offset[l_idx2]).max()22    tgt_leg_len = np.abs(tgt_offset[l_idx1]).max() + np.abs(tgt_offset[l_idx2]).max()23 24    scale_rt = tgt_leg_len / src_leg_len25    # print(scale_rt)26    src_root_pos = positions[:, 0]27    tgt_root_pos = src_root_pos * scale_rt28 29    '''Inverse Kinematics'''30    quat_params = src_skel.inverse_kinematics_np(positions, face_joint_indx)31    # print(quat_params.shape)32 33    '''Forward Kinematics'''34    src_skel.set_offset(target_offset)35    new_joints = src_skel.forward_kinematics_np(quat_params, tgt_root_pos)36    return new_joints37 38 39def extract_features(positions, feet_thre, n_raw_offsets, kinematic_chain, face_joint_indx, fid_r, fid_l):40    global_positions = positions.copy()41    """ Get Foot Contacts """42 43    def foot_detect(positions, thres):44        velfactor, heightfactor = np.array([thres, thres]), np.array([3.0, 2.0])45 46        feet_l_x = (positions[1:, fid_l, 0] - positions[:-1, fid_l, 0]) ** 247        feet_l_y = (positions[1:, fid_l, 1] - positions[:-1, fid_l, 1]) ** 248        feet_l_z = (positions[1:, fid_l, 2] - positions[:-1, fid_l, 2]) ** 249        #     feet_l_h = positions[:-1,fid_l,1]50        #     feet_l = (((feet_l_x + feet_l_y + feet_l_z) < velfactor) & (feet_l_h < heightfactor)).astype(np.float)51        feet_l = ((feet_l_x + feet_l_y + feet_l_z) < velfactor).astype(np.float)52 53        feet_r_x = (positions[1:, fid_r, 0] - positions[:-1, fid_r, 0]) ** 254        feet_r_y = (positions[1:, fid_r, 1] - positions[:-1, fid_r, 1]) ** 255        feet_r_z = (positions[1:, fid_r, 2] - positions[:-1, fid_r, 2]) ** 256        #     feet_r_h = positions[:-1,fid_r,1]57        #     feet_r = (((feet_r_x + feet_r_y + feet_r_z) < velfactor) & (feet_r_h < heightfactor)).astype(np.float)58        feet_r = (((feet_r_x + feet_r_y + feet_r_z) < velfactor)).astype(np.float)59        return feet_l, feet_r60 61    #62    feet_l, feet_r = foot_detect(positions, feet_thre)63    # feet_l, feet_r = foot_detect(positions, 0.002)64 65    '''Quaternion and Cartesian representation'''66    r_rot = None67 68    def get_rifke(positions):69        '''Local pose'''70        positions[..., 0] -= positions[:, 0:1, 0]71        positions[..., 2] -= positions[:, 0:1, 2]72        '''All pose face Z+'''73        positions = qrot_np(np.repeat(r_rot[:, None], positions.shape[1], axis=1), positions)74        return positions75 76    def get_quaternion(positions):77        skel = Skeleton(n_raw_offsets, kinematic_chain, "cpu")78        # (seq_len, joints_num, 4)79        quat_params = skel.inverse_kinematics_np(positions, face_joint_indx, smooth_forward=False)80 81        '''Fix Quaternion Discontinuity'''82        quat_params = qfix(quat_params)83        # (seq_len, 4)84        r_rot = quat_params[:, 0].copy()85        #     print(r_rot[0])86        '''Root Linear Velocity'''87        # (seq_len - 1, 3)88        velocity = (positions[1:, 0] - positions[:-1, 0]).copy()89        #     print(r_rot.shape, velocity.shape)90        velocity = qrot_np(r_rot[1:], velocity)91        '''Root Angular Velocity'''92        # (seq_len - 1, 4)93        r_velocity = qmul_np(r_rot[1:], qinv_np(r_rot[:-1]))94        quat_params[1:, 0] = r_velocity95        # (seq_len, joints_num, 4)96        return quat_params, r_velocity, velocity, r_rot97 98    def get_cont6d_params(positions):99        skel = Skeleton(n_raw_offsets, kinematic_chain, "cpu")100        # (seq_len, joints_num, 4)101        quat_params = skel.inverse_kinematics_np(positions, face_joint_indx, smooth_forward=True)102 103        '''Quaternion to continuous 6D'''104        cont_6d_params = quaternion_to_cont6d_np(quat_params)105        # (seq_len, 4)106        r_rot = quat_params[:, 0].copy()107        #     print(r_rot[0])108        '''Root Linear Velocity'''109        # (seq_len - 1, 3)110        velocity = (positions[1:, 0] - positions[:-1, 0]).copy()111        #     print(r_rot.shape, velocity.shape)112        velocity = qrot_np(r_rot[1:], velocity)113        '''Root Angular Velocity'''114        # (seq_len - 1, 4)115        r_velocity = qmul_np(r_rot[1:], qinv_np(r_rot[:-1]))116        # (seq_len, joints_num, 4)117        return cont_6d_params, r_velocity, velocity, r_rot118 119    cont_6d_params, r_velocity, velocity, r_rot = get_cont6d_params(positions)120    positions = get_rifke(positions)121 122    #     trejec = np.cumsum(np.concatenate([np.array([[0, 0, 0]]), velocity], axis=0), axis=0)123    #     r_rotations, r_pos = recover_ric_glo_np(r_velocity, velocity[:, [0, 2]])124 125    # plt.plot(positions_b[:, 0, 0], positions_b[:, 0, 2], marker='*')126    # plt.plot(ground_positions[:, 0, 0], ground_positions[:, 0, 2], marker='o', color='r')127    # plt.plot(trejec[:, 0], trejec[:, 2], marker='^', color='g')128    # plt.plot(r_pos[:, 0], r_pos[:, 2], marker='s', color='y')129    # plt.xlabel('x')130    # plt.ylabel('z')131    # plt.axis('equal')132    # plt.show()133 134    '''Root height'''135    root_y = positions[:, 0, 1:2]136 137    '''Root rotation and linear velocity'''138    # (seq_len-1, 1) rotation velocity along y-axis139    # (seq_len-1, 2) linear velovity on xz plane140    r_velocity = np.arcsin(r_velocity[:, 2:3])141    l_velocity = velocity[:, [0, 2]]142    #     print(r_velocity.shape, l_velocity.shape, root_y.shape)143    root_data = np.concatenate([r_velocity, l_velocity, root_y[:-1]], axis=-1)144 145    '''Get Joint Rotation Representation'''146    # (seq_len, (joints_num-1) *6) quaternion for skeleton joints147    rot_data = cont_6d_params[:, 1:].reshape(len(cont_6d_params), -1)148 149    '''Get Joint Rotation Invariant Position Represention'''150    # (seq_len, (joints_num-1)*3) local joint position151    ric_data = positions[:, 1:].reshape(len(positions), -1)152 153    '''Get Joint Velocity Representation'''154    # (seq_len-1, joints_num*3)155    local_vel = qrot_np(np.repeat(r_rot[:-1, None], global_positions.shape[1], axis=1),156                        global_positions[1:] - global_positions[:-1])157    local_vel = local_vel.reshape(len(local_vel), -1)158 159    data = root_data160    data = np.concatenate([data, ric_data[:-1]], axis=-1)161    data = np.concatenate([data, rot_data[:-1]], axis=-1)162    #     print(data.shape, local_vel.shape)163    data = np.concatenate([data, local_vel], axis=-1)164    data = np.concatenate([data, feet_l, feet_r], axis=-1)165 166    return data167 168 169def process_file(positions, feet_thre):170    # (seq_len, joints_num, 3)171    #     '''Down Sample'''172    #     positions = positions[::ds_num]173 174    '''Uniform Skeleton'''175    positions = uniform_skeleton(positions, tgt_offsets)176 177    '''Put on Floor'''178    floor_height = positions.min(axis=0).min(axis=0)[1]179    positions[:, :, 1] -= floor_height180    #     print(floor_height)181 182    #     plot_3d_motion("./positions_1.mp4", kinematic_chain, positions, 'title', fps=20)183 184    '''XZ at origin'''185    root_pos_init = positions[0]186    root_pose_init_xz = root_pos_init[0] * np.array([1, 0, 1])187    positions = positions - root_pose_init_xz188 189    # '''Move the first pose to origin '''190    # root_pos_init = positions[0]191    # positions = positions - root_pos_init[0]192 193    '''All initially face Z+'''194    r_hip, l_hip, sdr_r, sdr_l = face_joint_indx195    across1 = root_pos_init[r_hip] - root_pos_init[l_hip]196    across2 = root_pos_init[sdr_r] - root_pos_init[sdr_l]197    across = across1 + across2198    across = across / np.sqrt((across ** 2).sum(axis=-1))[..., np.newaxis]199 200    # forward (3,), rotate around y-axis201    forward_init = np.cross(np.array([[0, 1, 0]]), across, axis=-1)202    # forward (3,)203    forward_init = forward_init / np.sqrt((forward_init ** 2).sum(axis=-1))[..., np.newaxis]204 205    #     print(forward_init)206 207    target = np.array([[0, 0, 1]])208    root_quat_init = qbetween_np(forward_init, target)209    root_quat_init = np.ones(positions.shape[:-1] + (4,)) * root_quat_init210 211    positions_b = positions.copy()212 213    positions = qrot_np(root_quat_init, positions)214 215    #     plot_3d_motion("./positions_2.mp4", kinematic_chain, positions, 'title', fps=20)216 217    '''New ground truth positions'''218    global_positions = positions.copy()219 220    # plt.plot(positions_b[:, 0, 0], positions_b[:, 0, 2], marker='*')221    # plt.plot(positions[:, 0, 0], positions[:, 0, 2], marker='o', color='r')222    # plt.xlabel('x')223    # plt.ylabel('z')224    # plt.axis('equal')225    # plt.show()226 227    """ Get Foot Contacts """228 229    def foot_detect(positions, thres):230        velfactor, heightfactor = np.array([thres, thres]), np.array([3.0, 2.0])231 232        feet_l_x = (positions[1:, fid_l, 0] - positions[:-1, fid_l, 0]) ** 2233        feet_l_y = (positions[1:, fid_l, 1] - positions[:-1, fid_l, 1]) ** 2234        feet_l_z = (positions[1:, fid_l, 2] - positions[:-1, fid_l, 2]) ** 2235        #     feet_l_h = positions[:-1,fid_l,1]236        #     feet_l = (((feet_l_x + feet_l_y + feet_l_z) < velfactor) & (feet_l_h < heightfactor)).astype(np.float)237        feet_l = ((feet_l_x + feet_l_y + feet_l_z) < velfactor).astype(np.float)238 239        feet_r_x = (positions[1:, fid_r, 0] - positions[:-1, fid_r, 0]) ** 2240        feet_r_y = (positions[1:, fid_r, 1] - positions[:-1, fid_r, 1]) ** 2241        feet_r_z = (positions[1:, fid_r, 2] - positions[:-1, fid_r, 2]) ** 2242        #     feet_r_h = positions[:-1,fid_r,1]243        #     feet_r = (((feet_r_x + feet_r_y + feet_r_z) < velfactor) & (feet_r_h < heightfactor)).astype(np.float)244        feet_r = (((feet_r_x + feet_r_y + feet_r_z) < velfactor)).astype(np.float)245        return feet_l, feet_r246    #247    feet_l, feet_r = foot_detect(positions, feet_thre)248    # feet_l, feet_r = foot_detect(positions, 0.002)249 250    '''Quaternion and Cartesian representation'''251    r_rot = None252 253    def get_rifke(positions):254        '''Local pose'''255        positions[..., 0] -= positions[:, 0:1, 0]256        positions[..., 2] -= positions[:, 0:1, 2]257        '''All pose face Z+'''258        positions = qrot_np(np.repeat(r_rot[:, None], positions.shape[1], axis=1), positions)259        return positions260 261    def get_quaternion(positions):262        skel = Skeleton(n_raw_offsets, kinematic_chain, "cpu")263        # (seq_len, joints_num, 4)264        quat_params = skel.inverse_kinematics_np(positions, face_joint_indx, smooth_forward=False)265 266        '''Fix Quaternion Discontinuity'''267        quat_params = qfix(quat_params)268        # (seq_len, 4)269        r_rot = quat_params[:, 0].copy()270        #     print(r_rot[0])271        '''Root Linear Velocity'''272        # (seq_len - 1, 3)273        velocity = (positions[1:, 0] - positions[:-1, 0]).copy()274        #     print(r_rot.shape, velocity.shape)275        velocity = qrot_np(r_rot[1:], velocity)276        '''Root Angular Velocity'''277        # (seq_len - 1, 4)278        r_velocity = qmul_np(r_rot[1:], qinv_np(r_rot[:-1]))279        quat_params[1:, 0] = r_velocity280        # (seq_len, joints_num, 4)281        return quat_params, r_velocity, velocity, r_rot282 283    def get_cont6d_params(positions):284        skel = Skeleton(n_raw_offsets, kinematic_chain, "cpu")285        # (seq_len, joints_num, 4)286        quat_params = skel.inverse_kinematics_np(positions, face_joint_indx, smooth_forward=True)287 288        '''Quaternion to continuous 6D'''289        cont_6d_params = quaternion_to_cont6d_np(quat_params)290        # (seq_len, 4)291        r_rot = quat_params[:, 0].copy()292        #     print(r_rot[0])293        '''Root Linear Velocity'''294        # (seq_len - 1, 3)295        velocity = (positions[1:, 0] - positions[:-1, 0]).copy()296        #     print(r_rot.shape, velocity.shape)297        velocity = qrot_np(r_rot[1:], velocity)298        '''Root Angular Velocity'''299        # (seq_len - 1, 4)300        r_velocity = qmul_np(r_rot[1:], qinv_np(r_rot[:-1]))301        # (seq_len, joints_num, 4)302        return cont_6d_params, r_velocity, velocity, r_rot303 304    cont_6d_params, r_velocity, velocity, r_rot = get_cont6d_params(positions)305    positions = get_rifke(positions)306 307    #     trejec = np.cumsum(np.concatenate([np.array([[0, 0, 0]]), velocity], axis=0), axis=0)308    #     r_rotations, r_pos = recover_ric_glo_np(r_velocity, velocity[:, [0, 2]])309 310    # plt.plot(positions_b[:, 0, 0], positions_b[:, 0, 2], marker='*')311    # plt.plot(ground_positions[:, 0, 0], ground_positions[:, 0, 2], marker='o', color='r')312    # plt.plot(trejec[:, 0], trejec[:, 2], marker='^', color='g')313    # plt.plot(r_pos[:, 0], r_pos[:, 2], marker='s', color='y')314    # plt.xlabel('x')315    # plt.ylabel('z')316    # plt.axis('equal')317    # plt.show()318 319    '''Root height'''320    root_y = positions[:, 0, 1:2]321 322    '''Root rotation and linear velocity'''323    # (seq_len-1, 1) rotation velocity along y-axis324    # (seq_len-1, 2) linear velovity on xz plane325    r_velocity = np.arcsin(r_velocity[:, 2:3])326    l_velocity = velocity[:, [0, 2]]327    #     print(r_velocity.shape, l_velocity.shape, root_y.shape)328    root_data = np.concatenate([r_velocity, l_velocity, root_y[:-1]], axis=-1)329 330    '''Get Joint Rotation Representation'''331    # (seq_len, (joints_num-1) *6) quaternion for skeleton joints332    rot_data = cont_6d_params[:, 1:].reshape(len(cont_6d_params), -1)333 334    '''Get Joint Rotation Invariant Position Represention'''335    # (seq_len, (joints_num-1)*3) local joint position336    ric_data = positions[:, 1:].reshape(len(positions), -1)337 338    '''Get Joint Velocity Representation'''339    # (seq_len-1, joints_num*3)340    local_vel = qrot_np(np.repeat(r_rot[:-1, None], global_positions.shape[1], axis=1),341                        global_positions[1:] - global_positions[:-1])342    local_vel = local_vel.reshape(len(local_vel), -1)343 344    data = root_data345    data = np.concatenate([data, ric_data[:-1]], axis=-1)346    data = np.concatenate([data, rot_data[:-1]], axis=-1)347    #     print(data.shape, local_vel.shape)348    data = np.concatenate([data, local_vel], axis=-1)349    data = np.concatenate([data, feet_l, feet_r], axis=-1)350 351    return data, global_positions, positions, l_velocity352 353 354# Recover global angle and positions for rotation data355# root_rot_velocity (B, seq_len, 1)356# root_linear_velocity (B, seq_len, 2)357# root_y (B, seq_len, 1)358# ric_data (B, seq_len, (joint_num - 1)*3)359# rot_data (B, seq_len, (joint_num - 1)*6)360# local_velocity (B, seq_len, joint_num*3)361# foot contact (B, seq_len, 4)362def recover_root_rot_pos(data):363    rot_vel = data[..., 0]364    r_rot_ang = torch.zeros_like(rot_vel).to(data.device)365    '''Get Y-axis rotation from rotation velocity'''366    r_rot_ang[..., 1:] = rot_vel[..., :-1]367    r_rot_ang = torch.cumsum(r_rot_ang, dim=-1)368 369    r_rot_quat = torch.zeros(data.shape[:-1] + (4,)).to(data.device)370    r_rot_quat[..., 0] = torch.cos(r_rot_ang)371    r_rot_quat[..., 2] = torch.sin(r_rot_ang)372 373    r_pos = torch.zeros(data.shape[:-1] + (3,)).to(data.device)374    r_pos[..., 1:, [0, 2]] = data[..., :-1, 1:3]375    '''Add Y-axis rotation to root position'''376    r_pos = qrot(qinv(r_rot_quat), r_pos)377 378    r_pos = torch.cumsum(r_pos, dim=-2)379 380    r_pos[..., 1] = data[..., 3]381    return r_rot_quat, r_pos382 383 384def recover_from_rot(data, joints_num, skeleton):385    r_rot_quat, r_pos = recover_root_rot_pos(data)386 387    r_rot_cont6d = quaternion_to_cont6d(r_rot_quat)388 389    start_indx = 1 + 2 + 1 + (joints_num - 1) * 3390    end_indx = start_indx + (joints_num - 1) * 6391    cont6d_params = data[..., start_indx:end_indx]392    #     print(r_rot_cont6d.shape, cont6d_params.shape, r_pos.shape)393    cont6d_params = torch.cat([r_rot_cont6d, cont6d_params], dim=-1)394    cont6d_params = cont6d_params.view(-1, joints_num, 6)395 396    positions = skeleton.forward_kinematics_cont6d(cont6d_params, r_pos)397 398    return positions399 400 401def recover_from_ric(data, joints_num):402    r_rot_quat, r_pos = recover_root_rot_pos(data)403    positions = data[..., 4:(joints_num - 1) * 3 + 4]404    positions = positions.view(positions.shape[:-1] + (-1, 3))405 406    '''Add Y-axis rotation to local joints'''407    positions = qrot(qinv(r_rot_quat[..., None, :]).expand(positions.shape[:-1] + (4,)), positions)408 409    '''Add root XZ to joints'''410    positions[..., 0] += r_pos[..., 0:1]411    positions[..., 2] += r_pos[..., 2:3]412 413    '''Concate root and joints'''414    positions = torch.cat([r_pos.unsqueeze(-2), positions], dim=-2)415 416    return positions417'''418For Text2Motion Dataset419'''420'''421if __name__ == "__main__":422    example_id = "000021"423    # Lower legs424    l_idx1, l_idx2 = 5, 8425    # Right/Left foot426    fid_r, fid_l = [8, 11], [7, 10]427    # Face direction, r_hip, l_hip, sdr_r, sdr_l428    face_joint_indx = [2, 1, 17, 16]429    # l_hip, r_hip430    r_hip, l_hip = 2, 1431    joints_num = 22432    # ds_num = 8433    data_dir = '../dataset/pose_data_raw/joints/'434    save_dir1 = '../dataset/pose_data_raw/new_joints/'435    save_dir2 = '../dataset/pose_data_raw/new_joint_vecs/'436 437    n_raw_offsets = torch.from_numpy(t2m_raw_offsets)438    kinematic_chain = t2m_kinematic_chain439 440    # Get offsets of target skeleton441    example_data = np.load(os.path.join(data_dir, example_id + '.npy'))442    example_data = example_data.reshape(len(example_data), -1, 3)443    example_data = torch.from_numpy(example_data)444    tgt_skel = Skeleton(n_raw_offsets, kinematic_chain, 'cpu')445    # (joints_num, 3)446    tgt_offsets = tgt_skel.get_offsets_joints(example_data[0])447    # print(tgt_offsets)448 449    source_list = os.listdir(data_dir)450    frame_num = 0451    for source_file in tqdm(source_list):452        source_data = np.load(os.path.join(data_dir, source_file))[:, :joints_num]453        try:454            data, ground_positions, positions, l_velocity = process_file(source_data, 0.002)455            rec_ric_data = recover_from_ric(torch.from_numpy(data).unsqueeze(0).float(), joints_num)456            np.save(pjoin(save_dir1, source_file), rec_ric_data.squeeze().numpy())457            np.save(pjoin(save_dir2, source_file), data)458            frame_num += data.shape[0]459        except Exception as e:460            print(source_file)461            print(e)462 463    print('Total clips: %d, Frames: %d, Duration: %fm' %464          (len(source_list), frame_num, frame_num / 20 / 60))465'''466 467if __name__ == "__main__":468    example_id = "03950_gt"469    # Lower legs470    l_idx1, l_idx2 = 17, 18471    # Right/Left foot472    fid_r, fid_l = [14, 15], [19, 20]473    # Face direction, r_hip, l_hip, sdr_r, sdr_l474    face_joint_indx = [11, 16, 5, 8]475    # l_hip, r_hip476    r_hip, l_hip = 11, 16477    joints_num = 21478    # ds_num = 8479    data_dir = '../dataset/kit_mocap_dataset/joints/'480    save_dir1 = '../dataset/kit_mocap_dataset/new_joints/'481    save_dir2 = '../dataset/kit_mocap_dataset/new_joint_vecs/'482 483    n_raw_offsets = torch.from_numpy(kit_raw_offsets)484    kinematic_chain = kit_kinematic_chain485 486    '''Get offsets of target skeleton'''487    example_data = np.load(os.path.join(data_dir, example_id + '.npy'))488    example_data = example_data.reshape(len(example_data), -1, 3)489    example_data = torch.from_numpy(example_data)490    tgt_skel = Skeleton(n_raw_offsets, kinematic_chain, 'cpu')491    # (joints_num, 3)492    tgt_offsets = tgt_skel.get_offsets_joints(example_data[0])493    # print(tgt_offsets)494 495    source_list = os.listdir(data_dir)496    frame_num = 0497    '''Read source data'''498    for source_file in tqdm(source_list):499        source_data = np.load(os.path.join(data_dir, source_file))[:, :joints_num]500        try:501            name = ''.join(source_file[:-7].split('_')) + '.npy'502            data, ground_positions, positions, l_velocity = process_file(source_data, 0.05)503            rec_ric_data = recover_from_ric(torch.from_numpy(data).unsqueeze(0).float(), joints_num)504            if np.isnan(rec_ric_data.numpy()).any():505                print(source_file)506                continue507            np.save(pjoin(save_dir1, name), rec_ric_data.squeeze().numpy())508            np.save(pjoin(save_dir2, name), data)509            frame_num += data.shape[0]510        except Exception as e:511            print(source_file)512            print(e)513 514    print('Total clips: %d, Frames: %d, Duration: %fm' %515          (len(source_list), frame_num, frame_num / 12.5 / 60))