| 7 | |
| 8 | |
| 9 | class MapDataset(data.Dataset): |
| 10 | |
| 11 | def __init__(self, cfg, dataset, map_func, is_train): |
| 12 | self._cfg = cfg |
| 13 | self._dataset = dataset |
| 14 | self._map_func = PicklableWrapper(map_func) # wrap so that a lambda will work |
| 15 | |
| 16 | self._rng = random.Random(42) |
| 17 | self._fallback_candidates = set(range(len(dataset))) |
| 18 | |
| 19 | self.num_samples = cfg.DATALOADER.NUM_SAMPLES_TRAIN if is_train else cfg.DATALOADER.NUM_SAMPLES_TEST |
| 20 | self.max_frame_dist = cfg.DATALOADER.MAX_FRAME_DIST_TRAIN if is_train else cfg.DATALOADER.MAX_FRAME_DIST_TEST |
| 21 | self.min_frame_dist = cfg.DATALOADER.MIN_FRAME_DIST_TRAIN if is_train else cfg.DATALOADER.MIN_FRAME_DIST_TEST |
| 22 | self.pre_sample_flag = cfg.DATALOADER.PRE_SAMPLE_FLAG |
| 23 | if not is_train and self.pre_sample_flag: ## pre-sampling for only test-phase (sparse video testing environment) |
| 24 | self.max_frame_dist, self.min_frame_dist = 1, 1 |
| 25 | assert self.max_frame_dist >= self.min_frame_dist, "Max frame dist ({}) must be >= Min frame dist ({})".format(self.max_frame_dist, self.min_frame_dist) |
| 26 | |
| 27 | def __len__(self): |
| 28 | return len(self._dataset) |
| 29 | |
| 30 | def __getitem__(self, idx): |
| 31 | retry_count = 0 |
| 32 | cur_idx = int(idx) |
| 33 | |
| 34 | while True: |
| 35 | data = self._map_func(self.get_data(cur_idx)) |
| 36 | if data is not None: |
| 37 | self._fallback_candidates.add(cur_idx) |
| 38 | return data |
| 39 | |
| 40 | # _map_func fails for this idx, use a random new index from the pool |
| 41 | retry_count += 1 |
| 42 | self._fallback_candidates.discard(cur_idx) |
| 43 | cur_idx = self._rng.sample(self._fallback_candidates, k=1)[0] |
| 44 | |
| 45 | if retry_count >= 3: |
| 46 | logger = logging.getLogger(__name__) |
| 47 | logger.warning( |
| 48 | "Failed to apply `_map_func` for idx: {}, retry count: {}".format( |
| 49 | idx, retry_count |
| 50 | ) |
| 51 | ) |
| 52 | |
| 53 | def get_data(self, cur_idx): |
| 54 | data = [] |
| 55 | if self.num_samples == 1: |
| 56 | data.append(self._dataset[cur_idx]) |
| 57 | else: |
| 58 | sequence_name = self._dataset[cur_idx]['sequence_name'] |
| 59 | frame_offsets = self.get_frame_offsets() |
| 60 | for i in frame_offsets: |
| 61 | next_idx = cur_idx + i |
| 62 | if next_idx < len(self._dataset) and self._dataset[next_idx]['sequence_name'] == sequence_name: |
| 63 | data.append(self._dataset[next_idx]) |
| 64 | else: |
| 65 | data.append(data[-1]) |
| 66 | return data |
no outgoing calls
no test coverage detected