Spaces:
Running on Zero
Running on Zero
File size: 11,484 Bytes
2680bd5 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 | import warnings, math
from os.path import join as pjoin
from pathlib import Path
import numpy as np
import torch
from tqdm import tqdm
from typing import Literal
from torch.utils.data import Dataset
from .. import dist as dist_utils
from .tensors import motion_action_collate, motion_text_collate, motion_text_treble_collate
from ..metric.interaction import give_contact_label
from ...feature.single_motioncode import MotionCoder
from ...constant import INTRA_TIP_CONTACT_THRESH, TIP_PALM_CONTACT_THRESH, PALM_PALM_CONTACT_THRESH, SKELETON_CHAIN
class HandXDataset(Dataset):
contact_label: bool
normalize: bool
repr: Literal['joint_pos', 'joint_rot', 'joint_pos_w_scalar_rot', "joint_pos_w_axisangle_rot"]
data_file_name: str
data_dir: str
def __init__(self, split: str, debug=False, *args, **kwargs):
super(HandXDataset, self).__init__()
self.debug = debug
self.data_dir = kwargs['data_dir']
self.data_file_name = kwargs['data_file_name']
self.repr = kwargs['repr']
self.normalize = kwargs['normalize']
self.contact_label = kwargs.get("contact_label", False)
self.ratio = kwargs.get("ratio", 1.0)
if split == 'train':
self.data_file_name = "train_" + self.data_file_name
self.mano_file_name = 'train_mano.npz'
elif split == 'val':
self.data_file_name = 'test_' + self.data_file_name
self.mano_file_name = 'test_mano.npz'
self.split = split
self.load_motion()
if self.normalize:
self._calc_mean_std()
self.collate_fn = motion_text_treble_collate
def _calc_mean_std(self):
if self.split == 'val':
dist_utils.barrier()
print("Val split doesn't need create mean and std files.")
mean_std_path = pjoin(self.data_dir, f'mean_std_{self.repr}')
assert Path(mean_std_path).exists(), f"Mean and std not found at {mean_std_path}. Please run the dataset preparation script."
self.mean = np.load(pjoin(mean_std_path, 'mean.npy'))
self.std = np.load(pjoin(mean_std_path, 'std.npy'))
return
total_frames = 0
first_motion = next(iter(self.data_dict.values()))['motion']
feature_shape = first_motion.shape[1:]
sum_of_data = np.zeros(feature_shape, dtype=np.float64)
sum_of_squares = np.zeros(feature_shape, dtype=np.float64)
for data in self.data_dict.values():
motion_data = data['motion']
if motion_data.ndim != 3:
raise ValueError(f"Expected motion data to be 3D, but got {motion_data.ndim}D shape.")
num_frames_in_batch = motion_data.shape[0]
total_frames += num_frames_in_batch
sum_of_data += np.sum(motion_data, axis=0)
sum_of_squares += np.sum(np.square(motion_data), axis=0)
if total_frames == 0:
warnings.warn("No frames found in the dataset. Mean and std will be zero.")
self.mean = np.zeros(feature_shape)
self.std = np.zeros(feature_shape)
else:
self.mean = sum_of_data / total_frames
variance = (sum_of_squares / total_frames) - np.square(self.mean)
variance[variance < 0] = 0
self.std = np.sqrt(variance)
self.std[self.std < 1e-4] = 1.0
if dist_utils.is_main_process():
save_path = (Path(self.data_dir) / f'mean_std_{self.repr}').as_posix()
Path(save_path).mkdir(parents=True, exist_ok=True)
# mylogger.info(f"Saving mean and std to {save_path}")
np.save((Path(save_path) / 'mean.npy').as_posix(), self.mean)
np.save((Path(save_path) / 'std.npy').as_posix(), self.std)
def inv_transform(self, data):
if isinstance(data, torch.Tensor):
tmp_mean = torch.from_numpy(self.mean).to(data.device).float()
tmp_std = torch.from_numpy(self.std).to(data.device).float()
ret = data * tmp_std.reshape(-1) + tmp_mean.reshape(-1) # [B, T, J*C]
return ret
elif isinstance(data, np.ndarray):
ret = data * self.std.reshape(-1) + self.mean.reshape(-1) # [B, T, J*C]
return ret
else:
raise TypeError(f"Unsupported data type: {type(data)}. Expected torch.Tensor or np.ndarray.")
def _get_axisangle_rotation(self, mano_pose:np.ndarray) -> np.ndarray:
'''
mano_pose: (T, 48)
return: (T, J, 3)
'''
mano_pose = mano_pose.reshape(mano_pose.shape[0], -1, 3) # (T, 16, 3)
zero_padding = np.zeros((mano_pose.shape[0], 21-16, 3)) # (T, 5, 3)
return np.concatenate([mano_pose, zero_padding], axis=1) # (T, 21, 3)
def _get_scalar_rotation(self, single_motion_seq:np.ndarray, side:Literal['left', 'right']):
'''
single_motion_seq: (T, J, 3)
'''
temp_motioncoder = MotionCoder(single_motion_seq, isright=(side=='right'))
temp_motioncoder.get_local_coordinate()
local_motion = temp_motioncoder.local_motion # (T, J, 3)
additional_scalar_rotation = np.zeros((single_motion_seq.shape[0], single_motion_seq.shape[1])) # (T, J)
for skeleton_chain in SKELETON_CHAIN:
additional_scalar_rotation[:, skeleton_chain[0]] = 0
additional_scalar_rotation[:, skeleton_chain[-1]] = 0
for s in range(1, len(skeleton_chain) - 1):
j = skeleton_chain[s]
pre_j = skeleton_chain[s - 1]
nxt_j = skeleton_chain[s + 1]
v1 = (local_motion[:, j, :] - local_motion[:, pre_j, :])[:, [0, 2]] # (T, 2)
v2 = (local_motion[:, nxt_j, :] - local_motion[:, j, :])[:, [0, 2]] # (T, 2)
v1_direction_angle = np.arctan2(v1[:, 1], v1[:, 0]) # (T,)
v2_direction_angle = np.arctan2(v2[:, 1], v2[:, 0]) # (T,)
angle_diff = v2_direction_angle - v1_direction_angle # (T,)
if side == 'right':
angle_diff = -angle_diff
additional_scalar_rotation[:, j] = angle_diff
return additional_scalar_rotation # (T, J)
def load_motion(self):
motion_data = dict(np.load(pjoin(self.data_dir, self.data_file_name), allow_pickle=True))
for key in motion_data:
motion_data[key] = motion_data[key].item()
# mano_data = dict(np.load(pjoin(self.data_dir, self.mano_file_name), allow_pickle=True))
# for key in tqdm(mano_data, desc=f"RANK {dist_utils.get_rank()} | loading mano data for {self.split} split"):
# mano_data[key] = mano_data[key].item()
self.data_dict = dict()
self.name_list = []
self.length_list = []
for clip_name in tqdm(sorted(motion_data.keys()), desc=f"RANK {dist_utils.get_rank()} | processing {self.split} data"):
motion = motion_data[clip_name]['motion'] # (T, 2, J, 3)
if self.repr == 'joint_pos_w_axisangle_rot':
left_mano_pose = mano_data[clip_name]['left_pose'] # (T, 48)
right_mano_pose = mano_data[clip_name]['right_pose'] # (T, 48)
left_rot = self._get_axisangle_rotation(left_mano_pose) # (T, J, 3)
right_rot = self._get_axisangle_rotation(right_mano_pose) # (T, J, 3)
motion = np.concatenate([
motion,
np.stack([left_rot, right_rot], axis=1) # (T, 2, J, 3)
], axis=-1) # (T, 2, J, 6)
motion = motion.reshape(motion.shape[0], -1, 6) # (T, 2J, 6)
transl = np.mean((motion[:, 0, :3] + motion[:, 21, :3]) / 2, axis=0) # (3,)
motion[:, :, :3] -= transl
elif self.repr == 'joint_pos_w_scalar_rot':
left_rot_scalar = self._get_scalar_rotation(motion[:, 0], side='left') # (T, J)
left_rot_scalar = np.nan_to_num(left_rot_scalar, nan=0.0, posinf=0.0, neginf=0.0)
right_rot_scalar = self._get_scalar_rotation(motion[:, 1], side='right') # (T, J)
right_rot_scalar = np.nan_to_num(right_rot_scalar, nan=0.0, posinf=0.0, neginf=0.0)
motion = np.concatenate([
motion,
np.stack([left_rot_scalar, right_rot_scalar], axis=1)[:, :, :, np.newaxis] # (T, 2, J, 1)
], axis=-1) # (T, 2, J, 4)
motion = motion.reshape(motion.shape[0], -1, 4) # (T, 2J, 4)
transl = np.mean((motion[:, 0, :3] + motion[:, 21, :3]) / 2, axis=0) # (3,)
motion[:, :, :3] -= transl
elif self.repr == 'joint_pos':
transl = np.mean((motion[:, 0] + motion[:, 21]) / 2, axis=0) # (3,)
motion[:, :, :3] -= transl
pass
else:
raise NotImplementedError(f"Representation {self.repr} not implemented.")
assert len(motion_data[clip_name]['left_annotation']) == len(motion_data[clip_name]['right_annotation']) == len(motion_data[clip_name]['interaction_annotation']), f"Annotation length mismatch for clip {clip_name}"
annotation_count = len(motion_data[clip_name]['left_annotation'])
for j in range(annotation_count):
name = f"{clip_name}_ann{j}"
self.name_list.append(name)
self.length_list.append(motion.shape[0])
self.data_dict[name] = dict(
motion=motion,
text={
'left': motion_data[clip_name]['left_annotation'][j],
'right': motion_data[clip_name]['right_annotation'][j],
'two_hands_relation': motion_data[clip_name]['interaction_annotation'][j]
}
)
self.length_list = np.array(self.length_list)
if self.split == 'train' and self.ratio < 1.0:
samples_count = int(math.ceil(len(self.name_list) * self.ratio))
random_indices = np.random.choice(len(self.name_list), size=samples_count, replace=False)
self.name_list = [self.name_list[i] for i in random_indices]
self.length_list = self.length_list[random_indices]
self.data_dict = {name: self.data_dict[name] for name in self.name_list}
def __getitem__(self, index):
if self.debug:
print(f"motion name: {self.name_list[index]}")
print(f"self.name_list is UNIQUE: {len(self.name_list) == len(set(self.name_list))}")
motion = self.data_dict[self.name_list[index]]['motion'] # (T, 2J, C)
m_length = self.length_list[index]
text = self.data_dict[self.name_list[index]]['text']
if self.contact_label == True:
contact_label = give_contact_label(
motion[:, :, :3].reshape(motion.shape[0], 2, -1, 3), # (T, 2, J, 3)
tip_tip_threshold=INTRA_TIP_CONTACT_THRESH,
tip_palm_threshold=TIP_PALM_CONTACT_THRESH,
palm_palm_threshold=PALM_PALM_CONTACT_THRESH
)
if self.normalize:
motion = (motion - self.mean[np.newaxis]) / self.std
if self.contact_label:
return motion.transpose((1, 2, 0)), m_length, text, contact_label
else:
return motion.transpose((1, 2, 0)), m_length, text
def __len__(self):
return len(self.data_dict)
|