MCPcopy Create free account
hub / github.com/HYUNJS/SGT / MapDataset

Class MapDataset

projects/Datasets/MOT/common.py:9–77  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

7
8
9class 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

Callers 2

build_mot_train_loaderFunction · 0.90
build_mot_test_loaderFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected