(self,
*data_key_batch,
crop_box=None,
seed=2024,
**kwargs)
| 238 | data_key, data_key2, crop_box=crop_box, seed=seed, **kwargs) |
| 239 | |
| 240 | def load_video_batch(self, |
| 241 | *data_key_batch, |
| 242 | crop_box=None, |
| 243 | seed=2024, |
| 244 | **kwargs): |
| 245 | rng = np.random.default_rng(seed + hash(data_key_batch[0]) % 10000) |
| 246 | # read video |
| 247 | import decord |
| 248 | decord.bridge.set_bridge('torch') |
| 249 | readers = [] |
| 250 | for data_k in data_key_batch: |
| 251 | reader = decord.VideoReader(data_k) |
| 252 | readers.append(reader) |
| 253 | |
| 254 | fps = readers[0].get_avg_fps() |
| 255 | length = min([len(r) for r in readers]) |
| 256 | frame_timestamps = [ |
| 257 | readers[0].get_frame_timestamp(i) for i in range(length) |
| 258 | ] |
| 259 | frame_timestamps = np.array(frame_timestamps, dtype=np.float32) |
| 260 | h, w = readers[0].next().shape[:2] |
| 261 | frame_ids, (x1, x2, y1, y2), (oh, ow), fps = self._get_frameid_bbox( |
| 262 | fps, frame_timestamps, h, w, crop_box, rng) |
| 263 | |
| 264 | # preprocess video |
| 265 | videos = [ |
| 266 | reader.get_batch(frame_ids)[:, y1:y2, x1:x2, :] |
| 267 | for reader in readers |
| 268 | ] |
| 269 | videos = [self._video_preprocess(video, oh, ow) for video in videos] |
| 270 | return *videos, frame_ids, (oh, ow), fps |
| 271 | # return videos if len(videos) > 1 else videos[0] |
| 272 | |
| 273 | |
| 274 | def prepare_source(src_video, src_mask, src_ref_images, num_frames, image_size, |
no test coverage detected