| 745 | |
| 746 | # Get parse |
| 747 | def get_arg_parser(): |
| 748 | parser = argparse.ArgumentParser() |
| 749 | parser.add_argument('--data_location', |
| 750 | help='Full path of train data', |
| 751 | required=False, |
| 752 | default='./data') |
| 753 | parser.add_argument('--steps', |
| 754 | help='set the number of steps on train dataset', |
| 755 | type=int, |
| 756 | default=0) |
| 757 | parser.add_argument('--batch_size', |
| 758 | help='Batch size to train. Default is 512', |
| 759 | type=int, |
| 760 | default=2048) |
| 761 | parser.add_argument('--output_dir', |
| 762 | help='Full path to model output directory. \ |
| 763 | Default to ./result. Covered by --checkpoint. ', |
| 764 | required=False, |
| 765 | default='./result') |
| 766 | parser.add_argument('--checkpoint', |
| 767 | help='Full path to checkpoints input/output. \ |
| 768 | Default to ./result/$MODEL_TIMESTAMP', |
| 769 | required=False) |
| 770 | parser.add_argument('--save_steps', |
| 771 | help='set the number of steps on saving checkpoints', |
| 772 | type=int, |
| 773 | default=0) |
| 774 | parser.add_argument('--seed', |
| 775 | help='set the random seed for tensorflow', |
| 776 | type=int, |
| 777 | default=2021) |
| 778 | parser.add_argument('--optimizer', |
| 779 | type=str, \ |
| 780 | choices=['adam', 'adamasync', 'adagraddecay', 'adagrad'], |
| 781 | default='adamasync') |
| 782 | parser.add_argument('--linear_learning_rate', |
| 783 | help='Learning rate for linear model', |
| 784 | type=float, |
| 785 | default=0.2) |
| 786 | parser.add_argument('--deep_learning_rate', |
| 787 | help='Learning rate for deep model', |
| 788 | type=float, |
| 789 | default=0.01) |
| 790 | parser.add_argument('--keep_checkpoint_max', |
| 791 | help='Maximum number of recent checkpoint to keep', |
| 792 | type=int, |
| 793 | default=1) |
| 794 | parser.add_argument('--timeline', |
| 795 | help='number of steps on saving timeline. Default 0', |
| 796 | type=int, |
| 797 | default=0) |
| 798 | parser.add_argument('--protocol', |
| 799 | type=str, |
| 800 | choices=['grpc', 'grpc++', 'star_server'], |
| 801 | default='grpc') |
| 802 | parser.add_argument('--inter', |
| 803 | help='set inter op parallelism threads.', |
| 804 | type=int, |