| 677 | |
| 678 | # Get parse |
| 679 | def get_arg_parser(): |
| 680 | parser = argparse.ArgumentParser() |
| 681 | parser.add_argument('--data_location', |
| 682 | help='Full path of train data', |
| 683 | required=False, |
| 684 | default='./data') |
| 685 | parser.add_argument('--steps', |
| 686 | help='set the number of steps on train dataset', |
| 687 | type=int, |
| 688 | default=0) |
| 689 | parser.add_argument('--batch_size', |
| 690 | help='Batch size to train. Default is 4096', |
| 691 | type=int, |
| 692 | default=2048) |
| 693 | parser.add_argument('--output_dir', |
| 694 | help='Full path to model output directory. \ |
| 695 | Default to ./result. Covered by --checkpoint. ', |
| 696 | required=False, |
| 697 | default='./result') |
| 698 | parser.add_argument('--checkpoint', |
| 699 | help='Full path to checkpoints input/output. \ |
| 700 | Default to ./result/$MODEL_TIMESTAMP', |
| 701 | required=False) |
| 702 | parser.add_argument('--save_steps', |
| 703 | help='set the number of steps on saving checkpoints', |
| 704 | type=int, |
| 705 | default=0) |
| 706 | parser.add_argument('--seed', |
| 707 | help='set the random seed for tensorflow', |
| 708 | type=int, |
| 709 | default=2021) |
| 710 | parser.add_argument('--optimizer', |
| 711 | type=str, \ |
| 712 | choices=['adam', 'adamasync', 'adagraddecay'], |
| 713 | default='adamasync') # TODO: change to adam or adamsync |
| 714 | parser.add_argument('--learning_rate', |
| 715 | help='Learning rate for model', |
| 716 | type=float, |
| 717 | default=0.001) |
| 718 | parser.add_argument('--deep_dropout', |
| 719 | help='Dropout regularization for deep model', |
| 720 | type=float, |
| 721 | default=0.0) |
| 722 | parser.add_argument('--keep_checkpoint_max', |
| 723 | help='Maximum number of recent checkpoint to keep', |
| 724 | type=int, |
| 725 | default=1) |
| 726 | parser.add_argument('--timeline', |
| 727 | help='number of steps on saving timeline. Default 0', |
| 728 | type=int, |
| 729 | default=0) |
| 730 | parser.add_argument('--protocol', |
| 731 | type=str, |
| 732 | choices=['grpc', 'grpc++', 'star_server'], |
| 733 | default='grpc') |
| 734 | parser.add_argument('--inter', |
| 735 | help='set inter op parallelism threads.', |
| 736 | type=int, |