(
self, root_videos_dir, root_masks_dir, root_flows_dir, root_flowmasks_dir, root_outputs_dir,
dataset_args,
batch_size, shuffle, validation_split,
num_workers, video_names_filename=None, training=True, name=None
)
| 8 | |
| 9 | class MaskedFrameDataLoader(BaseDataLoader): |
| 10 | def __init__( |
| 11 | self, root_videos_dir, root_masks_dir, root_flows_dir, root_flowmasks_dir, root_outputs_dir, |
| 12 | dataset_args, |
| 13 | batch_size, shuffle, validation_split, |
| 14 | num_workers, video_names_filename=None, training=True, name=None |
| 15 | ): |
| 16 | # Input directories |
| 17 | self.rids = RootInputDirectories( |
| 18 | root_videos_dir, root_masks_dir, root_flows_dir=root_flows_dir, root_flowmasks_dir=root_flowmasks_dir, video_names_filename=video_names_filename) |
| 19 | |
| 20 | # Output directories |
| 21 | self.rods = RootOutputDirectories(root_outputs_dir) |
| 22 | self.name = name |
| 23 | |
| 24 | # Dataset, default is video |
| 25 | if 'type' not in dataset_args: |
| 26 | dataset_args['type'] = 'video' |
| 27 | if dataset_args['type'] == 'video': |
| 28 | Dataset = VideoFrameAndMaskDataset |
| 29 | elif dataset_args['type'] == 'CelebA': |
| 30 | Dataset = CelebAFrameAndMaskDataset |
| 31 | elif dataset_args['type'] == 'Places2': |
| 32 | Dataset = Places2FrameAndMaskDataset |
| 33 | elif dataset_args['type'] == 'super_resolution': |
| 34 | Dataset = VideoSuperResolutionDataset |
| 35 | elif dataset_args['type'] == 'video_flow': |
| 36 | Dataset = VideoFrameMaskAndFlowDataset |
| 37 | elif dataset_args['type'] == 'video_stationary': |
| 38 | Dataset = VideoFrameAndStationaryMaskDataset |
| 39 | elif dataset_args['type'] == 'video_flow_stationary': |
| 40 | Dataset = VideoFrameStationaryMaskAndFlowDataset |
| 41 | elif dataset_args['type'] == 'video_flow_mix': |
| 42 | Dataset = VideoFrameMaskAndFlowMixDataset |
| 43 | else: |
| 44 | raise NotImplementedError(f"Dataset type {dataset_args['type']}") |
| 45 | |
| 46 | self.dataset = Dataset( |
| 47 | self.rids, self.rods, dataset_args, |
| 48 | ) |
| 49 | super().__init__(self.dataset, batch_size, shuffle, validation_split, num_workers) |
nothing calls this directly
no test coverage detected