MCPcopy Create free account
hub / github.com/TorchSSL/TorchSSL / get_ssl_dset

Method get_ssl_dset

datasets/ssl_dataset.py:280–336  ·  view source on GitHub ↗

get_ssl_dset split training samples into labeled and unlabeled samples. The labeled data is balanced samples over classes. Args: num_labels: number of labeled data. index: If index of np.array is given, labeled data is not randomly sampled, b

(self, num_labels, index=None, include_lb_to_ulb=True,
                     strong_transform=None, onehot=False)

Source from the content-addressed store, hash-verified

278 is_ulb, strong_transform, onehot)
279
280 def get_ssl_dset(self, num_labels, index=None, include_lb_to_ulb=True,
281 strong_transform=None, onehot=False):
282 """
283 get_ssl_dset split training samples into labeled and unlabeled samples.
284 The labeled data is balanced samples over classes.
285
286 Args:
287 num_labels: number of labeled data.
288 index: If index of np.array is given, labeled data is not randomly sampled, but use index for sampling.
289 include_lb_to_ulb: If True, consistency regularization is also computed for the labeled data.
290 strong_transform: list of strong transform (RandAugment in FixMatch)
291 onehot: If True, the target is converted into onehot vector.
292
293 Returns:
294 BasicDataset (for labeled data), BasicDataset (for unlabeld data)
295 """
296 # Supervised top line using all data as labeled data.
297 if self.alg == 'fullysupervised':
298 lb_data, lb_targets = self.get_data()
299 lb_dset = BasicDataset(self.alg, lb_data, lb_targets, self.num_classes,
300 self.transform, False, None, onehot)
301 return lb_dset, None
302
303 if self.name.upper() == 'STL10':
304 lb_data, lb_targets, ulb_data = self.get_data()
305 if include_lb_to_ulb:
306 ulb_data = np.concatenate([ulb_data, lb_data], axis=0)
307 lb_data, lb_targets, _ = sample_labeled_data(self.args, lb_data, lb_targets, num_labels, self.num_classes)
308 ulb_targets = None
309 else:
310 data, targets = self.get_data()
311 lb_data, lb_targets, ulb_data, ulb_targets = split_ssl_data(self.args, data, targets,
312 num_labels, self.num_classes,
313 index, include_lb_to_ulb)
314 # output the distribution of labeled data for remixmatch
315 count = [0 for _ in range(self.num_classes)]
316 for c in lb_targets:
317 count[c] += 1
318 dist = np.array(count, dtype=float)
319 dist = dist / dist.sum()
320 dist = dist.tolist()
321 out = {"distribution": dist}
322 output_file = r"./data_statistics/"
323 output_path = output_file + str(self.name) + '_' + str(num_labels) + '.json'
324 if not os.path.exists(output_file):
325 os.makedirs(output_file, exist_ok=True)
326 with open(output_path, 'w') as w:
327 json.dump(out, w)
328 # print(Counter(ulb_targets.tolist()))
329 lb_dset = BasicDataset(self.alg, lb_data, lb_targets, self.num_classes,
330 self.transform, False, None, onehot)
331
332 ulb_dset = BasicDataset(self.alg, ulb_data, ulb_targets, self.num_classes,
333 self.transform, True, strong_transform, onehot)
334 # print(lb_data.shape)
335 # print(ulb_data.shape)
336 return lb_dset, ulb_dset

Callers 13

main_workerFunction · 0.95
main_workerFunction · 0.95
main_workerFunction · 0.95
main_workerFunction · 0.95
main_workerFunction · 0.95
main_workerFunction · 0.95
main_workerFunction · 0.95
main_workerFunction · 0.95
main_workerFunction · 0.95
main_workerFunction · 0.95
main_workerFunction · 0.95
main_workerFunction · 0.95

Calls 4

get_dataMethod · 0.95
BasicDatasetClass · 0.85
sample_labeled_dataFunction · 0.85
split_ssl_dataFunction · 0.85

Tested by

no test coverage detected