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)
| 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 |
no test coverage detected