(self, labelled_dataset, unlabelled_dataset=None)
| 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: |
nothing calls this directly
no outgoing calls
no test coverage detected