MCPcopy Create free account
hub / github.com/breeze-sys/Label-Only-MIA-Go / load_dataset

Function load_dataset

python_server/utils.py:19–104  ·  view source on GitHub ↗
(args, dataset, cluster=None, mode = 'target', max_num = 2000)

Source from the content-addressed store, hash-verified

17
18
19def load_dataset(args, dataset, cluster=None, mode = 'target', max_num = 2000):
20 kwargs = {'num_workers': 2, 'pin_memory': True}
21 # load trainset and testset
22
23 if mode == 'shadow' or mode == 'ChangeDataSize':
24 if dataset == 'GTSRB':
25 transform = transforms.Compose([Rand_Augment(), transforms.Resize((64,64)), transforms.ToTensor()])
26 else:
27 transform = transforms.Compose([Rand_Augment(), transforms.ToTensor()])
28 else:
29 if dataset == 'GTSRB':
30 transform = transforms.Compose([transforms.Resize((64,64)), transforms.ToTensor()])
31 else:
32 transform = transforms.Compose([transforms.ToTensor()])
33
34 if dataset == 'CIFAR10':
35 whole_set = datasets.CIFAR10('data', train=True, download=True, transform=transform)
36 max_cluster = 3000
37 test_size = 1000
38 elif dataset == 'CIFAR100':
39 whole_set = datasets.CIFAR100('data', train=True, download=True, transform=transform)
40 max_cluster = 7000
41 test_size = 1000
42 elif dataset == 'GTSRB':
43 whole_set = datasets.ImageFolder('data/GTSRB/', transform= transform)
44 max_cluster = 600
45 test_size = 500
46 elif dataset == 'Face':
47 whole_set = datasets.ImageFolder('data/lfw/', transform=transform)
48 max_cluster = 350
49 test_size = 100
50 # elif dataset == 'TinyImageNet':
51 # whole_set = datasets.ImageFolder('data/tiny-imagenet-200/train', transform=transform)
52 # max_cluster = 30000
53 # test_size = 2000
54 length = len(whole_set)
55 if mode == 'target':
56 train_size = cluster
57 remain_size = length - train_size - test_size
58 train_set, _, test_set = dataset_split(whole_set, [train_size, remain_size, test_size])
59 train_loader = DataLoader(train_set, batch_size=args.batch_size, shuffle=False, **kwargs)
60 test_loader = DataLoader(test_set, batch_size=args.batch_size, shuffle=False, **kwargs)
61 return train_loader, test_loader
62 elif mode == 'shadow':
63 train_size = length - max_cluster - test_size
64 _, train_set, test_set = dataset_split(whole_set, [max_cluster, train_size, test_size])
65 train_loader = DataLoader(train_set, batch_size=args.batch_size, shuffle=False, **kwargs)
66 #test_loader = DataLoader(test_set, batch_size=args.batch_size, shuffle=False, **kwargs)
67 return train_loader#, test_loader
68 elif mode == 'salem_unknown':
69 train_size = length - max_cluster - test_size
70 salme_train = int(train_size * 0.5)
71 salme_test = train_size - salme_train
72 _, train_set, test_set, _ = dataset_split(whole_set, [max_cluster, salme_train, salme_test, test_size])
73 train_loader = DataLoader(train_set, batch_size=args.batch_size, shuffle=False, **kwargs)
74 test_loader = DataLoader(test_set, batch_size=args.batch_size, shuffle=False, **kwargs)
75 return train_loader, test_loader
76 elif mode == 'salem_known':

Callers

nothing calls this directly

Calls 2

Rand_AugmentClass · 0.85
dataset_splitFunction · 0.85

Tested by

no test coverage detected