(self, lmdb_dir, n_poses, subdivision_stride, pose_resampling_fps, mean_pose, mean_dir_vec,
speaker_model=None, remove_word_timing=False, save_flag=False)
| 89 | |
| 90 | class SpeechMotionDataset(Dataset): |
| 91 | def __init__(self, lmdb_dir, n_poses, subdivision_stride, pose_resampling_fps, mean_pose, mean_dir_vec, |
| 92 | speaker_model=None, remove_word_timing=False, save_flag=False): |
| 93 | |
| 94 | self.lmdb_dir = lmdb_dir |
| 95 | self.n_poses = n_poses |
| 96 | self.subdivision_stride = subdivision_stride |
| 97 | self.skeleton_resampling_fps = pose_resampling_fps |
| 98 | self.mean_dir_vec = mean_dir_vec |
| 99 | self.remove_word_timing = remove_word_timing |
| 100 | |
| 101 | self.expected_audio_length = int(round(n_poses / pose_resampling_fps * 16000)) |
| 102 | self.expected_spectrogram_length = utils.data_utils.calc_spectrogram_length_from_motion_length( |
| 103 | n_poses, pose_resampling_fps) |
| 104 | |
| 105 | self.lang_model = None |
| 106 | self.save_flag = save_flag |
| 107 | |
| 108 | #self.beat_path = 'beat_resave' |
| 109 | self.beat_path = '../Gesture-Generation-from-Trimodal-Context/double_feat' |
| 110 | |
| 111 | logging.info("Reading data '{}'...".format(lmdb_dir)) |
| 112 | preloaded_dir = lmdb_dir + '_cache' |
| 113 | if not os.path.exists(preloaded_dir): |
| 114 | logging.info('Creating the dataset cache...') |
| 115 | assert mean_dir_vec is not None |
| 116 | if mean_dir_vec.shape[-1] != 3: |
| 117 | mean_dir_vec = mean_dir_vec.reshape(mean_dir_vec.shape[:-1] + (-1, 3)) |
| 118 | n_poses_extended = int(round(n_poses * 1.25)) # some margin |
| 119 | data_sampler = DataPreprocessor(lmdb_dir, preloaded_dir, n_poses_extended, |
| 120 | subdivision_stride, pose_resampling_fps, mean_pose, mean_dir_vec) |
| 121 | data_sampler.run() |
| 122 | else: |
| 123 | logging.info('Found the cache {}'.format(preloaded_dir)) |
| 124 | |
| 125 | # init lmdb |
| 126 | self.lmdb_env = lmdb.open(preloaded_dir, readonly=True, lock=False) |
| 127 | with self.lmdb_env.begin() as txn: |
| 128 | self.n_samples = txn.stat()['entries'] |
| 129 | |
| 130 | # make a speaker model |
| 131 | if speaker_model is None or speaker_model == 0: |
| 132 | precomputed_model = lmdb_dir + '_speaker_model.pkl' |
| 133 | if not os.path.exists(precomputed_model): |
| 134 | self._make_speaker_model(lmdb_dir, precomputed_model) |
| 135 | else: |
| 136 | with open(precomputed_model, 'rb') as f: |
| 137 | self.speaker_model = pickle.load(f) |
| 138 | else: |
| 139 | self.speaker_model = speaker_model |
| 140 | |
| 141 | def __len__(self): |
| 142 | return self.n_samples |
nothing calls this directly
no test coverage detected