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

Function main

generate_tap.py:72–162  ·  view source on GitHub ↗

init model / optim / loss func

(args, logger)

Source from the content-addressed store, hash-verified

70
71
72def main(args, logger):
73 ''' init model / optim / loss func '''
74 model = utils.get_arch(args.arch, args.dataset)
75 optim = utils.get_optim(
76 args.optim, model.parameters(),
77 lr=args.lr, weight_decay=args.weight_decay, momentum=args.momentum)
78 criterion = torch.nn.CrossEntropyLoss()
79
80 ''' get Tensor train loader '''
81 train_loader = utils.get_indexed_tensor_loader(
82 args.dataset, batch_size=args.batch_size, root=args.data_dir, train=True)
83
84 dataset = train_loader.loader.dataset
85 ascending = True
86 if args.targeted:
87 if args.dataset == 'cifar10': classes = 10
88 elif args.dataset == 'cifar100': classes = 100
89 elif args.dataset == 'imagenet-mini': classes = 100
90 else: raise ValueError
91 dataset = TargetedIndexedDataset(dataset, classes)
92 ascending = False
93
94 train_loader = utils.Loader(dataset, batch_size=args.batch_size, shuffle=False, drop_last=False)
95
96 ''' get train transforms '''
97 train_trans = utils.get_transforms(
98 args.dataset, train=True, is_tensor=True)
99
100 ''' get (original) test loader '''
101 test_loader = utils.get_indexed_loader(
102 args.dataset, batch_size=args.batch_size, root=args.data_dir, train=False)
103
104 if args.adv_type == 'robust-pgd':
105 defender = attacks.RobustPGDAttacker(
106 samp_num = args.samp_num,
107 trans = train_trans,
108 radius = args.pgd_radius,
109 steps = args.pgd_steps,
110 step_size = args.pgd_step_size,
111 random_start = args.pgd_random_start,
112 ascending = ascending,
113 )
114 elif args.adv_type == 'diff-aug-pgd':
115 defender = attacks.DiffAugPGDAttacker(
116 samp_num = args.samp_num,
117 trans = train_trans,
118 radius = args.pgd_radius,
119 steps = args.pgd_steps,
120 step_size = args.pgd_step_size,
121 random_start = args.pgd_random_start,
122 ascending = ascending,
123 )
124 else: raise ValueError
125
126 ''' initialize the defensive noise (for unlearnable examples) '''
127 data_nums = len( train_loader.loader.dataset )
128 if args.dataset == 'cifar10' or args.dataset == 'cifar100':
129 def_noise = np.zeros([data_nums, 3, 32, 32], dtype=np.float16)

Callers 1

generate_tap.pyFile · 0.70

Calls 3

regenerate_def_noiseFunction · 0.70
save_checkpointFunction · 0.70

Tested by

no test coverage detected