| 813 | |
| 814 | # Get parse |
| 815 | def get_arg_parser(): |
| 816 | parser = argparse.ArgumentParser() |
| 817 | parser.add_argument('--data_location', |
| 818 | help='Full path of train data', |
| 819 | required=False, |
| 820 | default='./data') |
| 821 | parser.add_argument('--steps', |
| 822 | help='set the number of steps on train dataset', |
| 823 | type=int, |
| 824 | default=0) |
| 825 | parser.add_argument('--batch_size', |
| 826 | help='Batch size to train. Default is 512', |
| 827 | type=int, |
| 828 | default=2048) |
| 829 | parser.add_argument('--output_dir', |
| 830 | help='Full path to model output directory. \ |
| 831 | Default to ./result. Covered by --checkpoint. ', |
| 832 | required=False, |
| 833 | default='./result') |
| 834 | parser.add_argument('--checkpoint', |
| 835 | help='Full path to checkpoints input/output. \ |
| 836 | Default to ./result/$MODEL_TIMESTAMP', |
| 837 | required=False) |
| 838 | parser.add_argument('--save_steps', |
| 839 | help='set the number of steps on saving checkpoints', |
| 840 | type=int, |
| 841 | default=0) |
| 842 | parser.add_argument('--seed', |
| 843 | help='set the random seed for tensorflow', |
| 844 | type=int, |
| 845 | default=2021) |
| 846 | parser.add_argument('--optimizer', |
| 847 | type=str, \ |
| 848 | choices=['adam', 'adamasync', 'adagraddecay'], |
| 849 | default='adam') |
| 850 | parser.add_argument('--learning_rate', |
| 851 | help='Learning rate for model', |
| 852 | type=float, |
| 853 | default=0.01) |
| 854 | parser.add_argument('--deep_dropout', |
| 855 | help='Dropout regularization for deep model', |
| 856 | type=float, |
| 857 | default=0.0) |
| 858 | parser.add_argument('--keep_checkpoint_max', |
| 859 | help='Maximum number of recent checkpoint to keep', |
| 860 | type=int, |
| 861 | default=1) |
| 862 | parser.add_argument('--timeline', |
| 863 | help='number of steps on saving timeline. Default 0', |
| 864 | type=int, |
| 865 | default=0) |
| 866 | parser.add_argument('--protocol', |
| 867 | type=str, |
| 868 | choices=['grpc', 'grpc++', 'star_server'], |
| 869 | default='grpc') |
| 870 | parser.add_argument('--inter', |
| 871 | help='set inter op parallelism threads.', |
| 872 | type=int, |