MCPcopy Create free account
hub / github.com/FoundationVision/ByteTrack / DetMOTDetection

Class DetMOTDetection

tutorials/motr/joint.py:26–192  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

24
25
26class DetMOTDetection:
27 def __init__(self, args, data_txt_path: str, seqs_folder, dataset2transform):
28 self.args = args
29 self.dataset2transform = dataset2transform
30 self.num_frames_per_batch = max(args.sampler_lengths)
31 self.sample_mode = args.sample_mode
32 self.sample_interval = args.sample_interval
33 self.vis = args.vis
34 self.video_dict = {}
35
36 with open(data_txt_path, 'r') as file:
37 self.img_files = file.readlines()
38 self.img_files = [osp.join(seqs_folder, x.strip()) for x in self.img_files]
39 self.img_files = list(filter(lambda x: len(x) > 0, self.img_files))
40
41 self.label_files = [(x.replace('images', 'labels_with_ids').replace('.png', '.txt').replace('.jpg', '.txt'))
42 for x in self.img_files]
43 # The number of images per sample: 1 + (num_frames - 1) * interval.
44 # The number of valid samples: num_images - num_image_per_sample + 1.
45 self.item_num = len(self.img_files) - (self.num_frames_per_batch - 1) * self.sample_interval
46
47 self._register_videos()
48
49 # video sampler.
50 self.sampler_steps: list = args.sampler_steps
51 self.lengths: list = args.sampler_lengths
52 print("sampler_steps={} lenghts={}".format(self.sampler_steps, self.lengths))
53 if self.sampler_steps is not None and len(self.sampler_steps) > 0:
54 # Enable sampling length adjustment.
55 assert len(self.lengths) > 0
56 assert len(self.lengths) == len(self.sampler_steps) + 1
57 for i in range(len(self.sampler_steps) - 1):
58 assert self.sampler_steps[i] < self.sampler_steps[i + 1]
59 self.item_num = len(self.img_files) - (self.lengths[-1] - 1) * self.sample_interval
60 self.period_idx = 0
61 self.num_frames_per_batch = self.lengths[0]
62 self.current_epoch = 0
63
64 def _register_videos(self):
65 for label_name in self.label_files:
66 video_name = '/'.join(label_name.split('/')[:-1])
67 if video_name not in self.video_dict:
68 print("register {}-th video: {} ".format(len(self.video_dict) + 1, video_name))
69 self.video_dict[video_name] = len(self.video_dict)
70 # assert len(self.video_dict) <= 300
71
72 def set_epoch(self, epoch):
73 self.current_epoch = epoch
74 if self.sampler_steps is None or len(self.sampler_steps) == 0:
75 # fixed sampling length.
76 return
77
78 for i in range(len(self.sampler_steps)):
79 if epoch >= self.sampler_steps[i]:
80 self.period_idx = i + 1
81 print("set epoch: epoch {} period_idx={}".format(epoch, self.period_idx))
82 self.num_frames_per_batch = self.lengths[self.period_idx]
83

Callers 1

buildFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected