| 22 | |
| 23 | |
| 24 | class ChangeSim(data.Dataset): |
| 25 | def __init__(self, crop_size=(320, 240), num_classes=5, set='train'): |
| 26 | """ |
| 27 | ChangeSim Dataloader |
| 28 | Please download ChangeSim Dataset in https://github.com/SAMMiCA/ChangeSim |
| 29 | |
| 30 | Args: |
| 31 | crop_size (tuple): Image resize shape (H,W) (default: (320, 240)) |
| 32 | num_classes (int): Number of target change detection class |
| 33 | 5 for multi-class change detection |
| 34 | 2 for binary change detection (default: 5) |
| 35 | set (str): 'train' or 'test' (defalut: 'train') |
| 36 | """ |
| 37 | self.crop_size = crop_size |
| 38 | self.num_classes = num_classes |
| 39 | self.set = set |
| 40 | self.blacklist=[] |
| 41 | train_list = ['Warehouse_0', 'Warehouse_1', 'Warehouse_2', 'Warehouse_3', 'Warehouse_4', 'Warehouse_5'] |
| 42 | test_list = ['Warehouse_6', 'Warehouse_7', 'Warehouse_8', 'Warehouse_9'] |
| 43 | self.image_total_files = [] |
| 44 | if set == 'train': |
| 45 | for map in train_list: |
| 46 | self.image_total_files += glob.glob('../Query/Query_Seq_Train/' + map + '/Seq_0/rgb/*.png') |
| 47 | self.image_total_files += glob.glob('../Query/Query_Seq_Train/' + map + '/Seq_1/rgb/*.png') |
| 48 | elif set == 'test': |
| 49 | for map in test_list: |
| 50 | self.image_total_files += glob.glob('../Query/Query_Seq_Test/' + map + '/Seq_0/rgb/*.png') |
| 51 | self.image_total_files += glob.glob('../Query/Query_Seq_Test/' + map + '/Seq_1/rgb/*.png') |
| 52 | # if not max_iters == None: |
| 53 | # self.image_total_files = self.image_total_files * int(np.ceil(float(max_iters) / len(self.image_total_files))) |
| 54 | # self.image_total_files = self.image_total_files[:max_iters] |
| 55 | |
| 56 | self.seg = Object_Labeling.SegHelper(idx2color_path='./utils/idx2color.txt', num_class=self.num_classes) |
| 57 | |
| 58 | #### Color Transform #### |
| 59 | self.color_transform = transforms.Compose([transforms.ColorJitter(0.4, 0.4, 0.4, 0.25), |
| 60 | transforms.ToTensor()]) |
| 61 | # self.transform = Compose([Resize(crop_size), ToTensor()]) |
| 62 | |
| 63 | def __len__(self): |
| 64 | return len(self.image_total_files) |
| 65 | |
| 66 | def __getitem__(self, index): |
| 67 | # Train set |
| 68 | if self.set == 'train': |
| 69 | loss = nn.L1Loss() |
| 70 | while True: |
| 71 | if index in self.blacklist: |
| 72 | index=random.randint(0,self.__len__()-1) |
| 73 | continue |
| 74 | |
| 75 | test_rgb_path = self.image_total_files[index] |
| 76 | file_idx = test_rgb_path.split('/')[-1].split('.')[0] # ~~ of ~~.png |
| 77 | |
| 78 | ref_pose_find_path = test_rgb_path.replace(f'rgb/{file_idx}.png',f't0/idx/{file_idx}.txt') |
| 79 | f = open(ref_pose_find_path,'r',encoding='utf8') |
| 80 | ref_pose_idx = int(f.readlines()[0]) |
| 81 | g2o_path = test_rgb_path.replace('/Query/Query_Seq_Train','/Reference/Ref_Seq_Train').replace(f'rgb/{file_idx}.png',f'raw/poses.g2o') |