init model / optim / loss func
(args, logger)
| 82 | |
| 83 | |
| 84 | def main(args, logger): |
| 85 | ''' init model / optim / loss func ''' |
| 86 | model = utils.get_arch(args.arch, args.dataset) |
| 87 | optim = utils.get_optim( |
| 88 | args.optim, model.parameters(), |
| 89 | lr=args.lr, weight_decay=args.weight_decay, momentum=args.momentum) |
| 90 | criterion = torch.nn.CrossEntropyLoss() |
| 91 | |
| 92 | ''' get Tensor train loader ''' |
| 93 | # trainset = utils.get_dataset(args.dataset, root=args.data_dir, train=True) |
| 94 | # trainset = utils.IndexedTensorDataset(trainset.x, trainset.y) |
| 95 | # train_loader = utils.Loader( |
| 96 | # trainset, batch_size=args.batch_size, shuffle=True, drop_last=True) |
| 97 | train_loader = utils.get_indexed_tensor_loader( |
| 98 | args.dataset, batch_size=args.batch_size, root=args.data_dir, train=True) |
| 99 | |
| 100 | ''' get train transforms ''' |
| 101 | train_trans = utils.get_transforms( |
| 102 | args.dataset, train=True, is_tensor=True) |
| 103 | |
| 104 | ''' get (original) test loader ''' |
| 105 | test_loader = utils.get_indexed_loader( |
| 106 | args.dataset, batch_size=args.batch_size, root=args.data_dir, train=False) |
| 107 | |
| 108 | defender = attacks.RobustMinimaxPGDDefender( |
| 109 | samp_num = args.samp_num, |
| 110 | trans = train_trans, |
| 111 | radius = args.pgd_radius, |
| 112 | steps = args.pgd_steps, |
| 113 | step_size = args.pgd_step_size, |
| 114 | random_start = args.pgd_random_start, |
| 115 | atk_radius = args.atk_pgd_radius, |
| 116 | atk_steps = args.atk_pgd_steps, |
| 117 | atk_step_size = args.atk_pgd_step_size, |
| 118 | atk_random_start = args.atk_pgd_random_start, |
| 119 | ) |
| 120 | |
| 121 | attacker = attacks.PGDAttacker( |
| 122 | radius = args.atk_pgd_radius, |
| 123 | steps = args.atk_pgd_steps, |
| 124 | step_size = args.atk_pgd_step_size, |
| 125 | random_start = args.atk_pgd_random_start, |
| 126 | norm_type = 'l-infty', |
| 127 | ascending = True, |
| 128 | ) |
| 129 | |
| 130 | ''' initialize the defensive noise (for unlearnable examples) ''' |
| 131 | data_nums = len( train_loader.loader.dataset ) |
| 132 | if args.dataset == 'cifar10' or args.dataset == 'cifar100': |
| 133 | # def_noise = np.zeros([data_nums, 3, 32, 32], dtype=np.float16) |
| 134 | def_noise = np.zeros([data_nums, 3, 32, 32], dtype=np.int8) |
| 135 | elif args.dataset == 'tiny-imagenet': |
| 136 | # def_noise = np.zeros([data_nums, 3, 64, 64], dtype=np.float16) |
| 137 | def_noise = np.zeros([data_nums, 3, 64, 64], dtype=np.int8) |
| 138 | elif args.dataset == 'imagenet-mini': |
| 139 | # def_noise = np.zeros([data_nums, 3, 256, 256], dtype=np.float16) |
| 140 | def_noise = np.zeros([data_nums, 3, 256, 256], dtype=np.int8) |
| 141 | else: |
no test coverage detected