MCPcopy Create free account
hub / github.com/MediaBrain-SJTU/MemoNet / preprocess

Class preprocess

ETH/data/preprocessor.py:7–194  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

5
6
7class preprocess(object):
8
9 def __init__(self, data_root, seq_name, parser, log, split='train', phase='training'):
10 self.parser = parser
11 self.dataset = parser.dataset
12 self.data_root = data_root
13 self.past_frames = parser.past_frames
14 self.future_frames = parser.future_frames
15 self.frame_skip = parser.get('frame_skip', 1)
16 self.min_past_frames = parser.get('min_past_frames', self.past_frames)
17 self.min_future_frames = parser.get('min_future_frames', self.future_frames)
18 self.traj_scale = parser.traj_scale
19
20 self.past_traj_scale = parser.traj_scale
21 # print('Scale of ETH datasets:', self.traj_scale)
22 self.load_map = parser.get('load_map', False)
23 self.map_version = parser.get('map_version', '0.1')
24 self.seq_name = seq_name
25 self.split = split
26 self.phase = phase
27 self.log = log
28
29 if parser.dataset == 'nuscenes_pred':
30 label_path = os.path.join(data_root, 'label/{}/{}.txt'.format(split, seq_name))
31 delimiter = ' '
32 elif parser.dataset in {'eth', 'hotel', 'univ', 'zara1', 'zara2'}:
33 label_path = f'{data_root}/{parser.dataset}/{seq_name}.txt'
34 delimiter = ' '
35 else:
36 assert False, 'error'
37
38 self.gt = np.genfromtxt(label_path, delimiter=delimiter, dtype=str)
39 frames = self.gt[:, 0].astype(np.float32).astype(np.int)
40 fr_start, fr_end = frames.min(), frames.max()
41 self.init_frame = fr_start
42 self.num_fr = fr_end + 1 - fr_start
43
44 if self.load_map:
45 self.load_scene_map()
46 else:
47 self.geom_scene_map = None
48
49 self.class_names = class_names = {'Pedestrian': 1, 'Car': 2, 'Cyclist': 3, 'Truck': 4, 'Van': 5, 'Tram': 6, 'Person': 7, \
50 'Misc': 8, 'DontCare': 9, 'Traffic_cone': 10, 'Construction_vehicle': 11, 'Barrier': 12, 'Motorcycle': 13, \
51 'Bicycle': 14, 'Bus': 15, 'Trailer': 16, 'Emergency': 17, 'Construction': 18}
52 for row_index in range(len(self.gt)):
53 self.gt[row_index][2] = class_names[self.gt[row_index][2]]
54 self.gt = self.gt.astype('float32')
55 self.xind, self.zind = 13, 15
56
57 def GetID(self, data):
58 id = []
59 for i in range(data.shape[0]):
60 id.append(data[i, 1].copy())
61 return id
62
63 def TotalFrame(self):
64 return self.num_fr

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected