| 25 | |
| 26 | |
| 27 | def parse_args(): |
| 28 | parser = argparse.ArgumentParser(description="train mnist") |
| 29 | parser.add_argument( |
| 30 | "--log_path", type=str, required=True, help="dir to place tensorboard logs from all trials" |
| 31 | ) |
| 32 | parser.add_argument( |
| 33 | "--hidden_size_1", type=int, required=True, help="hidden size layer 1" |
| 34 | ) |
| 35 | parser.add_argument( |
| 36 | "--hidden_size_2", type=int, required=True, help="hidden size layer 2" |
| 37 | ) |
| 38 | parser.add_argument("--learning_rate", type=float, required=True, help="learning rate") |
| 39 | parser.add_argument("--epochs", type=int, required=True, help="number of epochs") |
| 40 | parser.add_argument("--dropout", type=float, required=True, help="dropout probability") |
| 41 | parser.add_argument("--batch_size", type=int, required=True, help="batch size") |
| 42 | return parser.parse_args() |
| 43 | |
| 44 | args = parse_args() |
| 45 | |