(args, dataset, cluster=None, mode = 'target', max_num = 2000)
| 17 | |
| 18 | |
| 19 | def 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': |
nothing calls this directly
no test coverage detected