MCPcopy Create free account
hub / github.com/DanielShalam/BPA / get_args

Function get_args

train.py:15–112  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

13
14
15def get_args():
16 parser = argparse.ArgumentParser()
17
18 parser.add_argument('--seed', type=int, default=1,
19 help="""Random seed.""")
20
21 parser.add_argument('--root_path', type=str, default='./',
22 help=""" Path to project root directory. """)
23 parser.add_argument('--checkpoint_dir', type=str, default=None,
24 help=""" Where to save model checkpoints. If None, it will automatically created. """)
25 parser.add_argument('--dataset', type=str, default='miniimagenet',
26 choices=['miniimagenet', 'cifar'])
27 parser.add_argument('--data_path', type=str, default='./datasets/few_shot/miniimagenet',
28 help="""Path to dataset root directory.""")
29
30 parser.add_argument('--backbone', type=str, default='wrn',
31 help="""Define which backbone network to use. """)
32 parser.add_argument('--pretrained_path', type=str, default=False,
33 help=""" Path to pretrained model, used for testing/fine-tuning. """)
34
35 parser.add_argument('--eval', type=utils.bool_flag, default=False,
36 help=""" If true, make evaluation on the *test set*.
37 The amount of test episodes controlled by --test_episodes=<>""")
38 parser.add_argument('--eval_freq', type=int, default=1,
39 help=""" Evaluate training every n epochs. """)
40 parser.add_argument('--eval_first', type=utils.bool_flag, default=False,
41 help=""" Set to true to evaluate the model before training. Useful for fine-tuning. """)
42 parser.add_argument('--num_workers', type=int, default=8)
43
44 # wandb specific arguments
45 parser.add_argument('--wandb', type=utils.bool_flag, default=False,
46 help=""" Log data into wandb. """)
47 parser.add_argument('--project', type=str, default='BPA',
48 help=""" Project name in wandb. """)
49 parser.add_argument('--entity', type=str, default='',
50 help=""" Your wandb entity name. """)
51
52 # few-shot specific arguments
53 parser.add_argument('--method', type=str, default='pt_map_bpa',
54 choices=['proto', 'proto_bpa', 'pt_map', 'pt_map_bpa'],
55 help="""Specify which few-shot classifier to use.""")
56 parser.add_argument('--train_way', type=int, default=5,
57 help=""" Number of classes for each training task. """)
58 parser.add_argument('--val_way', type=int, default=5,
59 help=""" Number of classes for each validation/test task. """)
60 parser.add_argument('--num_shot', type=int, default=5,
61 help=""" Number of (labeled) support samples for each class. """)
62 parser.add_argument('--num_query', type=int, default=15,
63 help=""" Number of (un-labeled) query samples for each class. """)
64 parser.add_argument('--train_episodes', type=int, default=200,
65 help=""" Number of few-shot tasks for each epoch. """)
66 parser.add_argument('--eval_episodes', type=int, default=400,
67 help=""" Number of tasks to evaluate. """)
68 parser.add_argument('--test_episodes', type=int, default=10000,
69 help=""" Number of tasks to evaluate. """)
70 parser.add_argument('--temperature', type=float, default=1.,
71 help=""" Temperature for ProtoNet. """)
72

Callers 1

mainFunction · 0.70

Calls

no outgoing calls

Tested by

no test coverage detected