Return my rank for the tensor model parallel group.
()
| 274 | |
| 275 | |
| 276 | def get_tensor_model_parallel_rank(): |
| 277 | """Return my rank for the tensor model parallel group.""" |
| 278 | global _MPU_TENSOR_MODEL_PARALLEL_RANK |
| 279 | if _MPU_TENSOR_MODEL_PARALLEL_RANK is not None: |
| 280 | return _MPU_TENSOR_MODEL_PARALLEL_RANK |
| 281 | return torch.distributed.get_rank(group=get_tensor_model_parallel_group()) |
| 282 | |
| 283 | |
| 284 | def get_model_parallel_rank(): |
no test coverage detected