MCPcopy Create free account
hub / github.com/TPCD/DCCL / __init__

Method __init__

data/data_utils.py:20–47  ·  view source on GitHub ↗
(self, labelled_dataset, unlabelled_dataset=None)

Source from the content-addressed store, hash-verified

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:

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected