MCPcopy Create free account
hub / github.com/CausalLearning/robust-unlearnable-examples / main

Function main

generate_em.py:47–140  ·  view source on GitHub ↗

init model / optim / dataloader / loss func

(args, logger)

Source from the content-addressed store, hash-verified

45
46
47def main(args, logger):
48 ''' init model / optim / dataloader / loss func '''
49 model = utils.get_arch(args.arch, args.dataset)
50 optim = utils.get_optim(
51 args.optim, model.parameters(),
52 lr=args.lr, weight_decay=args.weight_decay, momentum=args.momentum)
53 train_loader = utils.get_indexed_loader(
54 args.dataset, batch_size=args.batch_size, root=args.data_dir, train=True)
55 test_loader = utils.get_indexed_loader(
56 args.dataset, batch_size=args.batch_size, root=args.data_dir, train=False)
57 criterion = torch.nn.CrossEntropyLoss()
58
59 defender = attacks.PGDAttacker(
60 radius = args.pgd_radius,
61 steps = args.pgd_steps,
62 step_size = args.pgd_step_size,
63 random_start = args.pgd_random_start,
64 norm_type = args.pgd_norm_type,
65 ascending = False,
66 )
67
68 if not args.cpu:
69 model.cuda()
70 criterion = criterion.cuda()
71
72 if args.parallel:
73 model = torch.nn.DataParallel(model)
74
75 log = dict()
76
77 ''' initialize the defensive noise (for unlearnable examples) '''
78 data_nums = len( train_loader.loader.dataset )
79 if args.dataset == 'cifar10' or args.dataset == 'cifar100':
80 def_noise = np.zeros([data_nums, 3, 32, 32], dtype=np.float16)
81 elif args.dataset == 'tiny-imagenet':
82 def_noise = np.zeros([data_nums, 3, 64, 64], dtype=np.float16)
83 elif args.dataset == 'imagenet-mini':
84 def_noise = np.zeros([data_nums, 3, 224, 224], dtype=np.float16)
85 else:
86 raise NotImplementedError
87
88 for step in range(args.train_steps):
89 lr = args.lr * (args.lr_decay_rate ** (step // args.lr_decay_freq))
90 for group in optim.param_groups:
91 group['lr'] = lr
92
93 x, y, ii = next(train_loader)
94 if not args.cpu:
95 x, y = x.cuda(), y.cuda()
96
97 if (step+1) % args.perturb_freq == 0:
98 def_x = defender.perturb(model, criterion, x, y)
99 def_noise[ii] = (def_x - x).cpu().numpy()
100
101 if args.cpu:
102 def_x = x + torch.tensor(def_noise[ii])
103 else:
104 def_x = x + torch.tensor(def_noise[ii]).cuda()

Callers 1

generate_em.pyFile · 0.70

Calls 3

perturbMethod · 0.95
save_checkpointFunction · 0.70
regenerate_def_noiseFunction · 0.70

Tested by

no test coverage detected