(self, datasets, image_size, is_train=True)
| 33 | |
| 34 | class MyData(data.Dataset): |
| 35 | def __init__(self, datasets, image_size, is_train=True): |
| 36 | self.size_train = image_size |
| 37 | self.size_test = image_size |
| 38 | self.keep_size = not config.size |
| 39 | self.data_size = config.size |
| 40 | self.is_train = is_train |
| 41 | self.load_all = config.load_all |
| 42 | self.device = config.device |
| 43 | valid_extensions = [".png", ".jpg", ".PNG", ".JPG", ".JPEG"] |
| 44 | |
| 45 | if self.is_train and config.auxiliary_classification: |
| 46 | self.cls_name2id = { |
| 47 | _name: _id for _id, _name in enumerate(class_labels_TR_sorted) |
| 48 | } |
| 49 | self.transform_image = transforms.Compose( |
| 50 | [ |
| 51 | transforms.Resize(self.data_size[::-1]), |
| 52 | transforms.ToTensor(), |
| 53 | transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]), |
| 54 | ][self.load_all or self.keep_size :] |
| 55 | ) |
| 56 | self.transform_label = transforms.Compose( |
| 57 | [ |
| 58 | transforms.Resize(self.data_size[::-1]), |
| 59 | transforms.ToTensor(), |
| 60 | ][self.load_all or self.keep_size :] |
| 61 | ) |
| 62 | dataset_root = os.path.join(config.data_root_dir, config.task) |
| 63 | # datasets can be a list of different datasets for training on combined sets. |
| 64 | self.image_paths = [] |
| 65 | for dataset in datasets.split("+"): |
| 66 | image_root = os.path.join(dataset_root, dataset, "im") |
| 67 | self.image_paths += [ |
| 68 | os.path.join(image_root, p) |
| 69 | for p in os.listdir(image_root) |
| 70 | if any(p.endswith(ext) for ext in valid_extensions) |
| 71 | ] |
| 72 | self.label_paths = [] |
| 73 | for p in self.image_paths: |
| 74 | for ext in valid_extensions: |
| 75 | ## 'im' and 'gt' may need modifying |
| 76 | p_gt = p.replace("/im/", "/gt/")[: -(len(p.split(".")[-1]) + 1)] + ext |
| 77 | file_exists = False |
| 78 | if os.path.exists(p_gt): |
| 79 | self.label_paths.append(p_gt) |
| 80 | file_exists = True |
| 81 | break |
| 82 | if not file_exists: |
| 83 | print("Not exists:", p_gt) |
| 84 | |
| 85 | if len(self.label_paths) != len(self.image_paths): |
| 86 | set_image_paths = set( |
| 87 | [os.path.splitext(p.split(os.sep)[-1])[0] for p in self.image_paths] |
| 88 | ) |
| 89 | set_label_paths = set( |
| 90 | [os.path.splitext(p.split(os.sep)[-1])[0] for p in self.label_paths] |
| 91 | ) |
| 92 | print("Path diff:", set_image_paths - set_label_paths) |
nothing calls this directly
no test coverage detected