| 13 | |
| 14 | |
| 15 | def 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 | |