MCPcopy Create free account
hub / github.com/M-Nauta/ProtoTree / get_args

Function get_args

util/args.py:13–142  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

11
12"""
13def 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')

Callers 3

run_ensembleFunction · 0.90
run_treeFunction · 0.90
main_tree.pyFile · 0.90

Calls 1

get_milestonesFunction · 0.85

Tested by

no test coverage detected