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