| 765 | return low_string == 'true' |
| 766 | |
| 767 | def get_arg_parser(): |
| 768 | parser = argparse.ArgumentParser(formatter_class=argparse.ArgumentDefaultsHelpFormatter) |
| 769 | parser.add_argument('--seed', help='set random seed', type=int, default=2020) |
| 770 | parser.add_argument('--data_location', |
| 771 | help='Full path of train data', |
| 772 | default='./data') |
| 773 | parser.add_argument('--steps', |
| 774 | help='set the number of steps on train dataset', |
| 775 | type=int, |
| 776 | default=0) |
| 777 | parser.add_argument('--batch_size', |
| 778 | help='Batch size to train', |
| 779 | type=int, |
| 780 | default=2048) |
| 781 | parser.add_argument('--output_dir', |
| 782 | help='Full path to logs & model output directory', |
| 783 | default='./result') |
| 784 | parser.add_argument('--checkpoint_dir', |
| 785 | help='Full path to checkpoints output directory') |
| 786 | parser.add_argument('--learning_rate', |
| 787 | help='Learning rate for model', |
| 788 | type=float, |
| 789 | default=0.1) |
| 790 | parser.add_argument('--l2_regularization', |
| 791 | help='L2 regularization for the model', |
| 792 | type=float) |
| 793 | parser.add_argument('--timeline', |
| 794 | help='number of steps on saving timeline', |
| 795 | type=int) |
| 796 | parser.add_argument('--save_steps', |
| 797 | help='set the number of steps on saving checkpoints', |
| 798 | type=int) |
| 799 | parser.add_argument('--keep_checkpoint_max', |
| 800 | help='Maximum number of recent checkpoint to keep', |
| 801 | type=int, |
| 802 | default=1) |
| 803 | parser.add_argument('--bf16', |
| 804 | help='enable DeepRec BF16 in deep model', |
| 805 | action='store_true') |
| 806 | parser.add_argument('--no_eval', |
| 807 | help='not evaluate trained model by eval dataset', |
| 808 | action='store_true') |
| 809 | parser.add_argument('--protocol', |
| 810 | type=str, |
| 811 | choices=['grpc', 'grpc++', 'star_server'], |
| 812 | default='grpc') |
| 813 | parser.add_argument('--inter', |
| 814 | help='set inter op parallelism threads', |
| 815 | type=int, |
| 816 | default=0) |
| 817 | parser.add_argument('--intra', |
| 818 | help='set intra op parallelism threads', |
| 819 | type=int, |
| 820 | default=0) |
| 821 | parser.add_argument('--input_layer_partitioner', |
| 822 | help='slice size of input layer partitioner. units MB', |
| 823 | type=int, |
| 824 | default=0) |