()
| 843 | |
| 844 | |
| 845 | def _get_sequence_data_parallel_group(): |
| 846 | global mpu |
| 847 | # When sequence parallelism is enabled, the process group for zero sharding and |
| 848 | # gradient allreduce must be across both dimensions of data and sequence parallelism. |
| 849 | if mpu is not None and hasattr(mpu, 'get_sequence_data_parallel_group'): |
| 850 | return mpu.get_sequence_data_parallel_group() |
| 851 | return _get_data_parallel_group() |
| 852 | |
| 853 | |
| 854 | def _get_expert_model_parallel_world_size(): |
no test coverage detected