| 24 | |
| 25 | |
| 26 | class 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 | |