MCPcopy Create free account
hub / github.com/HobbitLong/PyContrast / TrainOptions

Class TrainOptions

pycontrast/options/train_options.py:6–67  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

4
5
6class TrainOptions(BaseOptions):
7
8 def initialize(self, parser):
9 parser = BaseOptions.initialize(self, parser)
10
11 parser.add_argument('--aug', default='A', type=str,
12 help='data augmentation for training')
13 parser.add_argument('--beta', type=float, default=0.5,
14 help='balance between Jigsaw and InsDis')
15 parser.add_argument('--warm', action='store_true',
16 help='add warm-up setting')
17 parser.add_argument('--amp', action='store_true',
18 help='using mixed precision')
19 parser.add_argument('--opt_level', type=str, default='O2',
20 choices=['O1', 'O2'])
21
22 return parser
23
24 def modify_options(self, opt):
25 opt = self.override_options(opt)
26
27 iterations = opt.lr_decay_epochs.split(',')
28 opt.lr_decay_epochs = list([])
29 for it in iterations:
30 opt.lr_decay_epochs.append(int(it))
31
32 # set up saving name
33 opt.model_name = '{}_{}_{}_Jig_{}_{}_aug_{}_{}_{}'.format(
34 opt.method, opt.arch, opt.modal, opt.jigsaw, opt.mem,
35 opt.aug, opt.head, opt.nce_t
36 )
37 if opt.amp:
38 opt.model_name = '{}_amp_{}'.format(opt.model_name, opt.opt_level)
39 if opt.cosine:
40 opt.model_name = '{}_cosine'.format(opt.model_name)
41
42 # warm-up for large-batch training, e.g. 1024 with multiple nodes
43 if opt.batch_size > 256:
44 opt.warm = True
45 if opt.warm:
46 opt.model_name = '{}_warm'.format(opt.model_name)
47 opt.warmup_from = 0.01
48 if opt.epochs > 500:
49 opt.warm_epochs = 10
50 else:
51 opt.warm_epochs = 5
52 if opt.cosine:
53 eta_min = opt.learning_rate * (opt.lr_decay_rate ** 3)
54 opt.warmup_to = eta_min + (opt.learning_rate - eta_min) * (
55 1 + math.cos(math.pi * opt.warm_epochs / opt.epochs)) / 2
56 else:
57 opt.warmup_to = opt.learning_rate
58
59 # create folders
60 opt.model_folder = os.path.join(opt.model_path, opt.model_name)
61 if not os.path.isdir(opt.model_folder):
62 os.makedirs(opt.model_folder)
63 opt.tb_folder = os.path.join(opt.tb_path, opt.model_name)

Callers 1

mainFunction · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected