(config_dict, client_id=-1, n_clients=50, alpha=1e0, bsize=16,
linear_eval=False, hparam_eval=False, in_simulation=False, force_shuffle=False,
subset_proportion=1, subset_force_class_balanced=False, subset_seed=0)
| 302 | |
| 303 | ######### Dataloaders ######### |
| 304 | def load_data(config_dict, client_id=-1, n_clients=50, alpha=1e0, bsize=16, |
| 305 | linear_eval=False, hparam_eval=False, in_simulation=False, force_shuffle=False, |
| 306 | subset_proportion=1, subset_force_class_balanced=False, subset_seed=0): |
| 307 | |
| 308 | da_method = config_dict["da_method"] |
| 309 | train_mode = config_dict["train_mode"] |
| 310 | dataset_name = config_dict["dataset"] |
| 311 | data_dir = config_dict['data_dir'] |
| 312 | |
| 313 | # Define data augmentations |
| 314 | if(hparam_eval): |
| 315 | transform_train = SimCLRTransform(is_sup=False, image_size=32) |
| 316 | elif(linear_eval): |
| 317 | transform_train = BaseTransform(is_sup=True, image_size=32) |
| 318 | elif(da_method=="sup"): |
| 319 | transform_train = BaseTransform(is_sup=(train_mode=="sup"), image_size=32) |
| 320 | elif(da_method=="simclr" or da_method=="simsiam"): |
| 321 | transform_train = SimCLRTransform(is_sup=(train_mode=="sup"), image_size=32) |
| 322 | elif(da_method=="specloss"): |
| 323 | transform_train = SpecLossTransform(is_sup=(train_mode=="sup"), image_size=32) |
| 324 | elif(da_method=="byol"): |
| 325 | transform_train = BYOLTransform(is_sup=(train_mode=="sup"), image_size=32) |
| 326 | elif(da_method=="rotpred"): |
| 327 | transform_train = RotTransform(is_sup=(train_mode=="sup")) |
| 328 | elif(da_method=="orchestra"): |
| 329 | transform_train = OrchestraTransform(is_sup=(train_mode=="sup"), image_size=32) |
| 330 | |
| 331 | transform_test = T.Compose([ |
| 332 | T.ToTensor(), |
| 333 | T.Normalize((0.5, 0.5, 0.5), (0.5, 0.5, 0.5))] |
| 334 | ) |
| 335 | |
| 336 | # Load dataset |
| 337 | if(dataset_name=="CIFAR10"): |
| 338 | trainset = CIFAR10(f"{data_dir}/dataset/CIFAR10", train=True, download=False, transform=transform_train) |
| 339 | memset = CIFAR10(f"{data_dir}/dataset/CIFAR10", train=True, download=False, transform=transform_test) |
| 340 | testset = CIFAR10(f"{data_dir}/dataset/CIFAR10", train=False, download=False, transform=transform_test) |
| 341 | elif(dataset_name=="CIFAR100"): |
| 342 | trainset = CIFAR100(f"{data_dir}/dataset/CIFAR100", train=True, download=False, transform=transform_train) |
| 343 | memset = CIFAR100(f"{data_dir}/dataset/CIFAR100", train=True, download=False, transform=transform_test) |
| 344 | testset = CIFAR100(f"{data_dir}/dataset/CIFAR100", train=False, download=False, transform=transform_test) |
| 345 | else: |
| 346 | raise Exception("Dataset not recognized") |
| 347 | |
| 348 | # Dataloaders for given client |
| 349 | if(client_id > -1): |
| 350 | with open(f'{data_dir}/{n_clients}/{alpha}/{dataset_name}/train/' +dataset_name+"_"+str(client_id)+".pkl", "rb") as f: |
| 351 | train_ids = pkl.load(f).astype(np.int32) |
| 352 | with open(f'{data_dir}/{n_clients}/{alpha}/{dataset_name}/test/'+dataset_name+"_"+str(client_id)+".pkl", "rb") as f: |
| 353 | test_ids = pkl.load(f).astype(np.int32) |
| 354 | # Sanity check |
| 355 | train_deets, test_deets = np.unique(np.array(trainset.targets)[train_ids], return_counts=True), np.unique(np.array(testset.targets)[test_ids], return_counts=True) |
| 356 | |
| 357 | trainloader = DataLoader(clientDataset(trainset, train_ids), batch_size=bsize, shuffle=True, drop_last=True) |
| 358 | memloader = DataLoader(clientDataset(memset, train_ids), batch_size=bsize, shuffle=True, drop_last=True) |
| 359 | testloader = DataLoader(clientDataset(testset, test_ids), batch_size=bsize, shuffle=False, drop_last=True) |
| 360 | |
| 361 | # Sanity check |
nothing calls this directly
no test coverage detected
searching dependent graphs…