| 255 | |
| 256 | |
| 257 | def get_arg_parser(): |
| 258 | parser = argparse.ArgumentParser() |
| 259 | parser.add_argument('--data_location', |
| 260 | help='Full path of train data', |
| 261 | required=False, |
| 262 | default='./data') |
| 263 | parser.add_argument('--steps', |
| 264 | help='set the number of steps on train dataset', |
| 265 | type=int, |
| 266 | default=0) |
| 267 | parser.add_argument('--batch_size', |
| 268 | help='Batch size to train. Default is 512', |
| 269 | type=int, |
| 270 | default=512) |
| 271 | parser.add_argument('--output_dir', |
| 272 | help='Full path to logs & model output directory', |
| 273 | required=False, |
| 274 | default='./result') |
| 275 | parser.add_argument('--checkpoint', |
| 276 | help='Full path to checkpoints input/output directory', |
| 277 | required=False) |
| 278 | parser.add_argument('--deep_dropout', |
| 279 | help='Dropout regularization for deep model', |
| 280 | type=float, |
| 281 | default=0.0) |
| 282 | parser.add_argument("--optimizer", |
| 283 | type=str, |
| 284 | choices=["adam", "adagrad", "adamasync"], |
| 285 | default="adam") |
| 286 | parser.add_argument('--learning_rate', |
| 287 | help='Learning rate for model', |
| 288 | type=float, |
| 289 | default=0.001) |
| 290 | parser.add_argument('--save_steps', |
| 291 | help='set the number of steps on saving checkpoints', |
| 292 | type=int, |
| 293 | default=0) |
| 294 | parser.add_argument('--keep_checkpoint_max', |
| 295 | help='Maximum number of recent checkpoint to keep', |
| 296 | type=int, |
| 297 | default=1) |
| 298 | parser.add_argument('--bf16', |
| 299 | help='enable DeepRec BF16 in deep model. Default FP32', |
| 300 | action='store_true') |
| 301 | parser.add_argument('--no_eval', |
| 302 | help='not evaluate trained model by eval dataset.', |
| 303 | action='store_true') |
| 304 | parser.add_argument('--timeline', |
| 305 | help='number of steps on saving timeline. Default 0', |
| 306 | type=int, |
| 307 | default=0) |
| 308 | parser.add_argument("--protocol", |
| 309 | type=str, |
| 310 | choices=["grpc", "grpc++", "star_server"], |
| 311 | default="grpc") |
| 312 | parser.add_argument('--inter', |
| 313 | help='set inter op parallelism threads.', |
| 314 | type=int, |