(script)
| 5 | |
| 6 | |
| 7 | def parse_args(script): |
| 8 | parser = argparse.ArgumentParser(description='few-shot script %s' % (script)) |
| 9 | parser.add_argument('--dataset', default='miniImagenet', help='CUB/miniImagenet') |
| 10 | parser.add_argument('--model', default='WideResNet28_10', help='model: WideResNet28_10/ResNet{18}') |
| 11 | parser.add_argument('--method', default='S2M2_R', help='rotation/S2M2_R') |
| 12 | parser.add_argument('--train_aug', action='store_true', |
| 13 | help='perform data augmentation or not during training ') # still required for |
| 14 | # save_features.py and test.py to find the model path correctly |
| 15 | if script == 'train': |
| 16 | parser.add_argument('--num_classes', default=200, type=int, |
| 17 | help='total number of classes') # make it larger than the maximum label value in base class |
| 18 | parser.add_argument('--save_freq', default=10, type=int, help='Save frequency') |
| 19 | parser.add_argument('--start_epoch', default=0, type=int, help='Starting epoch') |
| 20 | parser.add_argument('--stop_epoch', default=400, type=int, |
| 21 | help='Stopping epoch') # for meta-learning methods, each epoch contains 100 episodes. |
| 22 | # The default epoch number is dataset dependent. See train.py |
| 23 | parser.add_argument('--resume', action='store_true', |
| 24 | help='continue from previous trained model with largest epoch') |
| 25 | parser.add_argument('--lr', default=0.001, type=int, help='learning rate') |
| 26 | parser.add_argument('--batch_size', default=16, type=int, help='batch size ') |
| 27 | parser.add_argument('--test_batch_size', default=2, type=int, help='batch size ') |
| 28 | parser.add_argument('--alpha', default=2.0, type=int, help='for S2M2 training ') |
| 29 | elif script == 'test': |
| 30 | parser.add_argument('--num_classes', default=200, type=int, help='total number of classes') |
| 31 | parser.add_argument('--model_dir', type=str, help='the pretrained model path ') |
| 32 | parser.add_argument('--file_name', type=str, help='where the features will be saved ') |
| 33 | parser.add_argument('--json_dir', type=str, default='./', help='') |
| 34 | |
| 35 | return parser.parse_args() |
| 36 | |
| 37 | |
| 38 | def get_assigned_file(checkpoint_dir, num): |
no outgoing calls
no test coverage detected