MCPcopy Create free account
hub / github.com/akhilmathurs/orchestra / load_data

Function load_data

utils.py:304–378  ·  view source on GitHub ↗
(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)

Source from the content-addressed store, hash-verified

302
303######### Dataloaders #########
304def 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

Callers

nothing calls this directly

Calls 8

SimCLRTransformClass · 0.85
BaseTransformClass · 0.85
SpecLossTransformClass · 0.85
BYOLTransformClass · 0.85
RotTransformClass · 0.85
OrchestraTransformClass · 0.85
clientDatasetClass · 0.85
get_dataset_subsetFunction · 0.85

Tested by

no test coverage detected

Used in the wild real call sites across dependent graphs

searching dependent graphs…