Set prarllel context in pipeline training process
(args_opt)
| 485 | |
| 486 | |
| 487 | def set_pipeline_parallel_context(args_opt): |
| 488 | """ |
| 489 | Set prarllel context in pipeline training process |
| 490 | """ |
| 491 | D.init() |
| 492 | device_num = D.get_group_size() |
| 493 | rank_id = D.get_rank() |
| 494 | print("rank_id is {}, device_num is {}".format(rank_id, device_num)) |
| 495 | context.reset_auto_parallel_context() |
| 496 | context.set_auto_parallel_context( |
| 497 | parallel_mode=ParallelMode.SEMI_AUTO_PARALLEL, gradients_mean=False, |
| 498 | full_batch=bool(args_opt.full_batch), loss_repeated_mean=True, |
| 499 | device_num=device_num, enable_parallel_optimizer=bool(args_opt.optimizer_shard), |
| 500 | pipeline_stages=args_opt.stage_num) |
| 501 | set_algo_parameters(elementwise_op_strategy_follow=True) |
| 502 | _set_multi_subgraphs() |
| 503 | return rank_id, device_num |
| 504 | |
| 505 | |
| 506 | def run_train_pipeline(args_opt): |
no outgoing calls
no test coverage detected