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

Class MergedDataset

data/data_utils.py:13–87  ·  view source on GitHub ↗

Takes two datasets (labelled_dataset, unlabelled_dataset) and merges them Allows you to iterate over them in parallel

Source from the content-addressed store, hash-verified

11 return subsample_indices
12import copy
13class 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])

Callers 5

G0_CUB200.pyFile · 0.90
get_datasetsFunction · 0.90
kmeans_subset.pyFile · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected