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

Function main

generate_robust_em.py:84–235  ·  view source on GitHub ↗

init model / optim / loss func

(args, logger)

Source from the content-addressed store, hash-verified

82
83
84def 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:

Callers 1

Calls 5

perturbMethod · 0.95
perturbMethod · 0.95
load_pretrained_modelFunction · 0.85
save_checkpointFunction · 0.70
regenerate_def_noiseFunction · 0.70

Tested by

no test coverage detected