MCPcopy Create free account
hub / github.com/LAMDA-CL/CVPR22-Fact / get_command_line_parser

Function get_command_line_parser

train.py:9–55  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

7PROJECT='base'
8
9def get_command_line_parser():
10 parser = argparse.ArgumentParser()
11
12 # about dataset and network
13 parser.add_argument('-project', type=str, default=PROJECT)
14 parser.add_argument('-dataset', type=str, default='cub200',
15 choices=['mini_imagenet', 'cub200', 'cifar100'])
16 parser.add_argument('-dataroot', type=str, default=DATA_DIR)
17
18 # about pre-training
19 parser.add_argument('-epochs_base', type=int, default=100)
20 parser.add_argument('-epochs_new', type=int, default=100)
21 parser.add_argument('-lr_base', type=float, default=0.1)
22 parser.add_argument('-lr_new', type=float, default=0.1)
23 parser.add_argument('-schedule', type=str, default='Step',
24 choices=['Step', 'Milestone','Cosine'])
25 parser.add_argument('-milestones', nargs='+', type=int, default=[60, 70])
26 parser.add_argument('-step', type=int, default=20)
27 parser.add_argument('-decay', type=float, default=0.0005)
28 parser.add_argument('-momentum', type=float, default=0.9)
29 parser.add_argument('-gamma', type=float, default=0.1)
30 parser.add_argument('-temperature', type=float, default=16)
31 parser.add_argument('-not_data_init', action='store_true', help='using average data embedding to init or not')
32 parser.add_argument('-batch_size_base', type=int, default=128)
33 parser.add_argument('-batch_size_new', type=int, default=0, help='set 0 will use all the availiable training image for new')
34 parser.add_argument('-test_batch_size', type=int, default=100)
35 parser.add_argument('-base_mode', type=str, default='ft_cos',
36 choices=['ft_dot', 'ft_cos']) # ft_dot means using linear classifier, ft_cos means using cosine classifier
37 parser.add_argument('-new_mode', type=str, default='avg_cos',
38 choices=['ft_dot', 'ft_cos', 'avg_cos']) # ft_dot means using linear classifier, ft_cos means using cosine classifier, avg_cos means using average data embedding and cosine classifier
39
40 #for fact
41 parser.add_argument('-balance', type=float, default=1.0)
42 parser.add_argument('-loss_iter', type=int, default=200)
43 parser.add_argument('-alpha', type=float, default=2.0)
44 parser.add_argument('-eta', type=float, default=0.1)
45
46 parser.add_argument('-start_session', type=int, default=0)
47 parser.add_argument('-model_dir', type=str, default=MODEL_DIR, help='loading model parameter from a specific dir')
48 parser.add_argument('-set_no_val', action='store_true', help='set validation using test set or no validation')
49
50 # about training
51 parser.add_argument('-gpu', default='0,1,2,3')
52 parser.add_argument('-num_workers', type=int, default=8)
53 parser.add_argument('-seed', type=int, default=1)
54 parser.add_argument('-debug', action='store_true')
55 return parser
56
57
58if __name__ == '__main__':

Callers 1

train.pyFile · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected