(self, *data_key_batch, crop_box=None, seed=2024, **kwargs)
| 225 | return self.load_video_batch(data_key, data_key2, crop_box=crop_box, seed=seed, **kwargs) |
| 226 | |
| 227 | def load_video_batch(self, *data_key_batch, crop_box=None, seed=2024, **kwargs): |
| 228 | rng = np.random.default_rng(seed + hash(data_key_batch[0]) % 10000) |
| 229 | # read video |
| 230 | import decord |
| 231 | decord.bridge.set_bridge('torch') |
| 232 | readers = [] |
| 233 | for data_k in data_key_batch: |
| 234 | reader = decord.VideoReader(data_k) |
| 235 | readers.append(reader) |
| 236 | |
| 237 | fps = readers[0].get_avg_fps() |
| 238 | length = min([len(r) for r in readers]) |
| 239 | frame_timestamps = [readers[0].get_frame_timestamp(i) for i in range(length)] |
| 240 | frame_timestamps = np.array(frame_timestamps, dtype=np.float32) |
| 241 | h, w = readers[0].next().shape[:2] |
| 242 | frame_ids, (x1, x2, y1, y2), (oh, ow), fps = self._get_frameid_bbox(fps, frame_timestamps, h, w, crop_box, rng) |
| 243 | |
| 244 | # preprocess video |
| 245 | videos = [reader.get_batch(frame_ids)[:, y1:y2, x1:x2, :] for reader in readers] |
| 246 | videos = [self._video_preprocess(video, oh, ow) for video in videos] |
| 247 | return *videos, frame_ids, (oh, ow), fps |
| 248 | # return videos if len(videos) > 1 else videos[0] |
| 249 | |
| 250 | |
| 251 | def prepare_source(src_video, src_mask, src_ref_images, num_frames, image_size, device): |
no test coverage detected