| 6 | |
| 7 | |
| 8 | def 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 | |
| 43 | def prepare_model(args): |
| 44 | if args.method == 'protonet': |