Takes two datasets (labelled_dataset, unlabelled_dataset) and merges them Allows you to iterate over them in parallel
| 11 | return subsample_indices |
| 12 | import copy |
| 13 | class MergedDataset(Dataset): |
| 14 | |
| 15 | """ |
| 16 | Takes two datasets (labelled_dataset, unlabelled_dataset) and merges them |
| 17 | Allows you to iterate over them in parallel |
| 18 | """ |
| 19 | |
| 20 | def __init__(self, labelled_dataset, unlabelled_dataset=None): |
| 21 | _labelled_dataset = copy.deepcopy(labelled_dataset) |
| 22 | _unlabelled_dataset = copy.deepcopy(unlabelled_dataset) |
| 23 | |
| 24 | self.labelled_dataset = _labelled_dataset |
| 25 | self.unlabelled_dataset = _unlabelled_dataset |
| 26 | |
| 27 | self.target_transform = None |
| 28 | if unlabelled_dataset is not None: |
| 29 | if hasattr(labelled_dataset, 'data'): |
| 30 | if isinstance(labelled_dataset.data, list): |
| 31 | self.data = _labelled_dataset.data + _unlabelled_dataset.data |
| 32 | if hasattr(_labelled_dataset, 'target'): |
| 33 | self.target = _labelled_dataset.target + _unlabelled_dataset.target |
| 34 | if hasattr(_labelled_dataset, 'uq_idxs'): |
| 35 | self.uq_idxs = _labelled_dataset.uq_idxs.tolist() + _unlabelled_dataset.uq_idxs.tolist() |
| 36 | elif _labelled_dataset.data.shape[1] == _unlabelled_dataset.data.shape[1]: |
| 37 | self.data = np.concatenate((_labelled_dataset.data, _unlabelled_dataset.data), axis=0) |
| 38 | else: |
| 39 | assert False, f'size not match: {_labelled_dataset.data.shape[1]} and {_unlabelled_dataset.data.shape[1]}' |
| 40 | elif hasattr(labelled_dataset, 'samples'): |
| 41 | self.data = _labelled_dataset.samples + _unlabelled_dataset.samples |
| 42 | if hasattr(_labelled_dataset, 'target'): |
| 43 | self.target = _labelled_dataset.target + _unlabelled_dataset.target |
| 44 | if hasattr(_labelled_dataset, 'uq_idxs'): |
| 45 | self.uq_idxs = _labelled_dataset.uq_idxs.tolist() + _unlabelled_dataset.uq_idxs.tolist() |
| 46 | else: |
| 47 | assert False, f'Unsuport {labelled_dataset}' |
| 48 | |
| 49 | def __getitem__(self, item): |
| 50 | if self.unlabelled_dataset is None: |
| 51 | _tuple = self.labelled_dataset[item] |
| 52 | if len(_tuple) > 3: |
| 53 | img, label, uq_idx, attr = _tuple |
| 54 | labeled_or_not = 1 |
| 55 | return img, label, uq_idx, np.array([labeled_or_not]), attr |
| 56 | else: |
| 57 | img, label, uq_idx = _tuple |
| 58 | labeled_or_not = 1 |
| 59 | return img, label, uq_idx, np.array([labeled_or_not]), np.array([0]) |
| 60 | else: |
| 61 | if item < len(self.labelled_dataset): |
| 62 | _tuple = self.labelled_dataset[item] |
| 63 | if len(_tuple) > 3: |
| 64 | img, label, uq_idx, attr = _tuple |
| 65 | labeled_or_not = 1 |
| 66 | return img, label, uq_idx, np.array([labeled_or_not]), attr |
| 67 | else: |
| 68 | img, label, uq_idx = _tuple |
| 69 | labeled_or_not = 1 |
| 70 | return img, label, uq_idx, np.array([labeled_or_not]), np.array([0]) |
no outgoing calls
no test coverage detected