()
| 11 | |
| 12 | """ |
| 13 | def get_args() -> argparse.Namespace: |
| 14 | |
| 15 | parser = argparse.ArgumentParser('Train a ProtoTree') |
| 16 | parser.add_argument('--dataset', |
| 17 | type=str, |
| 18 | default='CUB-200-2011', |
| 19 | help='Data set on which the ProtoTree should be trained') |
| 20 | parser.add_argument('--net', |
| 21 | type=str, |
| 22 | default='resnet50_inat', |
| 23 | help='Base network used in the tree. Pretrained network on iNaturalist is only available for resnet50_inat (default). Others are pretrained on ImageNet. Options are: resnet18, resnet34, resnet50, resnet50_inat, resnet101, resnet152, densenet121, densenet169, densenet201, densenet161, vgg11, vgg13, vgg16, vgg19, vgg11_bn, vgg13_bn, vgg16_bn or vgg19_bn') |
| 24 | parser.add_argument('--batch_size', |
| 25 | type=int, |
| 26 | default=64, |
| 27 | help='Batch size when training the model using minibatch gradient descent') |
| 28 | parser.add_argument('--depth', |
| 29 | type=int, |
| 30 | default=9, |
| 31 | help='The tree is initialized as a complete tree of this depth') |
| 32 | parser.add_argument('--epochs', |
| 33 | type=int, |
| 34 | default=100, |
| 35 | help='The number of epochs the tree should be trained') |
| 36 | parser.add_argument('--optimizer', |
| 37 | type=str, |
| 38 | default='AdamW', |
| 39 | help='The optimizer that should be used when training the tree') |
| 40 | parser.add_argument('--lr', |
| 41 | type=float, |
| 42 | default=0.001, |
| 43 | help='The optimizer learning rate for training the prototypes') |
| 44 | parser.add_argument('--lr_block', |
| 45 | type=float, |
| 46 | default=0.001, |
| 47 | help='The optimizer learning rate for training the 1x1 conv layer and last conv layer of the underlying neural network (applicable to resnet50 and densenet121)') |
| 48 | parser.add_argument('--lr_net', |
| 49 | type=float, |
| 50 | default=1e-5, |
| 51 | help='The optimizer learning rate for the underlying neural network') |
| 52 | parser.add_argument('--lr_pi', |
| 53 | type=float, |
| 54 | default=0.001, |
| 55 | help='The optimizer learning rate for the leaf distributions (only used if disable_derivative_free_leaf_optim flag is set') |
| 56 | parser.add_argument('--momentum', |
| 57 | type=float, |
| 58 | default=0.9, |
| 59 | help='The optimizer momentum parameter (only applicable to SGD)') |
| 60 | parser.add_argument('--weight_decay', |
| 61 | type=float, |
| 62 | default=0.0, |
| 63 | help='Weight decay used in the optimizer') |
| 64 | parser.add_argument('--disable_cuda', |
| 65 | action='store_true', |
| 66 | help='Flag that disables GPU usage if set') |
| 67 | parser.add_argument('--log_dir', |
| 68 | type=str, |
| 69 | default='./runs/run_prototree', |
| 70 | help='The directory in which train progress should be logged') |
no test coverage detected