MCPcopy Create free account
hub / github.com/PaddlePaddle/Research / parse_args

Function parse_args

CV/PaddleFSL/utils.py:8–41  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

6
7
8def parse_args():
9 parser = argparse.ArgumentParser()
10 parser.add_argument('--dataset', default='miniimagenet', type=str, help='miniimagenet/omniglot/cifarfs/fc100/cub/tieredimagenet')
11 parser.add_argument('--backbone', default='Conv4', type=str, help='model: Conv4/Resnet12')
12 parser.add_argument('--method', default='protonet', type=str, help='protonet/relationnet')
13 parser.add_argument('--log_dir', default='./logs/', type=str, help='Directory where to write event logs and checkpoints')
14 parser.add_argument('--epochs', default=100, type=int, help='number of training epochs')
15 parser.add_argument('--episodes', default=1000, type=int, help='number of episodes per epoch')
16 parser.add_argument('--lr', default=0.001, type=float, help='learning rate')
17 parser.add_argument('--lr_scheduler', default=False, type=bool, help='if using learning rate scheduler')
18 parser.add_argument('--weight_decay', default=0.0005, type=float, help='weight decay')
19 parser.add_argument('--if_dropout', default=False, type=bool, help='if using dropout in backbone')
20 parser.add_argument('--n_way', default=5, type=int, help='number of classes in a task')
21 parser.add_argument('--k_shot', default=1, type=int, help='number of training sample per class')
22 parser.add_argument('--n_query', default=15, type=int, help='number of queries per class')
23 # belowings are currently not used
24 # parser.add_argument('--train_aug', default=False, type=bool, help='perform data augmentation or not during training')
25 # parser.add_argument('--meta_batch', default=1, type=int, help='number of meta batch')
26
27 parser.add_argument('--use_gpu', default=1, type=int, help='whether gpu is used')
28 # for backbones
29 parser.add_argument('--num_filters', default=64, type=int, help='number of queries per class')
30 parser.add_argument('--pooling_type', default='max', type=str, help='max/avg')
31 parser.add_argument('--resnet12_num_filters', nargs='+', default=[64,128,256,512], type=int, help='number of conv channels in ResNet12 backbone, eg. --resnet12_num_filters 64 128 256 512')
32
33 # testing
34 parser.add_argument('--test_mode', default=False, type=bool, help='if in test mode')
35 parser.add_argument('--test_episodes', default=600, type=int, help='number of testing episodes')
36
37 # for Protonet
38 parser.add_argument('--distance_metric', default='euclidean', type=str, help='euclidean/cosine distance for protonet')
39 parser.add_argument('--temperature', default=1.0, type=float, help='distance temperature for protonet')
40
41 return parser.parse_args()
42
43def prepare_model(args):
44 if args.method == 'protonet':

Callers 2

trainFunction · 0.90
testFunction · 0.90

Calls 1

parse_argsMethod · 0.45

Tested by 1

testFunction · 0.72